-
Notifications
You must be signed in to change notification settings - Fork 900
[ENH] Implement NBEATS in v2 #2373
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 2 commits
3280392
0fdcee2
852a209
1b827d9
4fd6e27
4a34751
b41988c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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): | ||
| """Shared forward / training helpers for NBeats and NBeatsKAN (v2).""" | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. i think it will also be used for |
||
|
|
||
| def __init__( | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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( | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. are
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. yeah, |
||
| 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} | ||
| 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 |
| 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) |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
It should be
BaseModelno?