Skip to content

FalconTSTForecaster

FalconTSTForecaster

class FalconTSTForecaster(model_path='ant-intl/Falcon-TST_Large', config=None, device_map='cpu', quantization_config=None, revin=True)[source]

Falcon-TST forecaster via Hugging Face transformers.

This forecaster wraps Falcon-TST prediction models [1], [2] from Hugging Face and exposes them through the sktime forecasting interface.

The primary workflow is fit for zero-shot inference setup, which loads the model and stores history. It does not train or fine-tune model weights. Passing model_path=None initializes a Falcon-TST model from config instead of loading pretrained weights. These random weights cannot be trained through this estimator and are mainly useful for tests or local experimentation.

Parameters:
model_pathstr, default=”ant-intl/Falcon-TST_Large”

Hugging Face repository identifier or local path to a Falcon-TST checkpoint. If None, a model is created from config.

configFalconTSTConfig or dict, optional (default=None)

Model configuration used when model_path=None. If provided as a dict, it is converted with FalconTSTConfig.from_dict. If None and model_path=None, the default FalconTSTConfig is used. This path creates random weights; the estimator does not provide training for those weights.

device_mapstr, dict, int, or torch.device, default=”cpu”

Device placement following the transformers device_map naming convention, for example "cpu", "cuda", "cuda:0", or "auto".

quantization_configtransformers.quantizers.HfQuantizer, optional

Valid quantization configuration object compatible with transformers.PreTrainedModel.from_pretrained [3].

revinbool, default=True

Whether to use RevIN normalization during Falcon-TST prediction.

Attributes:
cutoff

Cut-off = “present time” state of forecaster.

fh

Forecasting horizon that was passed.

is_fitted

Whether fit has been called.

state

State of the estimator.

Notes

  • Falcon-TST training is not supported. The estimator performs only zero-shot forecasting from a loaded or randomly initialized model.

  • Falcon-TST supports multivariate targets, handled as independent channels by the model.

  • Exogenous data and quantile prediction are not supported.

  • Loaded models are cached via a multiton helper keyed by model-loading inputs to avoid repeated model instantiation.

  • For reduced-memory loading, use quantization_config or a pre-quantized checkpoint.

References

Examples

Simple zero-shot forecasting with the default Falcon-TST checkpoint:

>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.falcon_tst import FalconTSTForecaster
>>> y = load_airline()
>>> # By default, loads ant-intl/Falcon-TST_Large.
>>> forecaster = FalconTSTForecaster()
>>> forecaster.fit(y)
>>> y_pred = forecaster.predict(fh=[1, 2, 3])

Reduced-memory inference with device placement and quantization:

>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.falcon_tst import FalconTSTForecaster
>>> from transformers import BitsAndBytesConfig
>>> y = load_airline()
>>> forecaster = FalconTSTForecaster(
...     model_path="ant-intl/Falcon-TST_Large",
...     device_map="auto",
...     quantization_config=BitsAndBytesConfig(load_in_8bit=True),
... )
>>> forecaster.fit(y)
>>> y_pred = forecaster.predict(fh=[1, 2, 3])

Randomly initialized local model, useful for tests or local experimentation. This model is not trained by fit; the weights stay random and should not be used as a trained forecaster:

>>> from sktime.forecasting.falcon_tst import FalconTSTForecaster
>>> forecaster = FalconTSTForecaster(
...     model_path=None,
...     config={
...         "num_hidden_layers": 1,
...         "hidden_size": 4,
...         "ffn_hidden_size": 8,
...         "num_attention_heads": 1,
...         "seq_length": 8,
...         "shared_patch_size": 2,
...         "patch_size_list": [4],
...         "transformer_input_layernorm": True,
...         "expert_num_layers": 1,
...         "multi_forecast_head_list": [2],
...         "autoregressive_step_list": [1],
...         "num_experts": 1,
...         "moe_router_topk": 1,
...         "moe_ffn_hidden_size": 8,
...         "moe_shared_expert_intermediate_size": 8,
...         "use_cpu_initialization": True,
...     },
... )

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.