Skip to content

[ENH] Add forecaster support and remove _pkg class for new API - #2434

Open
phoeenniixx wants to merge 6 commits into
sktime:dmfrom
phoeenniixx:forecaster
Open

phoeenniixx wants to merge 6 commits into
sktime:dmfrom
phoeenniixx:forecaster

Conversation

@phoeenniixx

@phoeenniixx phoeenniixx commented Sep 20, 2026 •

Copy link
Copy Markdown
Member

Adds forecaster class as per #2407
stacks on #2423

Note

Need to change the target branch (even in workflow) once the base branch is merged (Datamodule).

Stack

(Still a WiP, I will update the desc once all the changes etc are made)

@phoeenniixx phoeenniixx self-assigned this Sep 20, 2026
@phoeenniixx phoeenniixx added enhancement New feature or request API design API design & software architecture module:models labels Sep 20, 2026
on:
push:
branches: [v2-dev]
branches: [dm]

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.

need to change this once #2423 is merged!

@codecov

codecov Bot commented Sep 28, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 95.42936% with 33 lines in your changes missing coverage. Please review.
✅ Project coverage is 88.42%. Comparing base (da46e88) to head (bbedfa8).

Files with missing lines Patch % Lines
...ytorch_forecasting/models/base/_base_forecaster.py 79.11% 33 Missing ⚠️
Additional details and impacted files
@@            Coverage Diff             @@
##               dm    #2434      +/-   ##
==========================================
+ Coverage   87.95%   88.42%   +0.46%     
==========================================
  Files         203      202       -1     
  Lines       11373    11652     +279     
==========================================
+ Hits        10003    10303     +300     
+ Misses       1370     1349      -21     
Flag Coverage Δ
cpu 88.42% <95.42%> (+0.46%) ⬆️
pytest 88.42% <95.42%> (+0.46%) ⬆️

Flags with carried forward coverage won't be shown. Click here to find out more.

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

…in Baseforecaster and not using skbase, .predict() now returns TimeSeries
@phoeenniixx
phoeenniixx marked this pull request as ready for review October 4, 2026 13:21
@phoeenniixx
phoeenniixx marked this pull request as draft October 4, 2026 13:22
@phoeenniixx

phoeenniixx commented Oct 4, 2026 •

Copy link
Copy Markdown
Member Author

