Skip to content
Open
Show file tree
Hide file tree
Changes from 4 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
239 changes: 239 additions & 0 deletions pytorch_forecasting/models/nbeats/_nbeats_adapter_v2.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,239 @@
"""Shared N-Beats adapter for pytorch-forecasting v2."""

from typing import Any, Optional, Union

import torch
from torch import nn
from torch.optim import Optimizer

from pytorch_forecasting.layers._nbeats._blocks import (
NBEATSSeasonalBlock,
NBEATSTrendBlock,
)
from pytorch_forecasting.metrics import Metric
from pytorch_forecasting.models.base._tslib_base_model_v2 import TslibBaseModel


class NBeatsAdapterV2(TslibBaseModel):

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.

It should be BaseModel no?

"""Shared forward / training helpers for NBeats and NBeatsKAN (v2)."""

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.

i think it will also be used for NBEATx or?


def __init__(

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.

please docstrings

self,
loss: Metric,
logging_metrics: list[nn.Module] | None = None,
optimizer: Optimizer | str | None = "adam",
optimizer_params: dict | None = None,
lr_scheduler: str | None = None,
lr_scheduler_params: dict | None = None,
metadata: dict | None = None,
backcast_loss_ratio: float = 0.0,
**kwargs: Any,
):
super().__init__(
loss=loss,
logging_metrics=logging_metrics,
optimizer=optimizer,
optimizer_params=optimizer_params,
lr_scheduler=lr_scheduler,
lr_scheduler_params=lr_scheduler_params,
metadata=metadata,
)
self.backcast_loss_ratio = backcast_loss_ratio

def _target_from_batch(self, x: dict[str, torch.Tensor]) -> torch.Tensor:
"""Extract univariate target history.

v1 used ``x["encoder_cont"][..., 0]``. v2 tslib batches keep the target
in ``history_target``.
"""
target = x["history_target"]
if target.ndim == 3:
target = target[..., 0]
return target

