diff --git a/test/models/timesfm3/reference/generate.py b/test/models/timesfm3/reference/generate.py new file mode 100644 index 000000000..47c224593 --- /dev/null +++ b/test/models/timesfm3/reference/generate.py @@ -0,0 +1,169 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Generate TimesFM-3 reference fixtures with pinned Google code.""" + +from __future__ import annotations + +import argparse +import importlib +import json +import subprocess +import sys +from pathlib import Path +from typing import Any + +import torch + +UPSTREAM_REVISION = "e31dadd84cb26bd5153fde6687502b8312e918fb" +WEIGHT_RECIPE = { + "offset_step": 3, + "modulus": 19, + "center": 9, + "divisor": 32, +} +INTERNAL_CONFIG: dict[str, Any] = { + "input_patch_len": 2, + "output_patch_len": 4, + "quantiles": [0.5], + "residual_block_config": { + "hidden_dims": 4, + "output_dims": 4, + "use_bias": False, + "activation": "relu", + }, + "transformer_config": { + "num_layers": 1, + "transformer": { + "model_dims": 4, + "hidden_dims": 6, + "num_heads": 2, + "attention_norm": "rms", + "feedforward_norm": "rms", + "qk_norm": "rms", + "use_rope_seq": True, + "use_rope_var": False, + "use_bias": False, + "ff_activation": "relu", + "deterministic": True, + }, + }, + "use_variate_attention": False, + "use_stitching": True, + "use_linear_detrending": True, + "use_frozen_running_stats": False, +} +FORWARD_INPUTS: dict[str, Any] = { + "values": [[[[1.0, 3.0], [0.0, 0.0], [0.0, 0.0]]]], + "masks": [[[[False, False], [True, True], [True, True]]]], + "patch_is_target": [[[True, True, True]]], + "patch_cpm_masks": [None, [[False, True, True]]], +} + + +def _fill_state(model: torch.nn.Module) -> list[dict[str, Any]]: + state = model.state_dict() + entries = [] + for index, key in enumerate(sorted(state)): + tensor = state[key] + values = ( + ( + torch.arange(tensor.numel()).reshape(tensor.shape) + + WEIGHT_RECIPE["offset_step"] * index + ) + % WEIGHT_RECIPE["modulus"] + - WEIGHT_RECIPE["center"] + ) / WEIGHT_RECIPE["divisor"] + state[key] = values.to(dtype=tensor.dtype) + entries.append({"key": key, "shape": list(tensor.shape)}) + model.load_state_dict(state, strict=True) + return entries + + +def _upstream_revision(checkout: Path) -> str: + result = subprocess.run( + ["git", "-C", str(checkout), "rev-parse", "HEAD"], + check=True, + capture_output=True, + text=True, + ) + return result.stdout.strip() + + +def _internal_fixture(upstream: Any) -> dict[str, Any]: + model = upstream.TimesFM3Torch(**INTERNAL_CONFIG).eval() + state = _fill_state(model) + values = torch.tensor(FORWARD_INPUTS["values"]) + masks = torch.tensor(FORWARD_INPUTS["masks"]) + patch_is_target = torch.tensor(FORWARD_INPUTS["patch_is_target"]) + forward_outputs = [] + with torch.inference_mode(): + for patch_cpm_mask in FORWARD_INPUTS["patch_cpm_masks"]: + cpm_mask = ( + None + if patch_cpm_mask is None + else torch.tensor(patch_cpm_mask) + ) + logits = model( + { + "values": values, + "masks": masks, + "patch_is_target": patch_is_target, + }, + patch_cpm_mask=cpm_mask, + )["logits"] + forward_outputs.append(logits.tolist()) + + return { + "config": INTERNAL_CONFIG, + "state": state, + "forward": { + "inputs": FORWARD_INPUTS, + "outputs": forward_outputs, + }, + } + + +def generate(checkout: Path) -> dict[str, Any]: + """Generate the fixtures using an exact Google TimesFM checkout.""" + revision = _upstream_revision(checkout) + if revision != UPSTREAM_REVISION: + raise RuntimeError( + f"Expected upstream revision {UPSTREAM_REVISION}, got {revision}." + ) + + sys.path.insert(0, str(checkout / "src")) + upstream = importlib.import_module("timesfm3.torch.model") + return { + "upstream": { + "repository": "https://github.com/google-research/timesfm", + "revision": revision, + }, + "weight_recipe": WEIGHT_RECIPE, + "internal": _internal_fixture(upstream), + } + + +def main() -> None: + """Generate and write the reference fixtures.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "checkout", + type=Path, + help="Path to the pinned google-research/timesfm checkout.", + ) + parser.add_argument( + "--output", + type=Path, + default=Path(__file__).with_name("golden.json"), + ) + args = parser.parse_args() + + fixture = generate(args.checkout.resolve()) + args.output.write_text( + json.dumps(fixture, indent=2, allow_nan=False) + "\n" + ) + + +if __name__ == "__main__": + main() diff --git a/test/models/timesfm3/reference/golden.json b/test/models/timesfm3/reference/golden.json new file mode 100644 index 000000000..aca0a5cab --- /dev/null +++ b/test/models/timesfm3/reference/golden.json @@ -0,0 +1,324 @@ +{ + "upstream": { + "repository": "https://github.com/google-research/timesfm", + "revision": "e31dadd84cb26bd5153fde6687502b8312e918fb" + }, + "weight_recipe": { + "offset_step": 3, + "modulus": 19, + "center": 9, + "divisor": 32 + }, + "internal": { + "config": { + "input_patch_len": 2, + "output_patch_len": 4, + "quantiles": [ + 0.5 + ], + "residual_block_config": { + "hidden_dims": 4, + "output_dims": 4, + "use_bias": false, + "activation": "relu" + }, + "transformer_config": { + "num_layers": 1, + "transformer": { + "model_dims": 4, + "hidden_dims": 6, + "num_heads": 2, + "attention_norm": "rms", + "feedforward_norm": "rms", + "qk_norm": "rms", + "use_rope_seq": true, + "use_rope_var": false, + "use_bias": false, + "ff_activation": "relu", + "deterministic": true + } + }, + "use_variate_attention": false, + "use_stitching": true, + "use_linear_detrending": true, + "use_frozen_running_stats": false + }, + "state": [ + { + "key": "output_head.bias", + "shape": [ + 4 + ] + }, + { + "key": "output_head.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "pre_transformer_resblock.hidden_layer.weight", + "shape": [ + 4, + 12 + ] + }, + { + "key": "pre_transformer_resblock.output_layer.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "pre_transformer_resblock.residual_layer.weight", + "shape": [ + 4, + 12 + ] + }, + { + "key": "transformer_stack.layers.0.ff0.weight", + "shape": [ + 6, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.ff1.weight", + "shape": [ + 4, + 6 + ] + }, + { + "key": "transformer_stack.layers.0.post_ff_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.post_seq_attn_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.pre_ff_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.pre_seq_attn_ln.weight", + "shape": [ + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.key_ln.weight", + "shape": [ + 2 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.key_proj.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.out_proj.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.per_dim_scale.per_dim_scale", + "shape": [ + 2 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.query_ln.weight", + "shape": [ + 2 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.query_proj.weight", + "shape": [ + 4, + 4 + ] + }, + { + "key": "transformer_stack.layers.0.seq_attn.value_proj.weight", + "shape": [ + 4, + 4 + ] + } + ], + "forward": { + "inputs": { + "values": [ + [ + [ + [ + 1, + 3 + ], + [ + 0, + 0 + ], + [ + 0, + 0 + ] + ] + ] + ], + "masks": [ + [ + [ + [ + false, + false + ], + [ + true, + true + ], + [ + true, + true + ] + ] + ] + ], + "patch_is_target": [ + [ + [ + true, + true, + true + ] + ] + ], + "patch_cpm_masks": [ + null, + [ + [ + false, + true, + true + ] + ] + ] + }, + "outputs": [ + [ + [ + [ + [ + [ + 1.7138645648956299 + ], + [ + 1.7143478393554688 + ], + [ + 1.7148311138153076 + ], + [ + 1.715314269065857 + ] + ], + [ + [ + 1.692699670791626 + ], + [ + 1.7234890460968018 + ], + [ + 1.754278302192688 + ], + [ + 1.7850675582885742 + ] + ], + [ + [ + 1.6920182704925537 + ], + [ + 1.7233555316925049 + ], + [ + 1.754692792892456 + ], + [ + 1.7860300540924072 + ] + ] + ] + ] + ], + [ + [ + [ + [ + [ + 1.7138645648956299 + ], + [ + 1.7143478393554688 + ], + [ + 1.7148311138153076 + ], + [ + 1.715314269065857 + ] + ], + [ + [ + 1.635363221168518 + ], + [ + 1.657575011253357 + ], + [ + 1.6797866821289062 + ], + [ + 1.7019984722137451 + ] + ], + [ + [ + 1.6271485090255737 + ], + [ + 1.6457258462905884 + ], + [ + 1.6643033027648926 + ], + [ + 1.6828805208206177 + ] + ] + ] + ] + ] + ] + } + } +} diff --git a/test/models/timesfm3/test_core.py b/test/models/timesfm3/test_core.py index 9f9878e31..33551a91f 100644 --- a/test/models/timesfm3/test_core.py +++ b/test/models/timesfm3/test_core.py @@ -1,6 +1,10 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import json +from pathlib import Path +from typing import Any, cast + import pytest import torch @@ -436,3 +440,89 @@ def test_internal_model_applies_cpm_revin_refinement( torch.testing.assert_close(refined[:, :, 0], frozen[:, :, 0]) assert not torch.equal(refined[:, :, 1:], frozen[:, :, 1:]) + + +def _reference_fixture() -> dict[str, Any]: + path = Path(__file__).with_name("reference") / "golden.json" + return cast(dict[str, Any], json.loads(path.read_text())) + + +def _model_config(config: dict[str, Any]) -> dict[str, Any]: + model_config = dict(config) + stack_config = dict(model_config["transformer_config"]) + transformer_config = dict(stack_config["transformer"]) + for name in ("attention_norm", "feedforward_norm", "deterministic"): + transformer_config.pop(name) + stack_config["transformer"] = transformer_config + model_config["transformer_config"] = stack_config + return model_config + + +def _load_reference_weights( + model: _TimesFM3Model, + model_fixture: dict[str, Any], + recipe: dict[str, int], +) -> None: + state = model.state_dict() + actual_state = [ + {"key": key, "shape": list(tensor.shape)} + for key, tensor in sorted(state.items()) + ] + assert actual_state == model_fixture["state"] + + for index, key in enumerate(sorted(state)): + tensor = state[key] + values = ( + ( + torch.arange(tensor.numel(), device=tensor.device).reshape( + tensor.shape + ) + + recipe["offset_step"] * index + ) + % recipe["modulus"] + - recipe["center"] + ) / recipe["divisor"] + state[key] = values.to(dtype=tensor.dtype) + model.load_state_dict(state, strict=True) + + +def _golden_model(device: torch.device) -> _TimesFM3Model: + fixture = _reference_fixture() + model_fixture = fixture["internal"] + model = _TimesFM3Model( + **_model_config(model_fixture["config"]), + device=device, + ).eval() + _load_reference_weights(model, model_fixture, fixture["weight_recipe"]) + return model + + +@pytest.mark.parametrize("case_index", [0, 1]) +@withCUDA +def test_internal_model_matches_pinned_upstream( + device: torch.device, + case_index: int, +) -> None: + fixture = _reference_fixture()["internal"]["forward"] + inputs = fixture["inputs"] + model = _golden_model(device) + values = torch.tensor(inputs["values"], device=device) + masks = torch.tensor(inputs["masks"], device=device) + patch_is_target = torch.tensor(inputs["patch_is_target"], device=device) + patch_cpm_mask = inputs["patch_cpm_masks"][case_index] + cpm_mask = ( + None + if patch_cpm_mask is None + else torch.tensor(patch_cpm_mask, device=device) + ) + + with torch.inference_mode(): + actual = model( + values, + masks, + patch_is_target, + patch_cpm_mask=cpm_mask, + )["logits"] + + expected = actual.new_tensor(fixture["outputs"][case_index]) + torch.testing.assert_close(actual, expected)