Skip to content
Draft
Show file tree
Hide file tree
Changes from 2 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
160 changes: 160 additions & 0 deletions pytorch_forecasting/models/nbeats/_nbeats_adapter_v2.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,160 @@
"""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 training_step(

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.

are test_step etc not needed?

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.

yeah, test_step and validation_step should be implemented considering the backcast_loss thing. added now, mb

self, batch: tuple[dict[str, torch.Tensor]], batch_idx: int
) -> dict[str, torch.Tensor]:
"""Training step with optional backcast loss (v1 ``step`` parity)."""
x, y = batch
out = self(x)
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

# Compute backcast term directly (avoid Metric.update state / shape quirks).
# v1 used self.loss(backcast, encoder_target); v2 BaseModel losses are
# wired for forecast horizon shapes only.
backcast_loss = (backcast - encoder_target).abs().mean() * backcast_weight
loss = loss * forecast_weight + backcast_loss

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}
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
130 changes: 130 additions & 0 deletions pytorch_forecasting/models/nbeats/_nbeats_v2.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
"""
N-Beats model for pytorch-forecasting v2 (no covariates).
"""

from typing import Any, Optional, Union

from torch import nn
from torch.optim import Optimizer

from pytorch_forecasting.layers._nbeats._blocks import (
NBEATSGenericBlock,
NBEATSSeasonalBlock,
NBEATSTrendBlock,
)
from pytorch_forecasting.metrics import MAE, MAPE, RMSE, SMAPE, Metric
from pytorch_forecasting.models.nbeats._nbeats_adapter_v2 import NBeatsAdapterV2


class NBeats(NBeatsAdapterV2):
"""
N-BEATS for pytorch-forecasting v2.

Based on
`N-BEATS: Neural basis expansion analysis for interpretable time series
forecasting <http://arxiv.org/abs/1905.10437>`_.

Network construction matches the v1 ``NBeats`` class; ``context_length`` /
``prediction_length`` come from datamodule ``metadata`` instead of
``from_dataset``.
"""

@classmethod
def _pkg(cls):
"""Package for the model."""
from pytorch_forecasting.models.nbeats._nbeats_pkg_v2 import NBeats_pkg_v2

return NBeats_pkg_v2

def __init__(
self,
loss: Metric,
stack_types: list[str] | None = None,
num_blocks: list[int] | None = None,
num_block_layers: list[int] | None = None,
widths: list[int] | None = None,
sharing: list[bool] | None = None,
expansion_coefficient_lengths: list[int] | None = None,
dropout: float = 0.1,
backcast_loss_ratio: float = 0.0,
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,
**kwargs: Any,
):
if expansion_coefficient_lengths is None:
expansion_coefficient_lengths = [3, 7]
if sharing is None:
sharing = [True, True]
if widths is None:
widths = [32, 512]
if num_block_layers is None:
num_block_layers = [3, 3]
if num_blocks is None:
num_blocks = [3, 3]
if stack_types is None:
stack_types = ["trend", "seasonality"]
if logging_metrics is None:
logging_metrics = [SMAPE(), MAE(), RMSE(), MAPE()]

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,
backcast_loss_ratio=backcast_loss_ratio,
)
self.save_hyperparameters(ignore=["loss", "logging_metrics", "metadata"])

self.stack_types = stack_types
self.num_blocks = num_blocks
self.num_block_layers = num_block_layers
self.widths = widths
self.sharing = sharing
self.expansion_coefficient_lengths = expansion_coefficient_lengths
self.dropout = dropout

self._init_network()

def _init_network(self):
"""Build N-BEATS stacks (same block wiring as v1)."""
self.net_blocks = nn.ModuleList()
for stack_id, stack_type in enumerate(self.stack_types):
for _ in range(self.num_blocks[stack_id]):
if stack_type == "generic":
net_block = NBEATSGenericBlock(
units=self.widths[stack_id],
thetas_dim=self.expansion_coefficient_lengths[stack_id],
num_block_layers=self.num_block_layers[stack_id],
backcast_length=self.context_length,
forecast_length=self.prediction_length,
dropout=self.dropout,
)
elif stack_type == "seasonality":
net_block = NBEATSSeasonalBlock(
units=self.widths[stack_id],
num_block_layers=self.num_block_layers[stack_id],
backcast_length=self.context_length,
forecast_length=self.prediction_length,
min_period=self.expansion_coefficient_lengths[stack_id],
dropout=self.dropout,
)
elif stack_type == "trend":
net_block = NBEATSTrendBlock(
units=self.widths[stack_id],
thetas_dim=self.expansion_coefficient_lengths[stack_id],
num_block_layers=self.num_block_layers[stack_id],
backcast_length=self.context_length,
forecast_length=self.prediction_length,
dropout=self.dropout,
)
else:
raise ValueError(f"Unknown stack type {stack_type}")

self.net_blocks.append(net_block)
Loading