from __future__ import annotations import unittest import torch from backend.kronos_core.model.kronos import ( _compat_load_state_dict, _module_has_meta_tensors, _move_module_to_device, ) class KronosMetaCompatibilityTests(unittest.TestCase): def test_meta_module_detection(self) -> None: meta_linear = torch.nn.Linear(2, 2, device="meta") self.assertTrue(_module_has_meta_tensors(meta_linear)) def test_compat_load_state_dict_assigns_into_meta_parameters(self) -> None: source = torch.nn.Linear(2, 2) meta_linear = torch.nn.Linear(2, 2, device="meta") _compat_load_state_dict(meta_linear, source.state_dict(), strict=True) self.assertFalse(_module_has_meta_tensors(meta_linear)) self.assertTrue(torch.allclose(meta_linear.weight.detach(), source.weight.detach())) self.assertTrue(torch.allclose(meta_linear.bias.detach(), source.bias.detach())) def test_move_module_to_device_handles_meta_modules(self) -> None: meta_linear = torch.nn.Linear(2, 2, device="meta") moved = _move_module_to_device(meta_linear, "cpu") self.assertFalse(_module_has_meta_tensors(moved)) self.assertEqual(str(next(moved.parameters()).device), "cpu") if __name__ == "__main__": unittest.main()