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