diff --git a/pytorch_forecasting/data/data_module/_encoder_decoder_data_module.py b/pytorch_forecasting/data/data_module/_encoder_decoder_data_module.py index c0fc21cc9..e82b31120 100644 --- a/pytorch_forecasting/data/data_module/_encoder_decoder_data_module.py +++ b/pytorch_forecasting/data/data_module/_encoder_decoder_data_module.py @@ -308,6 +308,7 @@ def _prepare_metadata(self): "max_prediction_length": self.max_prediction_length, "min_encoder_length": self._min_encoder_length, "min_prediction_length": self._min_prediction_length, + "target_normalizer": self.target_normalizer, } ) diff --git a/pytorch_forecasting/models/base/_base_model_v2.py b/pytorch_forecasting/models/base/_base_model_v2.py index 6e741a9d0..d095f9f88 100644 --- a/pytorch_forecasting/models/base/_base_model_v2.py +++ b/pytorch_forecasting/models/base/_base_model_v2.py @@ -4,7 +4,6 @@ # in the version-2. ######################################################################################## - from typing import Any, Optional, Union from warnings import warn @@ -18,7 +17,9 @@ from pytorch_forecasting.callbacks.predict import PredictCallback from pytorch_forecasting.metrics import ( + DistributionLoss, Metric, + QuantileLoss, coerce_to_pytorch_forecasting_metric, ) from pytorch_forecasting.utils._classproperty import classproperty @@ -107,6 +108,87 @@ def pkg(cls): """Package class for the model.""" return cls._pkg() + @property + def step_output_size(self) -> int: + """ + Number of outputs predicted per time horizon step. + + Returns + ------- + int + 1 for point predictions, + number of quantiles for QuantileLoss, + number of distribution parameters for DistributionLoss. + """ + if isinstance(self._loss, QuantileLoss): + return len(self._loss.quantiles) + elif isinstance(self._loss, DistributionLoss): + return len(self._loss.distribution_arguments) + return 1 + + def transform_output( + self, + prediction: torch.Tensor, + target_scale: dict[str, torch.Tensor] | torch.Tensor | None = None, + ) -> torch.Tensor: + """ + Transform raw network predictions to real scale and valid parameter domains. + + Parameters + ---------- + prediction : torch.Tensor + Raw network output tensor. + target_scale : dict or torch.Tensor, optional + Scale and center information from dataset/normalizer. + + Returns + ------- + torch.Tensor + Transformed predictions or distribution parameters. + """ + if target_scale is None: + return prediction + + if isinstance(target_scale, dict): + scale = target_scale["scale"] + center = target_scale["center"] + + if scale.dim() == 1: + scale = scale.unsqueeze(-1) + if center.dim() == 1: + center = center.unsqueeze(-1) + combined = torch.cat([center, scale], dim=-1) + else: + if target_scale.dim() == 1 or ( + target_scale.dim() == 2 and target_scale.size(-1) != 2 + ): + scale = target_scale + if scale.dim() == 1: + scale = scale.unsqueeze(-1) + center = torch.zeros_like(scale) + combined = torch.cat([center, scale], dim=-1) + else: + combined = target_scale + center = target_scale[..., 0:1] + scale = target_scale[..., 1:2] + + if isinstance(self._loss, DistributionLoss): + if self.target_normalizer is None: + raise ValueError( + f"{type(self._loss).__name__} requires a target_normalizer " + "in model metadata." + ) + return self._loss.rescale_parameters( + parameters=prediction, + target_scale=combined, + encoder=self.target_normalizer, + ) + + while scale.dim() < prediction.dim(): + scale = scale.unsqueeze(1) + center = center.unsqueeze(1) + return prediction * scale + center + def forward(self, x: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: """ Forward pass of the model. diff --git a/pytorch_forecasting/models/base/_tslib_base_model_v2.py b/pytorch_forecasting/models/base/_tslib_base_model_v2.py index 7ed92d8cd..30b4382b1 100644 --- a/pytorch_forecasting/models/base/_tslib_base_model_v2.py +++ b/pytorch_forecasting/models/base/_tslib_base_model_v2.py @@ -58,6 +58,7 @@ def __init__( ) self.save_hyperparameters(ignore=["loss", "logging_metrics", "metadata"]) self.metadata = metadata or {} + self.target_normalizer = self.metadata.get("target_normalizer", None) self.model_name = self.__class__.__name__ warn( @@ -152,44 +153,26 @@ def predict_step( def transform_output( self, - y_hat: torch.Tensor - | list[ - torch.Tensor - ], # evidenced from TimeXer implementation - in PR #1797 # noqa: E501 - target_scale: dict[str, torch.Tensor] | None, + y_hat: torch.Tensor | list[torch.Tensor], + target_scale: dict[str, torch.Tensor] | torch.Tensor | None, ) -> torch.Tensor | list[torch.Tensor]: """ Transform the output of the model to the original scale. + Delegates to ``BaseModel.transform_output``. + Parameters ---------- - y_hat : Union[torch.Tensor, list[torch.Tensor]] - Dictionary containing the model output. - target_scale : Optional[dict[str, torch.Tensor]] - Dictionary containing the target scale for inverse transformation. + y_hat : torch.Tensor or list[torch.Tensor] + Model output tensor or list of tensors to transform. + target_scale : dict[str, torch.Tensor] or torch.Tensor, optional + Target scale information for inverse transformation. Returns ------- - Union[torch.Tensor, list[torch.Tensor]] - Dictionary containing the transformed output. - - Notes - ----- - WARNING! : This is a temporary implementation and is meant to be replaced with - a more robust scaling and normalization module for v2 of PTF. + torch.Tensor or list[torch.Tensor] + Transformed output tensor or list of tensors. """ - - scale = None - center = None - - if "scale" in target_scale and "center" in target_scale: - scale = target_scale["scale"] - center = target_scale["center"] - else: - raise ValueError("Cannot transform output without scale and center.") - - while scale.dim() < y_hat.dim(): - scale = scale.unsqueeze(0) - center = center.unsqueeze(0) - - return y_hat * scale + center + if isinstance(y_hat, list): + return [super().transform_output(pred, target_scale) for pred in y_hat] + return super().transform_output(y_hat, target_scale) diff --git a/pytorch_forecasting/models/tide/_tide_dsipts/_tide_pkg_v2.py b/pytorch_forecasting/models/tide/_tide_dsipts/_tide_pkg_v2.py index bb3efa74f..6f9c6ba84 100644 --- a/pytorch_forecasting/models/tide/_tide_dsipts/_tide_pkg_v2.py +++ b/pytorch_forecasting/models/tide/_tide_dsipts/_tide_pkg_v2.py @@ -48,7 +48,7 @@ def get_test_train_params(cls): """ import torch.nn as nn - from pytorch_forecasting.metrics import MAE, MAPE + from pytorch_forecasting.metrics import MAE, MAPE, NormalDistributionLoss params = [ dict( @@ -77,6 +77,14 @@ def get_test_train_params(cls): datamodule_cfg=dict(max_encoder_length=4, max_prediction_length=2), loss=MAPE(), ), + dict( + hidden_size=16, + d_model=8, + n_add_enc=1, + n_add_dec=1, + dropout_rate=0.1, + loss=NormalDistributionLoss(), + ), ] default_dm_cfg = {"max_encoder_length": 4, "max_prediction_length": 3} diff --git a/pytorch_forecasting/models/tide/_tide_dsipts/_tide_v2.py b/pytorch_forecasting/models/tide/_tide_dsipts/_tide_v2.py index 80767ebad..97c7c5ebd 100644 --- a/pytorch_forecasting/models/tide/_tide_dsipts/_tide_v2.py +++ b/pytorch_forecasting/models/tide/_tide_dsipts/_tide_v2.py @@ -94,7 +94,8 @@ def __init__( self.past_channels = metadata["encoder_cont"] # psat_vars self.future_channels = metadata["decoder_cont"] # fut_vars self.output_channels = metadata["target"] # target_vars - self.mul = 1 + self.target_normalizer = metadata.get("target_normalizer", None) + self.mul = self.step_output_size self.use_quantiles = False self.outLinear = nn.Linear(d_model, self.output_channels) @@ -283,14 +284,31 @@ def forward(self, X: dict) -> dict: (dense_dec.view(B, self.future_steps, self.d_model), proj_fut), dim=2 ) temp_dec_output = self.temporal_decoder(temp_dec_input, False) - temp_dec_output = temp_dec_output.view( - B, self.future_steps, self.output_channels - ) - - linear_regr = self.linear_target(y_past.view(B, -1)) - linear_output = linear_regr.view(B, self.future_steps, self.output_channels) + if self.mul > 1: + if self.output_channels == 1: + temp_dec_output = temp_dec_output.view(B, self.future_steps, self.mul) + linear_regr = self.linear_target(y_past.view(B, -1)) + linear_output = linear_regr.view(B, self.future_steps, self.mul) + else: + temp_dec_output = temp_dec_output.view( + B, self.future_steps, self.output_channels, self.mul + ) + linear_regr = self.linear_target(y_past.view(B, -1)) + linear_output = linear_regr.view( + B, self.future_steps, self.output_channels, self.mul + ) + else: + temp_dec_output = temp_dec_output.view( + B, self.future_steps, self.output_channels + ) + linear_regr = self.linear_target(y_past.view(B, -1)) + linear_output = linear_regr.view(B, self.future_steps, self.output_channels) output = temp_dec_output + linear_output + + if "target_scale" in batch and hasattr(self, "transform_output"): + output = self.transform_output(output, batch["target_scale"]) + return {"prediction": output} # function to concat embedded categorical variables diff --git a/pytorch_forecasting/models/units/_units_pkg_v2.py b/pytorch_forecasting/models/units/_units_pkg_v2.py index 29b0d4b46..ba705d29c 100644 --- a/pytorch_forecasting/models/units/_units_pkg_v2.py +++ b/pytorch_forecasting/models/units/_units_pkg_v2.py @@ -50,7 +50,7 @@ def get_test_train_params(cls): instance. ``create_test_instance`` uses the first (or only) dictionary in ``params``. """ - from pytorch_forecasting.metrics import QuantileLoss + from pytorch_forecasting.metrics import NormalDistributionLoss, QuantileLoss params = [ {}, @@ -77,6 +77,11 @@ def get_test_train_params(cls): "stride": 4, "loss": QuantileLoss(quantiles=[0.1, 0.5, 0.9]), }, + { + "patch_len": 8, + "stride": 4, + "loss": NormalDistributionLoss(), + }, ] base_dm_cfg = {"max_encoder_length": 16, "max_prediction_length": 4} diff --git a/pytorch_forecasting/models/units/_units_v2.py b/pytorch_forecasting/models/units/_units_v2.py index b6e6b235e..8430dedb3 100644 --- a/pytorch_forecasting/models/units/_units_v2.py +++ b/pytorch_forecasting/models/units/_units_v2.py @@ -114,6 +114,7 @@ def __init__( self.context_length = self.metadata.get("max_encoder_length", 0) self.prediction_length = self.metadata.get("max_prediction_length", 0) self.target_dim = self.metadata.get("target", 1) + self.target_normalizer = self.metadata.get("target_normalizer", None) if d_model % n_heads != 0: raise ValueError( @@ -160,16 +161,7 @@ def _init_network(self): self.norm = nn.LayerNorm(self.d_model) - self.n_quantiles = None - # TODO: add DistributionLoss support - - if isinstance(self._loss, QuantileLoss): - self.n_quantiles = len(self._loss.quantiles) - - output_dim = self.prediction_length * self.target_dim - - if self.n_quantiles is not None: - output_dim = self.prediction_length * self.target_dim * self.n_quantiles + output_dim = self.prediction_length * self.target_dim * self.step_output_size self.head = nn.Sequential( nn.Flatten(start_dim=1), @@ -204,15 +196,17 @@ def forward(self, x: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: raw = self.head(patch_out) - if self.n_quantiles is not None: + if self.step_output_size > 1: if self.target_dim == 1: - out = raw.view(B, self.prediction_length, self.n_quantiles) + out = raw.view(B, self.prediction_length, self.step_output_size) else: out = raw.view( - B, self.prediction_length, self.target_dim, self.n_quantiles + B, self.prediction_length, self.target_dim, self.step_output_size ) - # TODO: add DistributionLoss output reshape else: out = raw.view(B, self.prediction_length, self.target_dim) + if "target_scale" in x and hasattr(self, "transform_output"): + out = self.transform_output(out, x["target_scale"]) + return {"prediction": out} diff --git a/pytorch_forecasting/tests/test_all_v2/utils.py b/pytorch_forecasting/tests/test_all_v2/utils.py index a8cb714dc..321736936 100644 --- a/pytorch_forecasting/tests/test_all_v2/utils.py +++ b/pytorch_forecasting/tests/test_all_v2/utils.py @@ -4,7 +4,8 @@ from pytorch_forecasting.base._base_pkg import Base_pkg from pytorch_forecasting.data import TimeSeries -from pytorch_forecasting.metrics import SMAPE +from pytorch_forecasting.data.encoders import EncoderNormalizer +from pytorch_forecasting.metrics import SMAPE, DistributionLoss def _setup_pkg_and_data( @@ -31,6 +32,10 @@ def _setup_pkg_and_data( if "loss" not in model_cfg: model_cfg["loss"] = SMAPE() + if isinstance(model_cfg.get("loss"), DistributionLoss): + if "target_normalizer" not in datamodule_cfg: + datamodule_cfg["target_normalizer"] = EncoderNormalizer() + default_datamodule_cfg = { "train_val_test_split": (0.8, 0.2), "add_relative_time_idx": True, diff --git a/tests/test_models/test_base_model_v2.py b/tests/test_models/test_base_model_v2.py index fe08996f1..5e0fd2232 100644 --- a/tests/test_models/test_base_model_v2.py +++ b/tests/test_models/test_base_model_v2.py @@ -3,7 +3,8 @@ import pytest import torch -from pytorch_forecasting.metrics import MAE +from pytorch_forecasting.data.encoders import EncoderNormalizer +from pytorch_forecasting.metrics import MAE, NormalDistributionLoss from pytorch_forecasting.models.base._base_model_v2 import BaseModel @@ -122,3 +123,46 @@ def test_optimizer_instance(): model.optimizer = opt cfg = model.configure_optimizers() assert cfg["optimizer"] is opt + + +def test_transform_output_dict_target_scale(): + """Dict target_scale applies affine denormalization correctly.""" + model = _make_model() + raw = torch.randn(4, 12, 1) + center = torch.tensor([10.0, 20.0, 30.0, 40.0]) + scale = torch.tensor([2.0, 3.0, 4.0, 5.0]) + target_scale = {"center": center, "scale": scale} + + result = model.transform_output(raw, target_scale) + + assert result.shape == raw.shape + assert torch.allclose(result[0], raw[0] * 2.0 + 10.0) + assert torch.allclose(result[3], raw[3] * 5.0 + 40.0) + + +def test_transform_output_plain_tensor_target_scale(): + """Plain tensor target_scale correctly applies affine denormalization.""" + model = _make_model() + raw = torch.randn(4, 12, 1) + target_scale = torch.tensor([[10.0, 2.0], [20.0, 3.0], [30.0, 4.0], [40.0, 5.0]]) + + result = model.transform_output(raw, target_scale) + + assert result.shape == raw.shape + assert torch.allclose(result[0], raw[0] * 2.0 + 10.0) + assert torch.allclose(result[2], raw[2] * 4.0 + 30.0) + + +def test_transform_output_distribution_loss(): + """DistributionLoss path rescales parameters via the loss function.""" + + model = _make_model(loss=NormalDistributionLoss()) + model.target_normalizer = EncoderNormalizer() + + raw = torch.randn(2, 4, 2) + target_scale = torch.tensor([[5.0, 2.0], [3.0, 1.5]]) + + result = model.transform_output(raw, target_scale) + + assert result.shape == (2, 4, 4) + assert not torch.equal(result[..., 2:], raw)