Files
DeepHealth/test_model_architectures.py

248 lines
9.0 KiB
Python
Raw Permalink Normal View History

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()