diff --git a/pytorch_forecasting/models/samformer/_samformer_v2.py b/pytorch_forecasting/models/samformer/_samformer_v2.py index 7d859921e..ee0fdc70c 100644 --- a/pytorch_forecasting/models/samformer/_samformer_v2.py +++ b/pytorch_forecasting/models/samformer/_samformer_v2.py @@ -30,6 +30,12 @@ class Samformer(BaseModel): Whether to use Reverse Instance Normalization. Default is True. persistence_weight : float, optional Weight for persistence baseline. Default is 0.0. + metadata : dict + Dataset metadata produced by + :class:`~pytorch_forecasting.data.data_module\ +.EncoderDecoderTimeSeriesDataModule`. + Must contain ``"max_encoder_length"``, ``"max_prediction_length"``, + and ``"encoder_cont"``. """ @classmethod @@ -58,6 +64,9 @@ def __init__( metadata: dict | None = None, **kwargs, ): + if metadata is None: + raise ValueError("metadata is required") + super().__init__( loss=loss, logging_metrics=logging_metrics, diff --git a/tests/test_models/test_samformer_v2.py b/tests/test_models/test_samformer_v2.py new file mode 100644 index 000000000..354af0d51 --- /dev/null +++ b/tests/test_models/test_samformer_v2.py @@ -0,0 +1,10 @@ +import pytest +import torch.nn as nn + +from pytorch_forecasting.models.samformer._samformer_v2 import Samformer + + +def test_samformer_requires_metadata(): + """Test that Samformer rejects missing metadata explicitly.""" + with pytest.raises(ValueError, match="metadata is required"): + Samformer(loss=nn.MSELoss(), hidden_size=512, use_revin=True, metadata=None)