covariate support

This commit is contained in:
siriuz42
2025-10-28 22:29:40 +00:00
parent 80f1a1eafe
commit c24db2afd2
2 changed files with 234 additions and 1 deletions
+4
View File
@@ -30,6 +30,10 @@ flax = [
"jaxtyping", "jaxtyping",
"jax[cuda]" "jax[cuda]"
] ]
xreg = [
"jax[cuda]",
"scikit-learn",
]
[tool.ruff] [tool.ruff]
line-length = 88 line-length = 88
+230 -1
View File
@@ -15,8 +15,9 @@
"""TimesFM 2p5 base implementation.""" """TimesFM 2p5 base implementation."""
import dataclasses import dataclasses
from typing import Any, Callable from typing import Any, Callable, Sequence
import collections
import numpy as np import numpy as np
from .. import configs from .. import configs
@@ -25,6 +26,8 @@ ResidualBlockConfig = configs.ResidualBlockConfig
StackedTransformersConfig = configs.StackedTransformersConfig StackedTransformersConfig = configs.StackedTransformersConfig
TransformerConfig = configs.TransformerConfig TransformerConfig = configs.TransformerConfig
ForecastConfig = configs.ForecastConfig ForecastConfig = configs.ForecastConfig
Category = int | str
XRegMode = str
def strip_leading_nans(arr): def strip_leading_nans(arr):
@@ -191,3 +194,229 @@ class TimesFM_2p5:
output_points = np.concatenate(output_points, axis=0) output_points = np.concatenate(output_points, axis=0)
output_quantiles = np.concatenate(output_quantiles, axis=0) output_quantiles = np.concatenate(output_quantiles, axis=0)
return output_points[:num_inputs], output_quantiles[:num_inputs] return output_points[:num_inputs], output_quantiles[:num_inputs]
def forecast_with_covariates(
self,
inputs: list[Sequence[float]],
dynamic_numerical_covariates: dict[str, Sequence[Sequence[float]]] | None = None,
dynamic_categorical_covariates: (
dict[str, Sequence[Sequence[Category]]] | None
) = None,
static_numerical_covariates: dict[str, Sequence[float]] | None = None,
static_categorical_covariates: dict[str, Sequence[Category]] | None = None,
xreg_mode: XRegMode = "xreg + timesfm",
normalize_xreg_target_per_input: bool = True,
ridge: float = 0.0,
max_rows_per_col: int = 0,
force_on_cpu: bool = False,
):
"""Forecasts on a list of time series with covariates.
To optimize inference speed, avoid string valued categorical covariates.
Args:
inputs: A list of time series forecast contexts. Each context time series
should be in a format convertible to JTensor by `jnp.array`.
dynamic_numerical_covariates: A dict of dynamic numerical covariates.
dynamic_categorical_covariates: A dict of dynamic categorical covariates.
static_numerical_covariates: A dict of static numerical covariates.
static_categorical_covariates: A dict of static categorical covariates.
xreg_mode: one of "xreg + timesfm" or "timesfm + xreg". "xreg + timesfm"
fits a model on the residuals of the TimesFM forecast. "timesfm + xreg"
fits a model on the targets then forecasts on the residuals via TimesFM.
normalize_xreg_target_per_input: whether to normalize the xreg target per
input in the given batch.
ridge: ridge penalty for the linear model.
max_rows_per_col: max number of rows per column for the linear model.
force_on_cpu: whether to force running on cpu for the linear model.
Returns:
A tuple of two lists. The first is the outputs of the model. The second is
the outputs of the xreg.
"""
if self.forecast_config is None:
raise ValueError("Model is not compiled. Please call compile() first.")
elif not self.forecast_config.return_backcast:
raise ValueError(
"For XReg, `return_backcast` must be set to True in the forecast config. Please recompile the model."
)
from ..utils import xreg_lib
# Verify and bookkeep covariates.
if not (
dynamic_numerical_covariates
or dynamic_categorical_covariates
or static_numerical_covariates
or static_categorical_covariates
):
raise ValueError(
"At least one of dynamic_numerical_covariates,"
" dynamic_categorical_covariates, static_numerical_covariates,"
" static_categorical_covariates must be set."
)
# Track the lengths of (1) each input, (2) the part that can be used in the
# linear model, and (3) the horizon.
input_lens, train_lens, test_lens = [], [], []
for i, input_ts in enumerate(inputs):
input_len = len(input_ts)
input_lens.append(input_len)
if xreg_mode == "timesfm + xreg":
# For fitting residuals, no TimesFM forecast on the first patch.
train_lens.append(max(0, input_len - self.model.p))
elif xreg_mode == "xreg + timesfm":
train_lens.append(input_len)
else:
raise ValueError(f"Unsupported mode: {xreg_mode}")
if dynamic_numerical_covariates:
test_lens.append(
len(list(dynamic_numerical_covariates.values())[0][i]) - input_len
)
elif dynamic_categorical_covariates:
test_lens.append(
len(list(dynamic_categorical_covariates.values())[0][i]) - input_len
)
else:
test_lens.append(self.forecast_config.max_horizon)
if test_lens[-1] > self.forecast_config.max_horizon:
raise ValueError(
"Forecast horizon length inferred from the dynamic covaraites is longer than the"
f"max_horizon defined in the forecast config: {test_lens[-1]} > {self.forecast_config.max_horizon=}."
)
# Prepare the covariates into train and test.
train_dynamic_numerical_covariates = collections.defaultdict(list)
test_dynamic_numerical_covariates = collections.defaultdict(list)
train_dynamic_categorical_covariates = collections.defaultdict(list)
test_dynamic_categorical_covariates = collections.defaultdict(list)
for covariates, train_covariates, test_covariates in (
(
dynamic_numerical_covariates,
train_dynamic_numerical_covariates,
test_dynamic_numerical_covariates,
),
(
dynamic_categorical_covariates,
train_dynamic_categorical_covariates,
test_dynamic_categorical_covariates,
),
):
if not covariates:
continue
for covariate_name, covariate_values in covariates.items():
for input_len, train_len, covariate_value in zip(
input_lens, train_lens, covariate_values
):
train_covariates[covariate_name].append(
covariate_value[(input_len - train_len) : input_len]
)
test_covariates[covariate_name].append(covariate_value[input_len:])
# Fit models.
if xreg_mode == "timesfm + xreg":
# Forecast via TimesFM then fit a model on the residuals.
point_outputs, quantile_outputs = self.forecast(
horizon=self.forecast_config.max_horizon, inputs=inputs
)
targets = [
(
np.array(input_ts)[-train_len:]
- point_output[: -self.forecast_config.max_horizon][-train_len:]
)
for input_ts, point_output, train_len in zip(inputs, point_outputs, train_lens)
]
per_instance_stats = None
if normalize_xreg_target_per_input:
targets, per_instance_stats = xreg_lib.normalize(targets)
xregs = xreg_lib.BatchedInContextXRegLinear(
targets=targets,
train_lens=train_lens,
test_lens=test_lens,
train_dynamic_numerical_covariates=train_dynamic_numerical_covariates,
test_dynamic_numerical_covariates=test_dynamic_numerical_covariates,
train_dynamic_categorical_covariates=train_dynamic_categorical_covariates,
test_dynamic_categorical_covariates=test_dynamic_categorical_covariates,
static_numerical_covariates=static_numerical_covariates,
static_categorical_covariates=static_categorical_covariates,
).fit(
ridge=ridge,
one_hot_encoder_drop=None if ridge > 0 else "first",
max_rows_per_col=max_rows_per_col,
force_on_cpu=force_on_cpu,
debug_info=False,
assert_covariates=True,
assert_covariate_shapes=True,
)
if normalize_xreg_target_per_input:
xregs = xreg_lib.renormalize(xregs, per_instance_stats)
xregs = np.array(xregs)
new_point_outputs = [
(point_output[-self.forecast_config.max_horizon :][:test_len] + xreg)
for point_output, test_len, xreg in zip(point_outputs, test_lens, xregs)
]
new_quantile_outputs = [
(
quantile_output[-self.forecast_config.max_horizon :][:test_len]
+ xreg[..., None]
)
for quantile_output, test_len, xreg in zip(quantile_outputs, test_lens, xregs)
]
else:
# Fit a model on the targets then forecast on the residuals via TimesFM.
targets = [
np.array(input_ts)[-train_len:]
for input_ts, train_len in zip(inputs, train_lens)
]
per_instance_stats = None
if normalize_xreg_target_per_input:
targets, per_instance_stats = xreg_lib.normalize(targets)
xregs, xregs_on_context, _, _, _ = xreg_lib.BatchedInContextXRegLinear(
targets=targets,
train_lens=train_lens,
test_lens=test_lens,
train_dynamic_numerical_covariates=train_dynamic_numerical_covariates,
test_dynamic_numerical_covariates=test_dynamic_numerical_covariates,
train_dynamic_categorical_covariates=train_dynamic_categorical_covariates,
test_dynamic_categorical_covariates=test_dynamic_categorical_covariates,
static_numerical_covariates=static_numerical_covariates,
static_categorical_covariates=static_categorical_covariates,
).fit(
ridge=ridge,
one_hot_encoder_drop=None if ridge > 0 else "first",
max_rows_per_col=max_rows_per_col,
force_on_cpu=force_on_cpu,
debug_info=True,
assert_covariates=True,
assert_covariate_shapes=True,
)
point_outputs, quantile_outputs = self.forecast(
horizon=self.forecast_config.max_horizon,
inputs=[
target - xreg_on_context
for target, xreg_on_context in zip(targets, xregs_on_context)
],
)
new_point_outputs = [
(point_output[-self.forecast_config.max_horizon :][:test_len] + xreg)
for point_output, test_len, xreg in zip(point_outputs, test_lens, xregs)
]
new_quantile_outputs = [
(
quantile_output[-self.forecast_config.max_horizon :][:test_len]
+ xreg[..., None]
)
for quantile_output, test_len, xreg in zip(quantile_outputs, test_lens, xregs)
]
if normalize_xreg_target_per_input:
new_point_outputs = xreg_lib.renormalize(new_point_outputs, per_instance_stats)
new_quantile_outputs = xreg_lib.renormalize(
new_quantile_outputs, per_instance_stats
)
return new_point_outputs, new_quantile_outputs