def forward(self, x: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
"""Pass forward of network.

Network steps match v1 ``NBeatsAdapter.forward``; only input assembly
and output packaging differ for the v2 API.
"""
# --- v2 batch adapter (v1: target = x["encoder_cont"][..., 0]) ---
target = self._target_from_batch(x)

# --- same as v1 from here ---
timesteps = self.context_length + self.prediction_length
generic_forecast = [
torch.zeros(
(target.size(0), timesteps), dtype=torch.float32, device=self.device
)
]
trend_forecast = [
torch.zeros(
(target.size(0), timesteps), dtype=torch.float32, device=self.device
)
]
seasonal_forecast = [
torch.zeros(
(target.size(0), timesteps), dtype=torch.float32, device=self.device
)
]
forecast = torch.zeros(
(target.size(0), self.prediction_length),
dtype=torch.float32,
device=self.device,
)

backcast = target # initialize backcast
for i, block in enumerate(self.net_blocks):
# evaluate block
backcast_block, forecast_block = block(backcast)

# add for interpretation
full = torch.cat([backcast_block.detach(), forecast_block.detach()], dim=1)
if isinstance(block, NBEATSTrendBlock):
trend_forecast.append(full)
elif isinstance(block, NBEATSSeasonalBlock):
seasonal_forecast.append(full)
else:
generic_forecast.append(full)

# update backcast and forecast
backcast = (
backcast - backcast_block
) # do not use backcast -= backcast_block as this signifies an inline operation # noqa: E501
forecast = forecast + forecast_block

prediction = forecast.unsqueeze(-1)
backcast_out = (target - backcast).unsqueeze(-1)
trend = torch.stack(trend_forecast, dim=0).sum(0).unsqueeze(-1)
seasonality = torch.stack(seasonal_forecast, dim=0).sum(0).unsqueeze(-1)
generic = torch.stack(generic_forecast, dim=0).sum(0).unsqueeze(-1)

# v1 applied transform_output via BaseModel; v2 tslib does so when scales exist
if "target_scale" in x:
prediction = self.transform_output(prediction, x["target_scale"])
backcast_out = self.transform_output(backcast_out, x["target_scale"])
trend = self.transform_output(trend, x["target_scale"])
seasonality = self.transform_output(seasonality, x["target_scale"])
generic = self.transform_output(generic, x["target_scale"])

# v1: to_network_output(...); v2: plain dict
return {
"prediction": prediction,
"backcast": backcast_out,
"trend": trend,
"seasonality": seasonality,
"generic": generic,
}

def _compute_loss(
self,
x: dict[str, torch.Tensor],
y: torch.Tensor,
out: dict[str, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
"""Forecast loss plus optional backcast term (v1 ``step`` parity).

Applied for train / val / test (not predict), matching v1's
``not self.predicting`` guard on the shared ``step()``.
"""
y_hat = out["prediction"]
loss = self.loss(y_hat, y)

if self.backcast_loss_ratio > 0:
backcast = out["backcast"].squeeze(-1)
encoder_target = self._target_from_batch(x)

backcast_weight = (
self.backcast_loss_ratio
* self.prediction_length
/ max(self.context_length, 1)
)
backcast_weight = backcast_weight / (backcast_weight + 1)
forecast_weight = 1 - backcast_weight

backcast_loss = (backcast - encoder_target).abs().mean() * backcast_weight
loss = loss * forecast_weight + backcast_loss

return loss, y_hat

def training_step(
self, batch: tuple[dict[str, torch.Tensor]], batch_idx: int
) -> dict[str, torch.Tensor]:
"""
Training step for the model with optional backcast loss.

Parameters
----------
batch : Tuple[Dict[str, torch.Tensor]]
Batch of data containing input and target tensors.
batch_idx : int
Index of the batch.

Returns
-------
STEP_OUTPUT
Dictionary containing the loss and other metrics.
"""
x, y = batch
out = self(x)
loss, y_hat = self._compute_loss(x, y, out)
self.log(
"train_loss", loss, on_step=True, on_epoch=True, prog_bar=True, logger=True
)
self.log_metrics(y_hat, y, prefix="train")
return {"loss": loss}

def validation_step(
self, batch: tuple[dict[str, torch.Tensor]], batch_idx: int
) -> dict[str, torch.Tensor]:
"""
Validation step for the model with optional backcast loss.

Parameters
----------
batch : Tuple[Dict[str, torch.Tensor]]
Batch of data containing input and target tensors.
batch_idx : int
Index of the batch.

Returns
-------
STEP_OUTPUT
Dictionary containing the loss and other metrics.
"""
x, y = batch
out = self(x)
loss, y_hat = self._compute_loss(x, y, out)
self.log(
"val_loss", loss, on_step=False, on_epoch=True, prog_bar=True, logger=True
)
self.log_metrics(y_hat, y, prefix="val")
return {"val_loss": loss}

def test_step(
self, batch: tuple[dict[str, torch.Tensor]], batch_idx: int
) -> dict[str, torch.Tensor]:
"""
Test step for the model with optional backcast loss.

Parameters
----------
batch : Tuple[Dict[str, torch.Tensor]]
Batch of data containing input and target tensors.
batch_idx : int
Index of the batch.

Returns
-------
STEP_OUTPUT
Dictionary containing the loss and other metrics.
"""
x, y = batch
out = self(x)
loss, y_hat = self._compute_loss(x, y, out)
self.log(
"test_loss", loss, on_step=False, on_epoch=True, prog_bar=True, logger=True
)
self.log_metrics(y_hat, y, prefix="test")
return {"test_loss": loss}
98 changes: 98 additions & 0 deletions pytorch_forecasting/models/nbeats/_nbeats_pkg_v2.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,98 @@
"""NBeats v2 package container."""

from pytorch_forecasting.base._base_pkg import Base_pkg


class NBeats_pkg_v2(Base_pkg):
"""NBeats v2 package container."""

_tags = {
"info:name": "NBeats",
"info:compute": 1,
"info:y_type": ["numeric"],
"authors": [
"dmitri-carpov" # paper author
"jdb78", # for v1
"Faakhir30",
],
"capability:exogenous": False,
"capability:multivariate": False,
"capability:pred_int": False,
"capability:flexible_history_length": False,
"capability:cold_start": False,
}

@classmethod
def get_cls(cls):
"""Get model class."""
from pytorch_forecasting.models.nbeats._nbeats_v2 import NBeats

return NBeats

@classmethod
def get_datamodule_cls(cls):
"""Get the underlying DataModule class."""
from pytorch_forecasting.data.data_module import TslibDataModule

return TslibDataModule

@classmethod
def get_test_train_params(cls):
"""Return testing parameter settings for the trainer."""
from pytorch_forecasting.metrics import MAE, MAPE, SMAPE

params = [
{
"widths": [16, 32],
"num_blocks": [1, 1],
"num_block_layers": [2, 2],
},
{
"backcast_loss_ratio": 1.0,
"widths": [16, 32],
"num_blocks": [1, 1],
"num_block_layers": [2, 2],
},
{
"stack_types": ["generic"],
"num_blocks": [1],
"num_block_layers": [2],
"widths": [16],
"expansion_coefficient_lengths": [8],
"sharing": [False],
},
{
"loss": MAE(),
"widths": [16, 32],
"num_blocks": [1, 1],
"num_block_layers": [2, 2],
},
{
"loss": MAPE(),
"logging_metrics": [SMAPE()],
"widths": [16, 32],
"num_blocks": [1, 1],
"num_block_layers": [2, 2],
},
{
"optimizer": "adamw",
"lr_scheduler": "cosine_annealing",
"lr_scheduler_params": {"T_max": 5},
"widths": [16, 32],
"num_blocks": [1, 1],
"num_block_layers": [2, 2],
},
]

default_dm_cfg = {
"context_length": 8,
"prediction_length": 3,
"add_relative_time_idx": False,
}

for param in params:
current_dm_cfg = param.get("datamodule_cfg", {})
default_dm_cfg.update(current_dm_cfg)
param["datamodule_cfg"] = default_dm_cfg.copy()

return params
Loading
Loading