Skip to content
Open
Show file tree
Hide file tree
Changes from 6 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
75 changes: 74 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,78 @@ 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:
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)
49 changes: 48 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,7 @@
import pytest
import torch

from pytorch_forecasting.metrics import MAE
from pytorch_forecasting.metrics import MAE, QuantileLoss
from pytorch_forecasting.models.base._base_model_v2 import BaseModel


Expand Down Expand Up @@ -122,3 +122,50 @@ def test_optimizer_instance():
model.optimizer = opt
cfg = model.configure_optimizers()
assert cfg["optimizer"] is opt


def test_step_output_size_point_loss():
"""Point loss (MAE) produces a single output per timestep."""
model = _make_model(loss=MAE())
assert model.step_output_size == 1


def test_step_output_size_quantile_loss():
"""QuantileLoss with 3 quantiles needs 3 outputs per timestep."""
model = _make_model(loss=QuantileLoss(quantiles=[0.1, 0.5, 0.9]))
assert model.step_output_size == 3


def test_transform_output_none_is_identity():
"""No target_scale means predictions pass through untouched."""
model = _make_model()
raw = torch.randn(2, 6, 1)
assert torch.equal(model.transform_output(raw, None), raw)


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