Few open Questions:

  • Should we add inverse transform feat in this PR or in new PR (will need some updates to ScalerAdapter etc)
  • Should mode=="raw" also be somehow wrapped inside TimeSeries or should we move it to predict_raw() or similar method? And we remove mode param. predict give 2D preds, predict_quantile in quantiles etc...
  • We also need to think of the shape of data - we return 3D tensors and a list of tensors for multi-target in forward(), i.e, the raw preds are 3D or a list of 3D tensors. I think for a cleaner design we should move to 4D tensors? To prevent isinstance checks everywhere to differentiate between single target and multi-target cases?
  • We need to update the serialization in v2 (see [ENH] Add implementation of save and load to v2 #2323) as [ENH] Add implementation of save and load to v2 #2323 is not merged, should move those changes to this PR (or some PR stacked on this?) as this change would make changes to BaseForecaster and Datamodules etc anyway!

@phoeenniixx
phoeenniixx marked this pull request as ready for review October 8, 2026 18:31
@phoeenniixx

Copy link
Copy Markdown
Member Author

still need to update docs/notebooks. As we are already on a large diff here, i think we should do this in another PR?


def get_model_params(self) -> dict[str, Any]:
"""Kwargs for ``get_cls()``, with ``None`` sentinels resolved."""
from pytorch_forecasting.metrics import MAE

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think either we put the import to the top of the file or we should import it only if necessary:

if self.loss is None:
            from pytorch_forecasting.metrics import MAE
            loss = MAE()
...
return dict(...,
             loss=self.loss or loss,
             ...)

dropout=self.dropout,
norm=self.norm,
activation_class=self.activation_class,
loss=MAE() if self.loss is None else self.loss,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same here as above regarding the imports

def _build_model(self, metadata: dict):
"""Instantiates the model, either from a checkpoint or from config."""
model_cls = self.get_cls()
if self.ckpt_path:

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Here is a problem, this only updates self.model but other properites are missed when reloading from checkpoint.

E.g.

"""Existing checkpoint limitation: reload loses fitted preprocessing state."""

import warnings
from pathlib import Path
from tempfile import TemporaryDirectory

from lightning.pytorch import Trainer
import numpy as np
import pandas as pd
import torch

from pytorch_forecasting.data import TimeSeries
from pytorch_forecasting.data.data_module import EncoderDecoderTimeSeriesDataModule
from pytorch_forecasting.models.mlp import DecoderMLPForecaster

warnings.filterwarnings("ignore", message=".*experimental.*")
torch.manual_seed(0)

frame = pd.DataFrame({
    "group": np.repeat(np.arange(10), 20),
    "time": np.tile(np.arange(20), 10),
    "target": 100.0 + np.tile(np.arange(20), 10),
    "feature": np.tile(np.arange(20), 10).astype(float),
})
data = TimeSeries(
    frame, time="time", target="target", group=["group"],
    num=["feature"], known=["feature"],
)

TRAINER_KWARGS = dict(
    accelerator="cpu", max_epochs=1, limit_train_batches=1,
    limit_val_batches=1, logger=False, enable_progress_bar=False,
    enable_model_summary=False, num_sanity_val_steps=0,
)

from sklearn.preprocessing import StandardScaler

with TemporaryDirectory(prefix="forecaster-review-") as directory:
    original = DecoderMLPForecaster(
        hidden_size=4, n_hidden_layers=1,
        datamodule=EncoderDecoderTimeSeriesDataModule(
            max_encoder_length=4, max_prediction_length=3,
            target_normalizer="auto", scalers={"feature": StandardScaler()},
        ),
    )
    checkpoint = original.fit(
        data, trainer=Trainer(**TRAINER_KWARGS, default_root_dir=directory),
        ckpt_dir=directory,
    )
    assert checkpoint.is_file()
    before_loader = original._load_dataloader(data)
    loaded = DecoderMLPForecaster(ckpt_path=checkpoint)
    after_loader = loaded._load_dataloader(data)
    breakpoint()
    before_dm = before_loader.dataset.data_module
    after_dm = after_loader.dataset.data_module
    before_x = next(iter(before_loader))[0]
    after_x = next(iter(after_loader))[0]
    print("Target fitted before / after reload:",
          before_dm._target_normalizer_fitted, after_dm._target_normalizer_fitted)
    print("Features fitted before / after reload:",
          before_dm._feature_scalers_fitted, after_dm._feature_scalers_fitted)
    for key in ("target_past", "decoder_cont"):
        print(f"{key} mean before / after reload:",
              before_x[key].mean().item(), after_x[key].mean().item())
        assert not torch.allclose(before_x[key], after_x[key])
    assert before_dm._target_normalizer_fitted and before_dm._feature_scalers_fitted
    assert not after_dm._target_normalizer_fitted and not after_dm._feature_scalers_fitted
    print("BUG reproduced: checkpoint reload changes preprocessing of identical input data.")

"""Instantiates the model, either from a checkpoint or from config."""
model_cls = self.get_cls()
if self.ckpt_path:
self.model = model_cls.load_from_checkpoint(

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I actually think there is quite something wrong with the serialization of this.

"""PR #2434: refitting a loaded forecaster saves incompatible model config."""

import warnings
from pathlib import Path
from tempfile import TemporaryDirectory

from lightning.pytorch import Trainer
import numpy as np
import pandas as pd
import torch

from pytorch_forecasting.data import TimeSeries
from pytorch_forecasting.data.data_module import EncoderDecoderTimeSeriesDataModule
from pytorch_forecasting.models.mlp import DecoderMLPForecaster

warnings.filterwarnings("ignore", message=".*experimental.*")
torch.manual_seed(0)

frame = pd.DataFrame({
    "group": np.repeat(np.arange(10), 20),
    "time": np.tile(np.arange(20), 10),
    "target": 100.0 + np.tile(np.arange(20), 10),
    "feature": np.tile(np.arange(20), 10).astype(float),
})
data = TimeSeries(
    frame, time="time", target="target", group=["group"],
    num=["feature"], known=["feature"],
)

TRAINER_KWARGS = dict(
    accelerator="cpu", max_epochs=1, limit_train_batches=1,
    limit_val_batches=1, logger=False, enable_progress_bar=False,
    enable_model_summary=False, num_sanity_val_steps=0,
)

with TemporaryDirectory(prefix="forecaster-review-") as directory:
    root = Path(directory)
    original = DecoderMLPForecaster(
        hidden_size=4, n_hidden_layers=1,
        datamodule=EncoderDecoderTimeSeriesDataModule(
            max_encoder_length=4, max_prediction_length=3,
        ),
    )
    first_path = original.fit(
        data, trainer=Trainer(**TRAINER_KWARGS, default_root_dir=root),
        ckpt_dir=root / "first",
    )
    assert first_path.is_file()
    loaded = DecoderMLPForecaster(ckpt_path=first_path)
    print("Loaded wrapper hidden_size:", loaded.hidden_size)
    print("Loaded model hidden_size:", loaded.model.hidden_size)
    assert loaded.hidden_size != loaded.model.hidden_size
    second_path = loaded.fit(
        data, trainer=Trainer(**TRAINER_KWARGS, default_root_dir=root),
        ckpt_dir=root / "second",
    )
    assert second_path.is_file()
    try:
        DecoderMLPForecaster(ckpt_path=second_path)
    except RuntimeError as error:
        assert "size mismatch" in str(error)
        print("BUG reproduced: reloading the second checkpoint fails:")
        print(error)
    else:
        raise AssertionError("Reload succeeded; the reported bug was not reproduced.")

Couldn't we simply store all init parameters as well and then set them again when loading as a first try to tackle serialization?

actual_batch_size,
MAX_PREDICTION_LENGTH_TEST,
)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

here's one more error codex found. Trainer does not work properly when reused.

"""PR #2434: a reused Trainer skips training a freshly rebuilt model."""

import warnings
from pathlib import Path
from tempfile import TemporaryDirectory

from lightning.pytorch import Trainer
import numpy as np
import pandas as pd
import torch

from pytorch_forecasting.data import TimeSeries
from pytorch_forecasting.data.data_module import EncoderDecoderTimeSeriesDataModule
from pytorch_forecasting.models.mlp import DecoderMLPForecaster

warnings.filterwarnings("ignore", message=".*experimental.*")
torch.manual_seed(0)

frame = pd.DataFrame({
    "group": np.repeat(np.arange(10), 20),
    "time": np.tile(np.arange(20), 10),
    "target": 100.0 + np.tile(np.arange(20), 10),
    "feature": np.tile(np.arange(20), 10).astype(float),
})
data = TimeSeries(
    frame, time="time", target="target", group=["group"],
    num=["feature"], known=["feature"],
)

TRAINER_KWARGS = dict(
    accelerator="cpu", max_epochs=1, limit_train_batches=1,
    limit_val_batches=1, logger=False, enable_progress_bar=False,
    enable_model_summary=False, num_sanity_val_steps=0,
)

from lightning.pytorch.callbacks import Callback


class CaptureTraining(Callback):
    def __init__(self):
        self.batches = 0
        self.before = {}

    def on_fit_start(self, trainer, model):
        self.before = {key: value.clone() for key, value in model.state_dict().items()}

    def on_train_batch_end(self, trainer, model, outputs, batch, batch_idx):
        self.batches += 1


with TemporaryDirectory(prefix="forecaster-review-") as directory:
    capture = CaptureTraining()
    trainer = Trainer(
        **TRAINER_KWARGS, callbacks=[capture], enable_checkpointing=False,
        default_root_dir=directory,
    )
    forecaster = DecoderMLPForecaster(
        hidden_size=4, n_hidden_layers=1, trainer=trainer,
        datamodule=EncoderDecoderTimeSeriesDataModule(
            max_encoder_length=4, max_prediction_length=3,
        ),
    )
    forecaster.fit(data, save_ckpt=False)
    first_model = forecaster.model
    first_batches = capture.batches
    first_step = trainer.global_step
    assert first_batches > 0
    forecaster.fit(data, save_ckpt=False)
    unchanged = all(
        torch.equal(value, capture.before[key])
        for key, value in forecaster.model.state_dict().items()
    )
    print("Cumulative training batches after first / second fit:", first_batches, capture.batches)
    print("Global step after first / second fit:", first_step, trainer.global_step)
    print("Second model unchanged from initialization:", unchanged)
    print("Forecaster reports fitted:", forecaster._is_fitted)
    assert forecaster.model is not first_model
    assert capture.batches == first_batches and trainer.global_step == first_step
    assert unchanged and forecaster._is_fitted
    print("BUG reproduced: the second fit marks a fresh, untrained model as fitted.")

@@ -0,0 +1,195 @@
"""SCINet forecaster."""

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Seems like this is incompatible with the data module:

"""PR #2434: default SCINet model and automatic datamodule are incompatible."""

import warnings

import numpy as np
import pandas as pd

from pytorch_forecasting.data import TimeSeries
from pytorch_forecasting.models.scinet import SCINetForecaster

warnings.filterwarnings("ignore", message=".*experimental.*")
frame = pd.DataFrame({
    "group": np.repeat(np.arange(10), 40),
    "time": np.tile(np.arange(40), 10),
    "target": 100.0 + np.tile(np.arange(40), 10),
    "feature": np.tile(np.arange(40), 10).astype(float),
})
data = TimeSeries(
    frame, time="time", target="target", group=["group"], num=["feature"],
)
forecaster = SCINetForecaster()
try:
    # Model construction fails before a Trainer or checkpoint directory is created.
    forecaster.fit(data, save_ckpt=False)
except ValueError as error:
    assert "context_length (30)" in str(error) and "divisible" in str(error)
    print(f"BUG reproduced: {type(error).__name__}: {error}")
else:
    raise AssertionError("Fit succeeded; the reported bug was not reproduced.")

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

API design API design & software architecture enhancement New feature or request module:models

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants