diff --git a/pytorch_forecasting/models/nbeats/_nbeats_adapter_v2.py b/pytorch_forecasting/models/nbeats/_nbeats_adapter_v2.py new file mode 100644 index 000000000..834e354b6 --- /dev/null +++ b/pytorch_forecasting/models/nbeats/_nbeats_adapter_v2.py @@ -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): + """Shared forward / training helpers for NBeats and NBeatsKAN (v2).""" + + def __init__( + 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} diff --git a/pytorch_forecasting/models/nbeats/_nbeats_pkg_v2.py b/pytorch_forecasting/models/nbeats/_nbeats_pkg_v2.py new file mode 100644 index 000000000..d93a17aba --- /dev/null +++ b/pytorch_forecasting/models/nbeats/_nbeats_pkg_v2.py @@ -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 diff --git a/pytorch_forecasting/models/nbeats/_nbeats_v2.py b/pytorch_forecasting/models/nbeats/_nbeats_v2.py new file mode 100644 index 000000000..bfb59a6ef --- /dev/null +++ b/pytorch_forecasting/models/nbeats/_nbeats_v2.py @@ -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 `_. + + 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)