Refactor TrajMixer to single residual
This commit is contained in:
@@ -57,10 +57,16 @@ class TrajMixerTest(unittest.TestCase):
|
||||
|
||||
x = torch.randn(2, 7, 120)
|
||||
self.assertEqual(mixer(x).shape, x.shape)
|
||||
self.assertEqual(sum(p.numel() for p in mixer.parameters()), 33_164)
|
||||
|
||||
expected = torch.eye(12).expand(10, 12, 12)
|
||||
torch.testing.assert_close(mixer.group_align.detach(), expected)
|
||||
self.assertEqual(sum(p.numel() for p in mixer.parameters()), 32_040)
|
||||
self.assertFalse(hasattr(mixer, "group_align"))
|
||||
self.assertFalse(hasattr(mixer, "intra_norm"))
|
||||
self.assertFalse(hasattr(mixer, "cross_norm"))
|
||||
self.assertEqual(tuple(mixer.norm.normalized_shape), (120,))
|
||||
self.assertEqual(tuple(mixer.intra_gate_logits.shape), (10, 12))
|
||||
torch.testing.assert_close(
|
||||
torch.sigmoid(mixer.intra_gate_logits.detach()),
|
||||
torch.full((10, 12), 0.1),
|
||||
)
|
||||
self.assertEqual(mixer.intra_hidden, 48)
|
||||
self.assertEqual(
|
||||
tuple(mixer.intra_gate_proj.shape),
|
||||
@@ -78,35 +84,42 @@ class TrajMixerTest(unittest.TestCase):
|
||||
self.assertEqual(tuple(mixer.gate_proj.shape), (12, 10, 40))
|
||||
self.assertEqual(tuple(mixer.value_proj.shape), (12, 10, 40))
|
||||
self.assertEqual(tuple(mixer.output_proj.shape), (12, 40, 10))
|
||||
self.assertEqual(tuple(mixer.intra_norm.normalized_shape), (12,))
|
||||
self.assertEqual(tuple(mixer.cross_norm.normalized_shape), (10,))
|
||||
|
||||
def test_zero_output_projections_make_both_stages_identity(self) -> None:
|
||||
def test_zero_final_output_projection_makes_mixer_identity(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
mixer = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
with torch.no_grad():
|
||||
mixer.intra_output_proj.zero_()
|
||||
mixer.output_proj.zero_()
|
||||
x = torch.randn(2, 5, 120)
|
||||
torch.testing.assert_close(mixer(x), x)
|
||||
|
||||
def test_forward_matches_single_outer_residual_formula(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
mixer = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
mixer.eval()
|
||||
x = torch.randn(2, 5, 120)
|
||||
|
||||
grouped = mixer.norm(x).reshape(2, 5, 10, 12)
|
||||
intra_output = mixer._intra_mix(grouped)
|
||||
static_gate = torch.sigmoid(mixer.intra_gate_logits).view(
|
||||
1, 1, 10, 12
|
||||
)
|
||||
mixed_input = grouped + static_gate * intra_output
|
||||
update = mixer._cross_mix(mixed_input).reshape(2, 5, 120)
|
||||
|
||||
torch.testing.assert_close(mixer(x), x + update)
|
||||
|
||||
def test_intra_stage_is_independent_across_groups(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
mixer = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
mixer.eval()
|
||||
with torch.no_grad():
|
||||
mixer.output_proj.zero_()
|
||||
|
||||
grouped = torch.randn(2, 4, 10, 12)
|
||||
changed = grouped.clone()
|
||||
changed[:, :, 3, :] += torch.randn_like(changed[:, :, 3, :])
|
||||
|
||||
original_out = mixer(grouped.reshape(2, 4, 120)).reshape(
|
||||
2, 4, 10, 12
|
||||
)
|
||||
changed_out = mixer(changed.reshape(2, 4, 120)).reshape(
|
||||
2, 4, 10, 12
|
||||
)
|
||||
original_out = mixer._intra_mix(grouped)
|
||||
changed_out = mixer._intra_mix(changed)
|
||||
unchanged_groups = torch.tensor([0, 1, 2, 4, 5, 6, 7, 8, 9])
|
||||
torch.testing.assert_close(
|
||||
original_out.index_select(2, unchanged_groups),
|
||||
@@ -117,7 +130,6 @@ class TrajMixerTest(unittest.TestCase):
|
||||
mixer = TrajMixer(6, n_head=3, dropout=0.0)
|
||||
mixer.eval()
|
||||
with torch.no_grad():
|
||||
mixer.intra_output_proj.zero_()
|
||||
mixer.gate_proj.zero_()
|
||||
mixer.value_proj.zero_()
|
||||
mixer.output_proj.zero_()
|
||||
@@ -138,8 +150,8 @@ class TrajMixerTest(unittest.TestCase):
|
||||
changed = grouped.clone()
|
||||
changed[0, 0, 0, 0] = 2.0
|
||||
|
||||
original_out = mixer(grouped.reshape(1, 1, 6)).reshape(1, 1, 3, 2)
|
||||
changed_out = mixer(changed.reshape(1, 1, 6)).reshape(1, 1, 3, 2)
|
||||
original_out = mixer._cross_mix(grouped)
|
||||
changed_out = mixer._cross_mix(changed)
|
||||
|
||||
self.assertNotEqual(
|
||||
original_out[0, 0, 1, 0].item(),
|
||||
@@ -178,12 +190,13 @@ class TrajMixerTest(unittest.TestCase):
|
||||
self.assertIsNotNone(parameter.grad, name)
|
||||
self.assertTrue(torch.isfinite(parameter.grad).all(), name)
|
||||
|
||||
def test_gpt_block_delegates_both_mixer_residuals_to_traj_mixer(self) -> None:
|
||||
def test_gpt_block_delegates_single_mixer_residual_to_traj_mixer(self) -> None:
|
||||
block = GPTBlock(n_embd=120, n_head=10)
|
||||
self.assertIsInstance(block.mlp, TrajMixer)
|
||||
self.assertFalse(hasattr(block, "ln2"))
|
||||
self.assertIsInstance(block.mlp.intra_norm, torch.nn.LayerNorm)
|
||||
self.assertIsInstance(block.mlp.cross_norm, torch.nn.LayerNorm)
|
||||
self.assertIsInstance(block.mlp.norm, torch.nn.LayerNorm)
|
||||
self.assertFalse(hasattr(block.mlp, "intra_norm"))
|
||||
self.assertFalse(hasattr(block.mlp, "cross_norm"))
|
||||
|
||||
x = torch.randn(2, 6, 120)
|
||||
self.assertEqual(block(x).shape, x.shape)
|
||||
@@ -200,6 +213,14 @@ class TrajMixerTest(unittest.TestCase):
|
||||
validate_traj_mixer_config(
|
||||
{"model_architecture": "traj_mixer_v2"}
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
|
||||
validate_traj_mixer_config(
|
||||
{"model_architecture": "traj_mixer_v3"}
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
|
||||
validate_traj_mixer_config(
|
||||
{"model_architecture": "traj_mixer_v4"}
|
||||
)
|
||||
|
||||
def test_checkpoint_must_contain_traj_mixer_parameters(self) -> None:
|
||||
block = GPTBlock(n_embd=120, n_head=10)
|
||||
@@ -222,8 +243,8 @@ class TrajMixerTest(unittest.TestCase):
|
||||
self.assertEqual(
|
||||
get_model_parameter_counts(mixer),
|
||||
{
|
||||
"model_parameter_count": 33_164,
|
||||
"trainable_parameter_count": 33_164,
|
||||
"model_parameter_count": 32_040,
|
||||
"trainable_parameter_count": 32_040,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user