Skip to content

Commit 470ed43

Browse files
Copilotthinkall
andcommitted
Fix log_training_metric issue for time series models
- Modified _eval_estimator in ml.py to skip computing training metrics for TimeSeriesDataset - Added test case test_log_training_metric_ts_models to validate the fix Co-authored-by: thinkall <3197038+thinkall@users.noreply.github.com>
1 parent 4daa784 commit 470ed43

2 files changed

Lines changed: 47 additions & 9 deletions

File tree

‎flaml/automl/ml.py‎

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -616,15 +616,18 @@ def _eval_estimator(
616616
logger.warning(f"ValueError {e} happened in `metric_loss_score`, set `val_loss` to `np.inf`")
617617
metric_for_logging = {"pred_time": pred_time}
618618
if log_training_metric:
619-
train_pred_y = get_y_pred(estimator, X_train, eval_metric, task)
620-
metric_for_logging["train_loss"] = metric_loss_score(
621-
eval_metric,
622-
train_pred_y,
623-
y_train,
624-
labels,
625-
fit_kwargs.get("sample_weight"),
626-
fit_kwargs.get("groups"),
627-
)
619+
# For time series tasks, skip computing training metrics as predicting on training data
620+
# doesn't work the same way as for regular ML models (TimeSeriesDataset needs proper test_data)
621+
if not isinstance(X_train, TimeSeriesDataset):
622+
train_pred_y = get_y_pred(estimator, X_train, eval_metric, task)
623+
metric_for_logging["train_loss"] = metric_loss_score(
624+
eval_metric,
625+
train_pred_y,
626+
y_train,
627+
labels,
628+
fit_kwargs.get("sample_weight"),
629+
fit_kwargs.get("groups"),
630+
)
628631
else: # customized metric function
629632
val_loss, metric_for_logging = eval_metric(
630633
X_val,

‎test/automl/test_forecast.py‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -681,6 +681,41 @@ def split_by_date(df: pd.DataFrame, dt: datetime.date):
681681
print("yahoo!")
682682

683683

684+
def test_log_training_metric_ts_models():
685+
"""Test that log_training_metric=True works with time series models (arima, sarimax, holt-winters)."""
686+
import statsmodels.api as sm
687+
688+
# Prepare data
689+
data = sm.datasets.co2.load_pandas().data["co2"].resample("MS").mean()
690+
data = data.bfill().ffill().to_frame().reset_index().rename(columns={"index": "ds", "co2": "y"})
691+
num_samples = data.shape[0]
692+
time_horizon = 12
693+
split_idx = num_samples - time_horizon
694+
df = data[:split_idx]
695+
696+
# Test each time series model with log_training_metric=True
697+
for estimator in ["arima", "sarimax", "holt-winters"]:
698+
print(f"\nTesting {estimator} with log_training_metric=True")
699+
automl = AutoML()
700+
settings = {
701+
"time_budget": 5,
702+
"metric": "mape",
703+
"task": "ts_forecast",
704+
"eval_method": "holdout",
705+
"label": "y",
706+
"log_training_metric": True, # This should not cause errors
707+
"estimator_list": [estimator],
708+
}
709+
710+
try:
711+
automl.fit(dataframe=df, **settings, period=time_horizon)
712+
print(f" ✅ {estimator} SUCCESS with log_training_metric=True")
713+
assert automl.best_estimator == estimator
714+
except Exception as e:
715+
print(f" ❌ {estimator} FAILED: {type(e).__name__}: {e}")
716+
raise
717+
718+
684719
if __name__ == "__main__":
685720
# test_forecast_automl(60)
686721
# test_multivariate_forecast_num(5)

0 commit comments

Comments
 (0)