diff --git a/darts/models/forecasting/torch_forecasting_model.py b/darts/models/forecasting/torch_forecasting_model.py index d81e415902..b126e0df76 100644 --- a/darts/models/forecasting/torch_forecasting_model.py +++ b/darts/models/forecasting/torch_forecasting_model.py @@ -1279,7 +1279,19 @@ def fit_from_dataset( ------- self Fitted model. + + Notes + ----- + Encoders passed via ``add_encoders`` at model creation are ignored by this method. Encoders are only + applied when training with :func:`fit()`. If you need encoded covariates, the training datasets must + already contain them. """ + if self.add_encoders: + logger.warning( + "Encoders (`add_encoders`) are ignored when training with `fit_from_dataset()`. " + "Encoders are only applied when calling `fit()`. If you need encoded covariates, " + "make sure the training datasets already contain them." + ) self._train( **self._setup_for_train( train_dataset=train_dataset, @@ -2057,7 +2069,19 @@ def predict_from_dataset( ------- Sequence[TimeSeries] Returns one or more forecasts for time series. + + Notes + ----- + Encoders passed via ``add_encoders`` at model creation are ignored by this method. Encoders are only + applied when predicting with :func:`predict()`. If you need encoded covariates, the inference dataset must + already contain them. """ + if self.add_encoders: + logger.warning( + "Encoders (`add_encoders`) are ignored when predicting with `predict_from_dataset()`. " + "Encoders are only applied when calling `predict()`. If you need encoded covariates, " + "make sure the inference dataset already contains them." + ) return self._predict( **self._setup_for_predict( n=n, diff --git a/darts/tests/models/forecasting/test_torch_forecasting_model.py b/darts/tests/models/forecasting/test_torch_forecasting_model.py index fc3d1447d5..f53490f0c7 100644 --- a/darts/tests/models/forecasting/test_torch_forecasting_model.py +++ b/darts/tests/models/forecasting/test_torch_forecasting_model.py @@ -2527,6 +2527,45 @@ def test_fit_with_stride(self, stride): assert len(train_set) == len(val_set) == math.ceil(3 / stride) assert train_set.stride == val_set.stride == stride + @patch("darts.models.forecasting.torch_forecasting_model.logger.warning") + def test_encoders_ignored_in_from_dataset_warns(self, mock_warning): + """`fit_from_dataset`/`predict_from_dataset` warn that `add_encoders` is ignored.""" + model = RNNModel( + 12, + "RNN", + 10, + 10, + **tfm_kwargs, + add_encoders={"cyclic": {"future": ["month"]}}, + ) + + with patch.object(model, "_setup_for_train", return_value={}): + with patch.object(model, "_train", return_value=model): + model.fit_from_dataset(train_dataset=object()) + mock_warning.assert_called_once() + assert "fit_from_dataset" in mock_warning.call_args.args[0] + assert "add_encoders" in mock_warning.call_args.args[0] + + mock_warning.reset_mock() + with patch.object(model, "_setup_for_predict", return_value={}): + with patch.object(model, "_predict", return_value=[]): + model.predict_from_dataset(n=1, dataset=object()) + mock_warning.assert_called_once() + assert "predict_from_dataset" in mock_warning.call_args.args[0] + assert "add_encoders" in mock_warning.call_args.args[0] + + def test_encoders_not_set_from_dataset_no_warn(self): + """`fit_from_dataset`/`predict_from_dataset` do not warn when no encoders are set.""" + model = RNNModel(12, "RNN", 10, 10, **tfm_kwargs) + with patch("darts.models.forecasting.torch_forecasting_model.logger.warning") as mock_warning: + with patch.object(model, "_setup_for_train", return_value={}): + with patch.object(model, "_train", return_value=model): + model.fit_from_dataset(train_dataset=object()) + with patch.object(model, "_setup_for_predict", return_value={}): + with patch.object(model, "_predict", return_value=[]): + model.predict_from_dataset(n=1, dataset=object()) + mock_warning.assert_not_called() + def test_predict_after_fit_from_dataset(self): """Test that the model can predict after being trained with `fit_from_dataset` using all covariates.""" icl, ocl = kwargs["input_chunk_length"], kwargs["output_chunk_length"]