SundialForecaster
SundialForecaster
- class 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)[source]
Sundial forecaster via Hugging Face
transformers.This forecaster wraps Sundial [1], [2], [3] and exposes forecasting through the
sktimeforecasting interface. Callingfitloads the model and stores the observed series as forecasting context. Callingpretrainfine-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.
- Parameters:
- 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].
- Attributes:
cutoffCut-off = “present time” state of forecaster.
fhForecasting horizon that was passed.
is_fittedWhether
fithas been called.stateState of the estimator.
References
[1]Sundial: A Family of Highly Capable Time Series Foundation Models: https://arxiv.org/abs/2502.00816
[2]Sundial repository: https://github.com/thuml/Sundial
[4]Transformers .generate(): https://huggingface.co/docs/transformers/main/en/main_classes/text_generation#transformers.GenerationMixin.generate
[5] (1,2,3)Trainer/TrainingArguments docs: https://huggingface.co/docs/transformers/en/main_classes/trainer
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, sopretrainis 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])
Methods
check_is_fitted([method_name])Check if the estimator has been fitted.
clone()Obtain a clone of the object with same hyper-parameters and config.
clone_tags(estimator[, tag_names])Clone tags from another object as dynamic override.
create_test_instance([parameter_set])Construct an instance of the class, using first test parameter set.
create_test_instances_and_names([parameter_set])Create list of all test instances and a list of names for them.
fit(y[, X, fh])Fit forecaster to training data.
fit_predict(y[, X, fh, X_pred])Fit and forecast time series at future horizon.
get_class_tag(tag_name[, tag_value_default])Get class tag value from class, with tag level inheritance from parents.
get_class_tags()Get class tags from class, with tag level inheritance from parent classes.
get_config()Get config flags for self.
get_fitted_params([deep])Get fitted parameters.
get_param_defaults()Get object's parameter defaults.
get_param_names([sort])Get object's parameter names.
get_params([deep])Get a dict of parameters values for this object.
get_pretrained_params([deep])Get pretrained parameters of this estimator.
get_tag(tag_name[, tag_value_default, ...])Get tag value from instance, with tag level inheritance and overrides.
get_tags()Get tags from instance, with tag level inheritance and overrides.
get_test_params([parameter_set])Return testing parameter settings for the estimator.
is_composite()Check if the object is composed of other BaseObjects.
load_from_path(serial)Load object from file location.
load_from_serial(serial)Load object from serialized memory container.
predict([fh, X])Forecast time series at future horizon.
predict_interval([fh, X, coverage])Compute/return prediction interval forecasts.
predict_proba([fh, X, marginal])Compute/return fully probabilistic forecasts.
predict_quantiles([fh, X, alpha])Compute/return quantile forecasts.
predict_residuals([y, X])Return residuals of time series forecasts.
predict_var([fh, X, cov])Compute/return variance forecasts.
pretrain(y[, X, fh])Pre-train forecaster on panel (global) data.
reset()Reset the object to a clean post-init state.
save([path, serialization_format])Save serialized self to bytes-like object or to (.zip) file.
score(y[, X, fh])Scores forecast against ground truth, using MAPE (non-symmetric).
set_config(**config_dict)Set config flags to given values.
set_params(**params)Set the parameters of this object.
set_random_state([random_state, deep, ...])Set random_state pseudo-random seed parameters for self.
set_tags(**tag_dict)Set instance level tag overrides to given values.
update(y[, X, update_params])Update cutoff value and, optionally, fitted parameters.
update_predict(y[, cv, X, update_params, ...])Make predictions and update model iteratively over the test set.
update_predict_single([y, fh, X, update_params])Update model with new data and make forecasts.

