248 lines
9.0 KiB
Python
248 lines
9.0 KiB
Python
|
|
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()
|