From eb5504a44d35e90a4036c5bbed0c1cad1ca67daf Mon Sep 17 00:00:00 2001 From: ZhangStudyLife <174326754+ZhangStudyLife@users.noreply.github.com> Date: Sun, 9 Aug 2026 21:16:03 +0800 Subject: [PATCH] docs: add DecoderMLP usage example --- pytorch_forecasting/models/mlp/_decodermlp.py | 47 +++++++++++++++++++ .../models/mlp/_decodermlp_pkg.py | 11 ++++- 2 files changed, 57 insertions(+), 1 deletion(-) diff --git a/pytorch_forecasting/models/mlp/_decodermlp.py b/pytorch_forecasting/models/mlp/_decodermlp.py index 1d85d6267..a84137b2a 100644 --- a/pytorch_forecasting/models/mlp/_decodermlp.py +++ b/pytorch_forecasting/models/mlp/_decodermlp.py @@ -27,6 +27,53 @@ class DecoderMLP(BaseModelWithCovariates): """MLP on the decoder. MLP that predicts output only based on information available in the decoder. + + Examples + -------- + Train on a small synthetic time series and make predictions: + + >>> import pandas as pd + >>> import torch + >>> from lightning.pytorch import Trainer + >>> from pytorch_forecasting import DecoderMLP, TimeSeriesDataSet + >>> _ = torch.manual_seed(0) + >>> data = pd.DataFrame( + ... { + ... "time_idx": range(12), + ... "series": ["A"] * 12, + ... "target": [float(i) for i in range(12)], + ... } + ... ) + >>> dataset = TimeSeriesDataSet( + ... data, + ... time_idx="time_idx", + ... target="target", + ... group_ids=["series"], + ... max_encoder_length=4, + ... max_prediction_length=2, + ... time_varying_known_reals=["time_idx"], + ... time_varying_unknown_reals=["target"], + ... ) + >>> dataloader = dataset.to_dataloader(train=True, batch_size=4, num_workers=0) + >>> model = DecoderMLP.from_dataset( + ... dataset, hidden_size=8, n_hidden_layers=1, dropout=0.0 + ... ) + >>> trainer_kwargs = dict( + ... accelerator="cpu", + ... logger=False, + ... enable_progress_bar=False, + ... enable_model_summary=False, + ... ) + >>> trainer = Trainer(fast_dev_run=True, **trainer_kwargs) + >>> trainer.fit(model, train_dataloaders=dataloader) + >>> predictions = model.predict( + ... dataset, + ... fast_dev_run=True, + ... batch_size=4, + ... trainer_kwargs=trainer_kwargs, + ... ) + >>> predictions.shape + torch.Size([4, 2]) """ @classmethod diff --git a/pytorch_forecasting/models/mlp/_decodermlp_pkg.py b/pytorch_forecasting/models/mlp/_decodermlp_pkg.py index 0da6b981c..730c3fc7c 100644 --- a/pytorch_forecasting/models/mlp/_decodermlp_pkg.py +++ b/pytorch_forecasting/models/mlp/_decodermlp_pkg.py @@ -4,7 +4,16 @@ class DecoderMLP_pkg(_BasePtForecaster): - """DecoderMLP package container.""" + """DecoderMLP package container. + + Examples + -------- + Resolve the package container to the user-facing model class: + + >>> from pytorch_forecasting.models.mlp import DecoderMLP_pkg + >>> DecoderMLP_pkg.get_cls().__name__ + 'DecoderMLP' + """ _tags = { "info:name": "DecoderMLP",