Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 31 additions & 0 deletions src/timesfm3/torch/timesfm3_forecaster.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -503,6 +514,26 @@ 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.shape(ctx)[-1] == 0
or (np.ndim(ctx) == 2 and np.shape(ctx)[0] == 0)
):
raise ValueError(
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]),
("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
Expand Down
104 changes: 104 additions & 0 deletions tests/test_forecaster_query_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
# 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)),
np.array([]),
np.empty((1, 0)),
],
)
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)
)
Loading