diff --git a/docs/source/m_layer_v2.rst b/docs/source/m_layer_v2.rst index b6f09f5ee..b0091e2a0 100644 --- a/docs/source/m_layer_v2.rst +++ b/docs/source/m_layer_v2.rst @@ -48,3 +48,4 @@ See the detailed API documentation for the V2 base classes and specific model im models.tide._tide_dsipts._tide_v2.TIDE models.timexer._timexer_v2.TimeXer models.mlp._decodermlp_v2.DecoderMLP_v2 + models.softs._softs_v2.SOFTS diff --git a/docs/source/pkg_v2.rst b/docs/source/pkg_v2.rst index 59731d854..750cdfba0 100644 --- a/docs/source/pkg_v2.rst +++ b/docs/source/pkg_v2.rst @@ -100,3 +100,4 @@ See the detailed API documentation for the available V2 Package classes below: models.tide._tide_dsipts._tide_v2_pkg.TIDE_pkg_v2 models.timexer._timexer_pkg_v2.TimeXer_pkg_v2 models.mlp._decodermlp_pkg_v2.DecoderMLP_pkg_v2 + models.softs._softs_pkg_v2.SOFTS_pkg_v2 diff --git a/pytorch_forecasting/layers/_blocks/__init__.py b/pytorch_forecasting/layers/_blocks/__init__.py index 512760a31..7f28a1d97 100644 --- a/pytorch_forecasting/layers/_blocks/__init__.py +++ b/pytorch_forecasting/layers/_blocks/__init__.py @@ -1,3 +1,9 @@ from pytorch_forecasting.layers._blocks._residual_block_dsipts import ResidualBlock +from pytorch_forecasting.layers._blocks._softs_block import ( + STADModule, +) -__all__ = ["ResidualBlock"] +__all__ = [ + "ResidualBlock", + "STADModule", +] diff --git a/pytorch_forecasting/layers/_blocks/_softs_block.py b/pytorch_forecasting/layers/_blocks/_softs_block.py new file mode 100644 index 000000000..7898687d1 --- /dev/null +++ b/pytorch_forecasting/layers/_blocks/_softs_block.py @@ -0,0 +1,62 @@ +""" +SOFTS Blocks for Star Aggregate-Dispatch Network. +""" + +import torch +import torch.nn as nn + + +class STADModule(nn.Module): + """ + Star Aggregate-Dispatch (STAD) Module for capturing inter-series dependencies. + + Uses a star-topology to aggregate all channels into a central node, + process it via an MLP, and dispatch back — achieving O(C) cross-channel + mixing instead of O(C²) self-attention. + + Parameters + ---------- + d_model : int + Embedding dimension per channel per time step. + d_core : int + Dimension of the central star node (information bottleneck). + dropout : float, default=0.0 + Dropout probability inside the channel-mixing MLP. + """ + + def __init__(self, d_model: int, d_core: int, dropout: float = 0.0): + super().__init__() + self.channel_mixing = nn.Sequential( + nn.Linear(d_model, d_model), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(d_model, d_model), + ) + self.gen_weight = nn.Linear(d_model, d_core) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Aggregate channel features into a star node and dispatch back. + + Parameters + ---------- + x : torch.Tensor + Shape ``(batch_size, n_channels, seq_len, d_model)``. + + Returns + ------- + torch.Tensor + Same shape as input, enriched with cross-channel context. + """ + + B, C, L, D = x.shape + + w = self.gen_weight(x).mean(dim=2) + w = torch.softmax(w, dim=1) + + x_pooled = x.mean(dim=2) + core_node = torch.einsum("bcd,bce->bed", x_pooled, w) + core_node = self.channel_mixing(core_node) + dispatch_out = torch.einsum("bed,bce->bcd", core_node, w) + + dispatch_out = dispatch_out.unsqueeze(2).repeat(1, 1, L, 1) + return x + dispatch_out diff --git a/pytorch_forecasting/layers/_encoders/_softs_encoder.py b/pytorch_forecasting/layers/_encoders/_softs_encoder.py new file mode 100644 index 000000000..e432c71db --- /dev/null +++ b/pytorch_forecasting/layers/_encoders/_softs_encoder.py @@ -0,0 +1,60 @@ +""" +Implementation of EncoderLayer for SOFTS from `nn.Module`. +""" + +import torch +import torch.nn as nn + +from pytorch_forecasting.layers._blocks._softs_block import STADModule + + +class SOFTSEncoderLayer(nn.Module): + """ + Single encoder layer for SOFTS, combining STAD and a Feed-Forward Network. + + Applies Pre-LayerNorm STAD (cross-channel) then FFN (within-channel) + with residual connections, following the Pre-LN Transformer convention. + + Parameters + ---------- + d_model : int + Embedding dimension per channel per time step. + d_core : int + Dimension of the central star node in the STAD sub-layer. + d_ff : int + Hidden dimension of the feed-forward network (typically 4 x d_model). + dropout : float, default=0.0 + Dropout probability applied after the STAD and FFN sub-layers. + """ + + def __init__(self, d_model: int, d_core: int, d_ff: int, dropout: float = 0.0): + super().__init__() + self.stad = STADModule(d_model=d_model, d_core=d_core, dropout=dropout) + self.ffn = nn.Sequential( + nn.Linear(d_model, d_ff), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(d_ff, d_model), + ) + self.norm1 = nn.LayerNorm(d_model) + self.norm2 = nn.LayerNorm(d_model) + self.dropout = nn.Dropout(dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """ + Apply one SOFTS encoder layer: STAD sub-layer then FFN sub-layer. + + Parameters + ---------- + x : torch.Tensor + Input tensor of shape ``(batch_size, n_channels, seq_len, d_model)``. + + Returns + ------- + torch.Tensor + Output tensor of shape ``(batch_size, n_channels, seq_len, d_model)``. + """ + x = x + self.dropout(self.stad(self.norm1(x))) + x = x + self.dropout(self.ffn(self.norm2(x))) + + return x diff --git a/pytorch_forecasting/models/__init__.py b/pytorch_forecasting/models/__init__.py index dc635b261..d81e5a4ea 100644 --- a/pytorch_forecasting/models/__init__.py +++ b/pytorch_forecasting/models/__init__.py @@ -15,6 +15,7 @@ from pytorch_forecasting.models.nhits import NHiTS from pytorch_forecasting.models.nn import GRU, LSTM, MultiEmbedding, get_rnn from pytorch_forecasting.models.rnn import RecurrentNetwork +from pytorch_forecasting.models.softs import SOFTS, SOFTS_pkg_v2 from pytorch_forecasting.models.temporal_fusion_transformer import ( TemporalFusionTransformer, ) @@ -42,4 +43,6 @@ "TiDEModel", "TimeXer", "xLSTMTime", + "SOFTS", + "SOFTS_pkg_v2", ] diff --git a/pytorch_forecasting/models/softs/__init__.py b/pytorch_forecasting/models/softs/__init__.py new file mode 100644 index 000000000..fd429fc99 --- /dev/null +++ b/pytorch_forecasting/models/softs/__init__.py @@ -0,0 +1,8 @@ +""" +SOFTS Model for Multivariate Time Series Forecasting. +""" + +from pytorch_forecasting.models.softs._softs_pkg_v2 import SOFTS_pkg_v2 +from pytorch_forecasting.models.softs._softs_v2 import SOFTS + +__all__ = ["SOFTS", "SOFTS_pkg_v2"] diff --git a/pytorch_forecasting/models/softs/_softs_pkg_v2.py b/pytorch_forecasting/models/softs/_softs_pkg_v2.py new file mode 100644 index 000000000..f2848cf5e --- /dev/null +++ b/pytorch_forecasting/models/softs/_softs_pkg_v2.py @@ -0,0 +1,99 @@ +""" +Packages container for SOFTS model. +""" + +from pytorch_forecasting.base._base_pkg import Base_pkg + + +class SOFTS_pkg_v2(Base_pkg): + """ + SOFTS package container. + Reference : https://arxiv.org/abs/2404.14197 + """ + + _tags = { + "info:name": "SOFTS", + "info:y_type": ["numeric"], + "info:compute": 2, + "authors": ["Secilia-Cxy", "Muhammad-Rebaal"], + "capability:exogenous": True, + "capability:multivariate": True, + "capability:pred_int": True, + "capability:flexible_history_length": True, + "capability:cold_start": False, + } + + @classmethod + def get_cls(cls): + """Get model class.""" + from pytorch_forecasting.models.softs._softs_v2 import SOFTS + + return SOFTS + + @classmethod + def get_datamodule_cls(cls): + """Get the underlying DataModule class.""" + from pytorch_forecasting.data.data_module import ( + EncoderDecoderTimeSeriesDataModule, + ) + + return EncoderDecoderTimeSeriesDataModule + + @classmethod + def get_test_train_params(cls): + """Return testing parameter settings for the trainer. + + Returns + ------- + list of dict + Each dict is a valid set of constructor arguments for ``SOFTS``. + The key ``datamodule_cfg`` is passed to the DataModule, not the model. + """ + from pytorch_forecasting.metrics import MAE, MAPE, RMSE, SMAPE + + params = [ + {}, + dict(hidden_size=64, d_core=64, d_ff=256, n_layers=1), + dict(hidden_size=128, n_layers=1, use_revin=False), + dict( + hidden_size=64, + n_layers=1, + loss=MAE(), + ), + dict( + hidden_size=64, + n_layers=1, + loss=MAPE(), + ), + dict( + hidden_size=64, + n_layers=1, + loss=RMSE(), + ), + dict( + hidden_size=64, + n_layers=1, + use_revin=False, + loss=MAE(), + ), + dict(hidden_size=64, dropout=0.0, n_layers=1), + dict(datamodule_cfg=dict(max_encoder_length=16, max_prediction_length=4)), + dict( + optimizer="adamw", + lr_scheduler="cosine_annealing", + lr_scheduler_params={"T_max": 5}, + ), + dict( + optimizer="adagrad", + optimizer_params={"lr": 1e-3}, + ), + dict(hidden_size=64, n_layers=1, logging_metrics=[SMAPE()]), + ] + + default_dm_cfg = {"max_encoder_length": 8, "max_prediction_length": 2} + + for param in params: + current_dm_cfg = param.get("datamodule_cfg", {}) + param["datamodule_cfg"] = {**default_dm_cfg, **current_dm_cfg} + + return params diff --git a/pytorch_forecasting/models/softs/_softs_v2.py b/pytorch_forecasting/models/softs/_softs_v2.py new file mode 100644 index 000000000..3c3e4f0f8 --- /dev/null +++ b/pytorch_forecasting/models/softs/_softs_v2.py @@ -0,0 +1,207 @@ +""" +SOFTS Model Implementation for PyTorch Forecasting v2. +------------------------------------------------------- +""" + +import torch +import torch.nn as nn +from torch.optim import Optimizer + +from pytorch_forecasting.layers._encoders._softs_encoder import SOFTSEncoderLayer +from pytorch_forecasting.layers._normalization import RevIN +from pytorch_forecasting.models.base._base_model_v2 import BaseModel + + +class SOFTS(BaseModel): + """ + SOFTS: Efficient Multivariate Time Series Forecasting with Series-Core Fusion. + + GitHub Link: https://github.com/Secilia-Cxy/SOFTS/ + + Research Paper: https://arxiv.org/abs/2404.14197 + + Parameters + ---------- + hidden_size: int + Embedding size of individual time series channel, default = 512 + d_core: int + Hidden dimension of the central core node, default = 512 + d_ff: int + Dimension of the feed-forward network, default = 2048 + n_layers: int + Number of encoder layers, default = 2 + dropout: float + Dropout rate, default = 0.1 + use_revin: bool + Whether to use RevIN, default = True + optimizer: Optimizer | str + Optimizer to use for training, default = "adam" + optimizer_params: dict | None + Parameters for the optimizer, default = None + lr_scheduler: str | None + Learning rate scheduler to use, default = None + lr_scheduler_params: dict | None + Parameters for the learning rate scheduler, default = None + """ + + @classmethod + def _pkg(cls): + from pytorch_forecasting.models.softs._softs_pkg_v2 import SOFTS_pkg_v2 + + return SOFTS_pkg_v2 + + def __init__( + self, + loss: nn.Module, + hidden_size: int = 512, + d_core: int = 512, + d_ff: int = 2048, + n_layers: int = 2, + dropout: float = 0.1, + use_revin: bool = True, + 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, + ): + super().__init__( + loss=loss, + logging_metrics=logging_metrics, + optimizer=optimizer, + optimizer_params=optimizer_params, + lr_scheduler=lr_scheduler, + lr_scheduler_params=lr_scheduler_params, + ) + self.save_hyperparameters(ignore=["loss", "logging_metrics", "metadata"]) + + self.metadata = metadata or {} + self.context_length = self.metadata.get("max_encoder_length", 0) + self.prediction_length = self.metadata.get("max_prediction_length", 0) + + self.cont_dim = self.metadata.get("encoder_cont", 0) + self.target_dim = self.metadata.get("target", 1) + + self.use_revin = use_revin + self.n_quantiles = ( + len(loss.quantiles) + if hasattr(loss, "quantiles") and loss.quantiles is not None + else 1 + ) + + self._init_network(hidden_size, d_core, d_ff, n_layers, dropout) + + def _init_network(self, d_model, d_core, d_ff, n_layers, dropout): + # Normalization + if self.use_revin: + self.revin = RevIN(num_features=self.cont_dim + self.target_dim) + + # Embedding Layer + self.embedding = nn.Linear(1, d_model) + + # Encoder Blocks + self.encoder = nn.ModuleList( + [ + SOFTSEncoderLayer( + d_model=d_model, d_core=d_core, d_ff=d_ff, dropout=dropout + ) + for _ in range(n_layers) + ] + ) + + # Final Projection + self.projection = nn.Linear( + self.context_length * d_model, self.prediction_length * self.n_quantiles + ) + + def forward(self, x: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + # Form Input: [Batch_Size, Context_Length, Features] + available_features = [] + target_indices = [] + current_idx = 0 + + if "encoder_cont" in x and x["encoder_cont"].size(-1) > 0: + available_features.append(x["encoder_cont"]) + current_idx += x["encoder_cont"].size(-1) + + if "target_past" in x and x["target_past"].size(-1) > 0: + target_data = x["target_past"] + if target_data.ndim == 2: + target_data = target_data.unsqueeze(-1) + n_targets = target_data.size(-1) + target_indices = list(range(current_idx, current_idx + n_targets)) + available_features.append(target_data) + + input_data = torch.cat(available_features, dim=-1) + + # RevIN + if self.use_revin: + input_data = self.revin(input_data, mode="norm") + + # Independent projection for channels: [B, C, L, d_model] + x_enc = input_data.permute(0, 2, 1).unsqueeze(-1) + x_enc = self.embedding(x_enc) + + # Process through SOFTS STAD Encoder + for layer in self.encoder: + x_enc = layer(x_enc) + + # Output projection + B, C, L, D = x_enc.shape + x_enc = x_enc.reshape(B, C, -1) + out = self.projection(x_enc) + + # Reshape for predictions + out = out.reshape(B, C, self.prediction_length, self.n_quantiles) + out = out.permute(0, 2, 1, 3) + + if self.n_quantiles == 1: + out = out.squeeze(-1) + + # De-normalize + if self.use_revin: + if out.ndim == 4: + # temporarily reshape to 3D for RevIN [B, Pred_len * quantiles, C] + out = out.permute(0, 1, 3, 2).reshape(B, -1, C) + out = self.revin(out, mode="denorm") + out = out.reshape( + B, self.prediction_length, self.n_quantiles, C + ).permute(0, 1, 3, 2) + else: + out = self.revin(out, mode="denorm") + + # Extract only the target features from output instead of + # passing all covariates to loss + if target_indices: + if out.ndim == 4: + out = out[:, :, target_indices, :] + else: + out = out[:, :, target_indices] + + if "target_scale" in x and hasattr(self, "transform_output"): + out = self.transform_output(out, x["target_scale"]) + + return {"prediction": out} + + def transform_output( + self, + y_hat: torch.Tensor | list[torch.Tensor], + target_scale: torch.Tensor | dict[str, torch.Tensor] | None, + ) -> torch.Tensor | list[torch.Tensor]: + """ + Transform the output of the model back to the original scale. + + Support: + - EncoderDecoderTimeSeriesDataModule: target_scale is a scalar tensor + """ + if target_scale is None: + return y_hat + + # EncoderDecoderTimeSeriesDataModule provides a plain tensor + if isinstance(target_scale, torch.Tensor): + scale = target_scale + while scale.dim() < y_hat.dim(): + scale = scale.unsqueeze(-1) + return y_hat * scale