Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions darts/models/forecasting/torch_forecasting_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
39 changes: 39 additions & 0 deletions darts/tests/models/forecasting/test_torch_forecasting_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down