from __future__ import annotations import logging import types import unittest from unittest.mock import patch from backend.forecasting.providers import timesfm_provider class _FakeParameter: def __init__(self, *, is_meta: bool) -> None: self.is_meta = is_meta class _FakeTimesfmModel: def __init__(self, *, is_meta: bool) -> None: self.device = "cpu" self._parameters = [_FakeParameter(is_meta=is_meta)] self.calls: list[tuple[object, ...]] = [] def parameters(self): # type: ignore[no-untyped-def] return iter(self._parameters) def load_state_dict(self, tensors, strict: bool = True, assign: bool = False) -> None: self.calls.append(("load_state_dict", tensors, strict, assign)) def to_empty(self, *, device: str) -> None: self.calls.append(("to_empty", device)) def to(self, device: str) -> None: self.calls.append(("to", device)) def eval(self) -> None: self.calls.append(("eval",)) class TimesfmProviderCompatibilityTests(unittest.TestCase): def setUp(self) -> None: timesfm_provider._TIMESFM_RUNTIME_PATCHED = False def test_runtime_patch_enables_assign_for_meta_models(self) -> None: fake_internal_module = types.SimpleNamespace( load_file=lambda path: {"path": path}, logging=logging, TimesFM_2p5_200M_torch_module=type("FakeModelModule", (), {}), ) with patch.object(timesfm_provider, "timesfm", object()), patch.object( timesfm_provider.importlib, "import_module", return_value=fake_internal_module, ): timesfm_provider._patch_timesfm_runtime_compatibility(logging.getLogger("test")) patched_cls = fake_internal_module.TimesFM_2p5_200M_torch_module fake_model = _FakeTimesfmModel(is_meta=True) patched_cls.load_checkpoint(fake_model, "checkpoint.safetensors", torch_compile=False) self.assertIn(("load_state_dict", {"path": "checkpoint.safetensors"}, True, True), fake_model.calls) self.assertIn(("to", "cpu"), fake_model.calls) self.assertIn(("eval",), fake_model.calls) self.assertTrue(getattr(patched_cls, "_aiforecast_meta_patch")) def test_runtime_patch_keeps_standard_loading_for_non_meta_models(self) -> None: fake_internal_module = types.SimpleNamespace( load_file=lambda path: {"path": path}, logging=logging, TimesFM_2p5_200M_torch_module=type("FakeModelModule", (), {}), ) with patch.object(timesfm_provider, "timesfm", object()), patch.object( timesfm_provider.importlib, "import_module", return_value=fake_internal_module, ): timesfm_provider._patch_timesfm_runtime_compatibility(logging.getLogger("test")) patched_cls = fake_internal_module.TimesFM_2p5_200M_torch_module fake_model = _FakeTimesfmModel(is_meta=False) patched_cls.load_checkpoint(fake_model, "checkpoint.safetensors", torch_compile=False) self.assertIn(("load_state_dict", {"path": "checkpoint.safetensors"}, True, False), fake_model.calls) self.assertNotIn(("to_empty", "cpu"), fake_model.calls) if __name__ == "__main__": unittest.main()