Fix: Correct cross-attention query routing in Post-LN TransformerDecoderLayer (#3382)

This commit is contained in:
Shantanu Suryawanshi
2026-04-07 09:16:12 -07:00
committed by GitHub
parent b98831ad0e
commit a8776b7bbd
2 changed files with 76 additions and 1 deletions
+1 -1
View File
@@ -228,7 +228,7 @@ class TransformerDecoderLayer(Module):
y = self.dropout1(y)
x = self.ln1(x + y)
y = self.cross_attention(y, memory, memory, memory_mask)
y = self.cross_attention(x, memory, memory, memory_mask)
y = self.dropout2(y)
x = self.ln2(x + y)
+75
View File
@@ -2006,6 +2006,81 @@ class TestLayers(mlx_tests.MLXTestCase):
out = attn(x, x, x)
self.assertEqual(out.shape, x.shape)
def test_transformer_encoder_layer(self):
# Test norm_first=True (default)
layer = nn.TransformerEncoderLayer(dims=32, num_heads=4)
x = mx.random.normal(shape=(2, 5, 32))
out = layer(x, mask=None)
self.assertEqual(out.shape, x.shape)
# Test norm_first=False
layer = nn.TransformerEncoderLayer(dims=32, num_heads=4, norm_first=False)
out = layer(x, mask=None)
self.assertEqual(out.shape, x.shape)
# Test with causal mask
mask = nn.MultiHeadAttention.create_additive_causal_mask(5)
out = layer(x, mask=mask)
self.assertEqual(out.shape, x.shape)
# Test with custom mlp_dims
layer = nn.TransformerEncoderLayer(dims=32, num_heads=4, mlp_dims=64)
out = layer(x, mask=None)
self.assertEqual(out.shape, x.shape)
def test_transformer_decoder_layer(self):
dims = 32
num_heads = 4
x = mx.random.normal(shape=(2, 5, dims))
memory = mx.random.normal(shape=(2, 8, dims))
# Test norm_first=True (default)
layer = nn.TransformerDecoderLayer(dims=dims, num_heads=num_heads)
out = layer(x, memory, x_mask=None, memory_mask=None)
self.assertEqual(out.shape, x.shape)
# Test norm_first=False
layer = nn.TransformerDecoderLayer(
dims=dims, num_heads=num_heads, norm_first=False
)
out = layer(x, memory, x_mask=None, memory_mask=None)
self.assertEqual(out.shape, x.shape)
# Test with masks
x_mask = nn.MultiHeadAttention.create_additive_causal_mask(5)
out = layer(x, memory, x_mask=x_mask, memory_mask=None)
self.assertEqual(out.shape, x.shape)
# Test with custom mlp_dims
layer = nn.TransformerDecoderLayer(dims=dims, num_heads=num_heads, mlp_dims=64)
out = layer(x, memory, x_mask=None, memory_mask=None)
self.assertEqual(out.shape, x.shape)
def test_transformer_encoder(self):
encoder = nn.TransformerEncoder(num_layers=2, dims=32, num_heads=4)
x = mx.random.normal(shape=(2, 5, 32))
out = encoder(x, mask=None)
self.assertEqual(out.shape, x.shape)
def test_transformer_decoder(self):
decoder = nn.TransformerDecoder(num_layers=2, dims=32, num_heads=4)
x = mx.random.normal(shape=(2, 5, 32))
memory = mx.random.normal(shape=(2, 8, 32))
out = decoder(x, memory, x_mask=None, memory_mask=None)
self.assertEqual(out.shape, x.shape)
def test_transformer(self):
model = nn.Transformer(
dims=32,
num_heads=4,
num_encoder_layers=2,
num_decoder_layers=2,
)
src = mx.random.normal(shape=(2, 8, 32))
tgt = mx.random.normal(shape=(2, 5, 32))
out = model(src, tgt, src_mask=None, tgt_mask=None, memory_mask=None)
self.assertEqual(out.shape, tgt.shape)
if __name__ == "__main__":
mlx_tests.MLXTestRunner()