Back to models
Forecaster

TiRex2Forecaster

Interface to the TiRex-2 zero-shot forecaster by NX-AI.

TiRex-2 is a pretrained xLSTM-based time series foundation model for zero-shot forecasting. It natively supports multivariate targets, past covariates, and known-future covariates, and outputs nine quantile levels (0.1 to 0.9) per prediction step.

The model is run in eager mode, via torch.compiler.set_stance, because the underlying package applies torch.compile unconditionally, which requires a working C++ toolchain and fails on machines without one.

Quickstart

python
from sktime.forecasting.tirex2 import TiRex2Forecaster

estimator = TiRex2Forecaster(model_path: str='NX-AI/TiRex-2', device: str='auto', revision: str=None, tta_sign_flip: bool=None, tta_diff: bool=None, hf_kwargs: dict=None, ignore_deps: bool=False)

Parameters(7)

model_pathstr, default=”NX-AI/TiRex-2”

Hugging Face repo id or local checkpoint directory.

The decontaminated variants NX-AI/TiRex-2-gifteval-zs, NX-AI/TiRex-2-gifteval-pretrain and NX-AI/TiRex-2-fevbench are gated on Hugging Face and require authentication to download.

device{“auto”, “cpu”, “cuda”, “mps”}, default=”auto”

Device used for inference. "auto" resolves to cuda, then mps, then cpu, depending on availability.

revisionstr, optional, default=None
Model repository revision, branch, tag, or commit.
tta_sign_flipbool, optional, default=None

Sign-flip test-time augmentation. None uses the checkpoint default. Roughly doubles inference cost when enabled.

tta_diffbool, optional, default=None

Postprocessor differencing. None uses the checkpoint default.

hf_kwargsdict, optional, default=None

Additional keyword arguments passed to huggingface_hub.snapshot_download.

ignore_depsbool, default=False
If True, soft dependency checks are skipped. Intended for tests and controlled environments.

Examples

>>> from sktime.datasets import load_airline
>>> from sktime.forecasting.tirex2 import TiRex2Forecaster
>>> y = load_airline ()
>>> forecaster = TiRex2Forecaster ()
>>> forecaster. fit (y)
>>> y_pred = forecaster. predict (fh = [1, 2, 3 ])

References

[2]

TiRex-2: Generalizing TiRex to Multivariate Data and Streaming, arXiv:2607.01204