Skip to content
Open
Show file tree
Hide file tree
Changes from 14 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
Original file line number Diff line number Diff line change
Expand Up @@ -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,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

should we pass a whole object in metadata? is this actually necessary? I think metadata should be lightweight and only should have the "metadata" and not objects unless it is actually 100% necessary

}
)

Expand Down
84 changes: 83 additions & 1 deletion pytorch_forecasting/models/base/_base_model_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@
# in the version-2.
########################################################################################


from typing import Any, Optional, Union
from warnings import warn

Expand All @@ -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
Expand Down Expand Up @@ -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.
Expand Down
45 changes: 14 additions & 31 deletions pytorch_forecasting/models/base/_tslib_base_model_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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)
10 changes: 9 additions & 1 deletion pytorch_forecasting/models/tide/_tide_dsipts/_tide_pkg_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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}

Expand Down
32 changes: 25 additions & 7 deletions pytorch_forecasting/models/tide/_tide_dsipts/_tide_v2.py

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why are we changing the forward logic of this model?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

There should be a cleaner and more robust way to handles the shapes?

Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand Down
7 changes: 6 additions & 1 deletion pytorch_forecasting/models/units/_units_pkg_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
{},
Expand All @@ -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}
Expand Down
22 changes: 8 additions & 14 deletions pytorch_forecasting/models/units/_units_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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}
7 changes: 6 additions & 1 deletion pytorch_forecasting/tests/test_all_v2/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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,
Expand Down
46 changes: 45 additions & 1 deletion tests/test_models/test_base_model_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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():
Comment on lines +128 to +143

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

how are these two tests useful? Why not direclty use the DistributionLoss?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These two tests only validate the affine denormalization path (pred * scale + center), which is the straightforward case.
I'll add those distribution loss tests as well

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

My doubt is - DistributionLoss does somethign similar no?
(i have not looked at the math of the losses, so i may be wrong)

@Muhammad-Rebaal Muhammad-Rebaal Sep 25, 2026 •

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These 2 tests are checking specifically making sure that Quantile Loss and Point Forecasts (like MAE) get un-scaled correctly, since I made changes in the transform_output function as well.

"""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)
Loading