SundialForecaster
Sundial forecaster via Hugging Face transformers.
This forecaster wraps Sundial [1], [2], [3] and exposes forecasting through the sktime forecasting interface. Calling fit loads the model and stores the observed series as forecasting context. Calling pretrain fine-tunes the model on panel or hierarchical data through the Sundial forward loss.
Sundial generates one or more sample paths. Point forecasts are computed as the empirical mean over generated samples. Quantile forecasts are computed as empirical quantiles over generated samples.
Quickstart
from sktime.forecasting.sundial import SundialForecaster
estimator = SundialForecaster(model_path='thuml/sundial-base-128m', config=None, device='cpu', dtype=None, forward_kwargs=None, random_state=None, validation_split=0.2, training_args=None, compute_metrics=None, callbacks=None)Parameters(10)
- model_pathstr, default=”thuml/sundial-base-128m”
Hugging Face repository identifier or local path to a Sundial checkpoint. If
None, a model is created fromconfigwith random weights and should be pretrained before it is used for meaningful forecasting.- configSundialConfig or dict, optional (default=None)
Model configuration used to initialize or override the Sundial model. If provided as a
dict, it is converted withSundialConfig.from_dict. Ifmodel_path=None, the model is initialized from this config with random weights. Ifmodel_pathis a checkpoint andconfigchanges parameter shapes, incompatible checkpoint weights are initialized from scratch; pretraining or fine-tuning is recommended before forecasting.- devicestr, int, or torch.device, default=”cpu”
Device on which to place the model, for example
"cpu","cuda", or"cuda:0".- dtypetorch.dtype or str, optional (default=None)
Data type used for model loading, following the
transformersdtypeconvention, for exampletorch.float16,torch.bfloat16, or"auto".- forward_kwargsdict, optional (default=None)
Additional keyword arguments forwarded to
model.generate(...)during prediction. Sundial-specific options include arguments such asnum_samplesandrevin; standard generation options supported bytransformers.GenerationMixin.generatemay also be passed. See the Sundial model card [3] and Transformers generation docs [4] for details.- random_stateint, RandomState instance or None, default=None
- Random seed used for Sundial sampling during prediction. If set, repeated predictions from the same fitted state are reproducible.
- validation_splitfloat or None, default=0.2
Fraction of data reserved for evaluation when
pretrainis used. IfNone, no evaluation dataset is created.- training_argsdict, optional (default=None)
Keyword arguments used to construct
transformers.TrainingArgumentsinpretrain[5].- compute_metricscallable or dict, optional (default=None)
Metrics callback(s) passed to
transformers.Trainer[5].- callbackslist, optional (default=None)
Trainer callbacks passed to
transformers.Trainer[5].
Examples
Simple zero-shot forecasting with thuml/sundial-base-128m:
>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.sundial import SundialForecaster
>>> y = load_airline ()
>>> forecaster = SundialForecaster ()
>>> forecaster. fit (y)
>>> y_pred = forecaster. predict (fh = [1, 2, 3 ]) Running with explicit device, dtype, and sampling settings:
>>> import torch
>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.sundial import SundialForecaster
>>> y = load_airline ()
>>> forecaster = SundialForecaster (
... device = "cuda",
... dtype = torch. bfloat16,
... forward_kwargs = { "num_samples": 20 },
... random_state = 42,
... )
>>> y_pred = forecaster. fit (y). predict (fh = [1, 2, 3 ]) Passing Sundial and Transformers generation options through forward_kwargs:
>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.sundial import SundialForecaster
>>> y = load_airline ()
>>> forecaster = SundialForecaster (
... forward_kwargs = { "num_samples": 20, "revin": False },
... )
>>> y_pred = forecaster. fit (y). predict (fh = [1, 2, 3 ]) Quantile prediction from generated samples:
>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.sundial import SundialForecaster
>>> y = load_airline ()
>>> forecaster = SundialForecaster (
... forward_kwargs = { "num_samples": 50 },
... )
>>> forecaster. fit (y)
>>> y_pred = forecaster. predict_quantiles (
... fh = [1, 2, 3 ],
... alpha = [0.1, 0.5, 0.9 ],
... ) Global training on panel data before forecasting a single series: The example below changes output_token_lens. This can initialize the affected prediction-head weights from scratch, so pretrain is run before forecasting.
>>> import torch
>>> from sktime.datasets import load_airline, load_tecator
>>> from sktime.forecasting.sundial import SundialForecaster
>>> device = "cuda" if torch. cuda. is_available () else None
>>> y_panel = load_tecator (
... return_type = "pd-multiindex",
... return_X_y = False,
... )
>>> y_panel = y_panel. drop (["class_val" ], axis = 1)
>>> y = load_airline ()
>>> forecaster = SundialForecaster (
... training_args = {
... "output_dir": "sundial-output",
... "do_train": True,
... "do_eval": True,
... "evaluation_strategy": "epoch",
... "num_train_epochs": 10,
... "per_device_train_batch_size": 8,
... "per_device_eval_batch_size": 8,
... "learning_rate": 5e-5,
... },
... config = { "output_token_lens": [8 ]},
... device = device,
... dtype = torch. bfloat16,
... validation_split = 0.3,
... )
>>> forecaster. pretrain (y_panel)
>>> forecaster. fit (y)
>>> y_pred = forecaster. predict (fh = [1, 2, 3 ])References
Sundial: A Family of Highly Capable Time Series Foundation Models: https://arxiv.org/abs/2502.00816
Sundial repository: https://github.com/thuml/Sundial
Transformers.generate(): https://huggingface.co/docs/transformers/main/en/main_classes/text_generation#transformers.GenerationMixin.generate
Trainer/TrainingArguments docs: https://huggingface.co/docs/transformers/en/main_classes/trainer