import unittest import torch from backbones import ( SwiGLU, TrajMixer, TrajMixerBlock, TransformerFFNBlock, build_backbone_block, ) from model_architectures import ( TRAJ_MIXER_ARCHITECTURE, TRANSFORMER_FFN_ARCHITECTURE, detect_model_architecture_from_state_dict, resolve_model_architecture, ) from models import DeepHealth def _build_block(model_architecture: str): return build_backbone_block( model_architecture, n_embd=12, n_head=3, use_time_rope=False, use_rbf_bias=False, mlp_dropout=0.0, ) def _as_model_state_dict(block: torch.nn.Module) -> dict[str, torch.Tensor]: return { f"blocks.0.{name}": value.detach().clone() for name, value in block.state_dict().items() } def _build_model( model_architecture: str | None, *, n_layer: int = 1, ) -> DeepHealth: return DeepHealth( vocab_size=8, n_embd=12, n_head=3, n_layer=n_layer, n_types=2, n_cont_types=0, n_categories=2, cont_type_ids=[], time_mode="absolute", model_architecture=model_architecture, ) class ModelArchitectureFactoryTest(unittest.TestCase): def test_factory_builds_both_architectures_with_expected_topology(self) -> None: ffn_block = _build_block(TRANSFORMER_FFN_ARCHITECTURE) self.assertIsInstance(ffn_block, TransformerFFNBlock) self.assertIsInstance(ffn_block.mlp, SwiGLU) self.assertTrue(hasattr(ffn_block, "ln1")) self.assertTrue(hasattr(ffn_block, "ln2")) traj_block = _build_block(TRAJ_MIXER_ARCHITECTURE) self.assertIsInstance(traj_block, TrajMixerBlock) self.assertIsInstance(traj_block.mlp, TrajMixer) self.assertTrue(hasattr(traj_block, "ln1")) self.assertFalse(hasattr(traj_block, "ln2")) def test_both_architectures_forward_and_backward(self) -> None: for architecture in ( TRANSFORMER_FFN_ARCHITECTURE, TRAJ_MIXER_ARCHITECTURE, ): with self.subTest(architecture=architecture): torch.manual_seed(0) block = _build_block(architecture) x = torch.randn(2, 5, 12, requires_grad=True) output = block(x) self.assertEqual(output.shape, x.shape) output.square().mean().backward() self.assertIsNotNone(x.grad) self.assertTrue(torch.isfinite(x.grad).all()) self.assertGreater(x.grad.abs().sum().item(), 0.0) self.assertIsNotNone(block.attn.qkv.weight.grad) self.assertGreater( block.attn.qkv.weight.grad.abs().sum().item(), 0.0, ) if architecture == TRANSFORMER_FFN_ARCHITECTURE: branch_parameters = ( block.mlp.w1.weight, block.mlp.w2.weight, block.mlp.w3.weight, ) else: branch_parameters = ( block.mlp.intra_gate_proj, block.mlp.intra_value_proj, block.mlp.output_proj, ) for parameter in branch_parameters: self.assertIsNotNone(parameter.grad) self.assertTrue(torch.isfinite(parameter.grad).all()) self.assertGreater(parameter.grad.abs().sum().item(), 0.0) def test_unknown_architecture_is_rejected(self) -> None: with self.assertRaises(ValueError): _build_block("unknown_architecture") with self.assertRaisesRegex(ValueError, "model_architecture is required"): _build_model(None) def test_deephealth_rejects_fewer_than_one_layer(self) -> None: for n_layer in (0, -1): with self.subTest(n_layer=n_layer): with self.assertRaisesRegex(ValueError, "n_layer must be >= 1"): _build_model( TRANSFORMER_FFN_ARCHITECTURE, n_layer=n_layer, ) def test_deephealth_uses_factory_and_strictly_reloads_both_models(self) -> None: for architecture, block_class in ( (TRANSFORMER_FFN_ARCHITECTURE, TransformerFFNBlock), (TRAJ_MIXER_ARCHITECTURE, TrajMixerBlock), ): with self.subTest(architecture=architecture): model = _build_model(architecture) self.assertEqual(model.model_architecture, architecture) self.assertIsInstance(model.blocks[0], block_class) self.assertEqual( detect_model_architecture_from_state_dict( model.state_dict() ), architecture, ) reloaded = _build_model(architecture) incompatible = reloaded.load_state_dict( model.state_dict(), strict=True, ) self.assertEqual(incompatible.missing_keys, []) self.assertEqual(incompatible.unexpected_keys, []) class ModelArchitectureResolutionTest(unittest.TestCase): def setUp(self) -> None: self.ffn_block = _build_block(TRANSFORMER_FFN_ARCHITECTURE) self.traj_block = _build_block(TRAJ_MIXER_ARCHITECTURE) self.ffn_state = _as_model_state_dict(self.ffn_block) self.traj_state = _as_model_state_dict(self.traj_block) def test_state_dict_detection_recognizes_both_architectures(self) -> None: self.assertEqual( detect_model_architecture_from_state_dict(self.ffn_state), TRANSFORMER_FFN_ARCHITECTURE, ) self.assertEqual( detect_model_architecture_from_state_dict(self.traj_state), TRAJ_MIXER_ARCHITECTURE, ) def test_explicit_markers_resolve_when_checkpoint_matches(self) -> None: for architecture, state_dict in ( (TRANSFORMER_FFN_ARCHITECTURE, self.ffn_state), (TRAJ_MIXER_ARCHITECTURE, self.traj_state), ): with self.subTest(architecture=architecture): self.assertEqual( resolve_model_architecture( {"model_architecture": architecture}, state_dict, ), architecture, ) def test_architecture_marker_is_required_for_checkpoint_loading(self) -> None: with self.assertRaisesRegex(ValueError, "model_architecture is required"): resolve_model_architecture({}, self.ffn_state) with self.assertRaisesRegex(ValueError, "model_architecture is required"): resolve_model_architecture(None, self.traj_state) def test_explicit_marker_conflicting_with_state_dict_is_rejected(self) -> None: conflicts = ( (TRANSFORMER_FFN_ARCHITECTURE, self.traj_state), (TRAJ_MIXER_ARCHITECTURE, self.ffn_state), ) for architecture, state_dict in conflicts: with self.subTest(architecture=architecture): with self.assertRaises(ValueError): resolve_model_architecture( {"model_architecture": architecture}, state_dict, ) def test_unknown_marker_and_ambiguous_state_dict_are_rejected(self) -> None: with self.assertRaises(ValueError): resolve_model_architecture( {"model_architecture": "traj_mixer_v4"} ) ambiguous_state = dict(self.ffn_state) ambiguous_state.update(self.traj_state) with self.assertRaises(ValueError): detect_model_architecture_from_state_dict(ambiguous_state) with self.assertRaises(ValueError): detect_model_architecture_from_state_dict( {"token_embedding.weight": torch.empty(2, 2)} ) def test_ffn_block_schema_is_stable_and_strictly_loadable(self) -> None: expected_keys = { "attn.time_bias_scale", "attn.qkv.weight", "attn.out_proj.weight", "attn.rbf_proj.weight", "mlp.w1.weight", "mlp.w1.bias", "mlp.w2.weight", "mlp.w2.bias", "mlp.w3.weight", "mlp.w3.bias", "ln1.weight", "ln1.bias", "ln2.weight", "ln2.bias", } state = self.ffn_block.state_dict() self.assertSetEqual(set(state), expected_keys) self.assertEqual(tuple(state["mlp.w1.weight"].shape), (30, 12)) self.assertEqual(tuple(state["mlp.w3.weight"].shape), (12, 30)) reloaded = _build_block(TRANSFORMER_FFN_ARCHITECTURE) incompatible = reloaded.load_state_dict(state, strict=True) self.assertEqual(incompatible.missing_keys, []) self.assertEqual(incompatible.unexpected_keys, []) if __name__ == "__main__": unittest.main()