Unify FFN and TrajMixer model architectures
This commit is contained in:
247
test_model_architectures.py
Normal file
247
test_model_architectures.py
Normal file
@@ -0,0 +1,247 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user