Files
DeepHealth/test_traj_mixer.py

108 lines
3.9 KiB
Python

import unittest
import torch
from backbones import GPTBlock, TrajMixer
from models import (
TRAJ_MIXER_ARCHITECTURE,
validate_traj_mixer_config,
validate_traj_mixer_state_dict,
)
from train_util import get_model_parameter_counts
class TrajMixerTest(unittest.TestCase):
def test_default_shape_parameters_and_initialization(self) -> None:
mixer = TrajMixer(
n_embd=120,
n_head=10,
hidden_group=20,
dropout=0.0,
)
x = torch.randn(2, 7, 120)
self.assertEqual(mixer(x).shape, x.shape)
self.assertEqual(sum(p.numel() for p in mixer.parameters()), 8_640)
expected = torch.eye(12).expand(10, 12, 12)
torch.testing.assert_close(mixer.group_align.detach(), expected)
self.assertEqual(tuple(mixer.gate_proj.shape), (12, 10, 20))
self.assertEqual(tuple(mixer.value_proj.shape), (12, 10, 20))
self.assertEqual(tuple(mixer.output_proj.shape), (12, 20, 10))
def test_mixer_does_not_mix_sequence_positions(self) -> None:
torch.manual_seed(0)
mixer = TrajMixer(120, n_head=10, hidden_group=20, dropout=0.0)
mixer.eval()
x = torch.randn(2, 5, 120)
changed = x.clone()
changed[:, 3, :] += torch.randn_like(changed[:, 3, :])
original_out = mixer(x)
changed_out = mixer(changed)
unchanged_positions = torch.tensor([0, 1, 2, 4])
torch.testing.assert_close(
original_out.index_select(1, unchanged_positions),
changed_out.index_select(1, unchanged_positions),
)
def test_gradients_reach_all_projection_families(self) -> None:
torch.manual_seed(1)
mixer = TrajMixer(120, n_head=10, hidden_group=20, dropout=0.0)
x = torch.randn(2, 4, 120, requires_grad=True)
mixer(x).square().mean().backward()
self.assertIsNotNone(x.grad)
for name, parameter in mixer.named_parameters():
self.assertIsNotNone(parameter.grad, name)
self.assertTrue(torch.isfinite(parameter.grad).all(), name)
def test_gpt_block_defaults_to_traj_mixer_and_standard_layer_norm(self) -> None:
block = GPTBlock(n_embd=120, n_head=10)
self.assertIsInstance(block.mlp, TrajMixer)
self.assertIsInstance(block.ln2, torch.nn.LayerNorm)
self.assertEqual(tuple(block.ln2.normalized_shape), (120,))
x = torch.randn(2, 6, 120)
self.assertEqual(block(x).shape, x.shape)
def test_architecture_marker_is_required(self) -> None:
validate_traj_mixer_config(
{"model_architecture": TRAJ_MIXER_ARCHITECTURE}
)
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
validate_traj_mixer_config({})
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
validate_traj_mixer_config({"model_architecture": "delphi_swiglu"})
def test_checkpoint_must_contain_traj_mixer_parameters(self) -> None:
block = GPTBlock(n_embd=120, n_head=10)
state_dict = {
f"blocks.0.{key}": value
for key, value in block.state_dict().items()
}
validate_traj_mixer_state_dict(state_dict)
state_dict.pop("blocks.0.mlp.group_align")
with self.assertRaisesRegex(ValueError, "not a TrajMixer checkpoint"):
validate_traj_mixer_state_dict(state_dict)
def test_invalid_group_partition_is_rejected(self) -> None:
with self.assertRaisesRegex(ValueError, "divisible"):
TrajMixer(n_embd=121, n_head=10, hidden_group=20)
def test_parameter_counts_match_traj_mixer_parameters(self) -> None:
mixer = TrajMixer(n_embd=120, n_head=10, hidden_group=20)
self.assertEqual(
get_model_parameter_counts(mixer),
{
"model_parameter_count": 8_640,
"trainable_parameter_count": 8_640,
},
)
if __name__ == "__main__":
unittest.main()