From e803508ad12b2d658693b9ab7ae8e6a52088765a Mon Sep 17 00:00:00 2001 From: Ayoub Chikri Date: Wed, 16 Sep 2026 19:58:41 +0200 Subject: [PATCH 1/2] Validate forecast batch metadata and covariate shapes --- src/timesfm3/torch/timesfm3_forecaster.py | 28 +++++++ tests/test_forecaster_query_validation.py | 97 +++++++++++++++++++++++ 2 files changed, 125 insertions(+) create mode 100644 tests/test_forecaster_query_validation.py diff --git a/src/timesfm3/torch/timesfm3_forecaster.py b/src/timesfm3/torch/timesfm3_forecaster.py index 5b4ca91d..e0237077 100644 --- a/src/timesfm3/torch/timesfm3_forecaster.py +++ b/src/timesfm3/torch/timesfm3_forecaster.py @@ -21,6 +21,7 @@ import math import os from collections.abc import Iterator +from numbers import Integral from typing import Any import numpy as np @@ -477,6 +478,16 @@ def predict_batch( padding_mode: str = "none", ) -> Iterator[ForecastOutput]: """Runs inference on a batch of time series with optional covariates.""" + if isinstance(horizon, bool) or not isinstance(horizon, Integral) or horizon <= 0: + raise ValueError("horizon must be a positive integer.") + for name, values in ( + ("ts_ids", ts_ids), + ("past_only_covariates", past_only_covariates), + ("past_future_covariates", past_future_covariates), + ): + if values is not None and len(values) != len(contexts): + raise ValueError(f"{name} must contain one entry per context.") + global_horizon = ( math.ceil(horizon / self.config.output_patch_length) * self.config.output_patch_length @@ -503,6 +514,23 @@ def predict_batch( pf_2d: list[np.ndarray | None] = [] for idx, ctx in enumerate(contexts): + if np.ndim(ctx) not in (1, 2) or ( + np.ndim(ctx) == 2 and np.shape(ctx)[0] == 0 + ): + raise ValueError( + f"contexts[{idx}] must be a 1D series or a 2D array with at least one variate." + ) + for name, covariate, expected_length in ( + ("past_only_covariates", po_cov_list[idx], np.shape(ctx)[-1]), + ("past_future_covariates", pf_cov_list[idx], np.shape(ctx)[-1] + horizon), + ): + if covariate is not None and ( + np.ndim(covariate) not in (1, 2) + or np.shape(covariate)[-1] != expected_length + ): + raise ValueError( + f"{name}[{idx}] must be a 1D or 2D array with time length {expected_length}." + ) target_clean = np.atleast_2d(np.array(ctx, dtype=np.float32)) po = po_cov_list[idx] po_arr = np.atleast_2d(np.array(po, dtype=np.float32)) if po is not None else None diff --git a/tests/test_forecaster_query_validation.py b/tests/test_forecaster_query_validation.py new file mode 100644 index 00000000..735d4ca5 --- /dev/null +++ b/tests/test_forecaster_query_validation.py @@ -0,0 +1,97 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Reject misaligned forecast inputs before silently truncating or decoding.""" + +from unittest import mock + +import numpy as np +import pytest +import torch + +from timesfm3.torch import timesfm3_forecaster as api +from timesfm3.torch.timesfm3_forecaster_test import _RecordingFakeModel + + +@pytest.fixture +def forecaster(): + with mock.patch.object(api.TimesFM3Forecaster, "_init_model"): + result = api.TimesFM3Forecaster( + api.ModelConfig( + per_core_batch_size=1, + input_patch_length=8, + output_patch_length=8, + median_quantile_index=1, + ) + ) + result.model = _RecordingFakeModel() + result.device = torch.device("cpu") + return result + + +@pytest.mark.parametrize( + "name", ["ts_ids", "past_only_covariates", "past_future_covariates"] +) +@pytest.mark.parametrize("length", [0, 1, 3]) +def test_batch_metadata_lengths(forecaster, name, length): + values = ["series"] * length if name == "ts_ids" else [None] * length + with pytest.raises(ValueError, match=name): + list(forecaster.predict_batch([np.arange(8)] * 2, 4, **{name: values})) + assert not forecaster.model.calls + + +@pytest.mark.parametrize( + "name,width", + [ + ("past_only_covariates", 7), + ("past_only_covariates", 9), + ("past_future_covariates", 11), + ("past_future_covariates", 13), + ], +) +def test_covariates_align_with_context_and_horizon(forecaster, name, width): + with pytest.raises(ValueError, match=name): + list(forecaster.predict_batch([np.arange(8)], 4, **{name: [np.arange(width)]})) + assert not forecaster.model.calls + + +@pytest.mark.parametrize("horizon", [0, -1, 1.5, True]) +def test_invalid_horizon(forecaster, horizon): + with pytest.raises(ValueError, match="horizon"): + list(forecaster.predict_batch([np.arange(8)], horizon)) + assert not forecaster.model.calls + + +@pytest.mark.parametrize( + "context", [np.array(1.0), np.ones((1, 1, 8)), np.ones((0, 8))] +) +def test_invalid_context_shape(forecaster, context): + with pytest.raises(ValueError, match="contexts"): + list(forecaster.predict_batch([context], 4)) + assert not forecaster.model.calls + + +def test_validation_precedes_leading_nan_trimming(forecaster): + context = np.array([np.nan, np.nan, 1, 2, 3, 4, 5, 6], np.float32) + out = forecaster.predict( + context, + 4, + past_only_covariates=np.arange(8), + past_future_covariates=np.arange(12), + padding_mode="edge", + ) + assert out.forecast.shape == (4,) + np.testing.assert_array_equal( + forecaster.model.calls[0]["past_only_covariates"][0, 0, -6:], np.arange(2, 8) + ) From b92e02e95f0e5a3b398707d442e9a2e386c02d68 Mon Sep 17 00:00:00 2001 From: Ayoub Chikri Date: Fri, 18 Sep 2026 14:19:39 +0200 Subject: [PATCH 2/2] Reject zero-length forecast contexts --- src/timesfm3/torch/timesfm3_forecaster.py | 9 ++++++--- tests/test_forecaster_query_validation.py | 9 ++++++++- 2 files changed, 14 insertions(+), 4 deletions(-) diff --git a/src/timesfm3/torch/timesfm3_forecaster.py b/src/timesfm3/torch/timesfm3_forecaster.py index e0237077..956122e4 100644 --- a/src/timesfm3/torch/timesfm3_forecaster.py +++ b/src/timesfm3/torch/timesfm3_forecaster.py @@ -514,11 +514,14 @@ def predict_batch( pf_2d: list[np.ndarray | None] = [] for idx, ctx in enumerate(contexts): - if np.ndim(ctx) not in (1, 2) or ( - np.ndim(ctx) == 2 and np.shape(ctx)[0] == 0 + if ( + np.ndim(ctx) not in (1, 2) + or np.shape(ctx)[-1] == 0 + or (np.ndim(ctx) == 2 and np.shape(ctx)[0] == 0) ): raise ValueError( - f"contexts[{idx}] must be a 1D series or a 2D array with at least one variate." + f"contexts[{idx}] must have at least one time step and, for 2D" + " inputs, at least one variate." ) for name, covariate, expected_length in ( ("past_only_covariates", po_cov_list[idx], np.shape(ctx)[-1]), diff --git a/tests/test_forecaster_query_validation.py b/tests/test_forecaster_query_validation.py index 735d4ca5..ae05bcf5 100644 --- a/tests/test_forecaster_query_validation.py +++ b/tests/test_forecaster_query_validation.py @@ -74,7 +74,14 @@ def test_invalid_horizon(forecaster, horizon): @pytest.mark.parametrize( - "context", [np.array(1.0), np.ones((1, 1, 8)), np.ones((0, 8))] + "context", + [ + np.array(1.0), + np.ones((1, 1, 8)), + np.ones((0, 8)), + np.array([]), + np.empty((1, 0)), + ], ) def test_invalid_context_shape(forecaster, context): with pytest.raises(ValueError, match="contexts"):