Scale TrajMixer hidden width with head count
This commit is contained in:
@@ -16,23 +16,23 @@ class TrajMixerTest(unittest.TestCase):
|
||||
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)
|
||||
self.assertEqual(sum(p.numel() for p in mixer.parameters()), 15_840)
|
||||
|
||||
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))
|
||||
self.assertEqual(mixer.hidden_group, 40)
|
||||
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))
|
||||
|
||||
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 = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
mixer.eval()
|
||||
x = torch.randn(2, 5, 120)
|
||||
changed = x.clone()
|
||||
@@ -48,7 +48,7 @@ class TrajMixerTest(unittest.TestCase):
|
||||
|
||||
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)
|
||||
mixer = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
x = torch.randn(2, 4, 120, requires_grad=True)
|
||||
|
||||
mixer(x).square().mean().backward()
|
||||
@@ -90,15 +90,15 @@ class TrajMixerTest(unittest.TestCase):
|
||||
|
||||
def test_invalid_group_partition_is_rejected(self) -> None:
|
||||
with self.assertRaisesRegex(ValueError, "divisible"):
|
||||
TrajMixer(n_embd=121, n_head=10, hidden_group=20)
|
||||
TrajMixer(n_embd=121, n_head=10)
|
||||
|
||||
def test_parameter_counts_match_traj_mixer_parameters(self) -> None:
|
||||
mixer = TrajMixer(n_embd=120, n_head=10, hidden_group=20)
|
||||
mixer = TrajMixer(n_embd=120, n_head=10)
|
||||
self.assertEqual(
|
||||
get_model_parameter_counts(mixer),
|
||||
{
|
||||
"model_parameter_count": 8_640,
|
||||
"trainable_parameter_count": 8_640,
|
||||
"model_parameter_count": 15_840,
|
||||
"trainable_parameter_count": 15_840,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user