(self)
| 2007 | self.assertEqual(out.shape, x.shape) |
| 2008 | |
| 2009 | def test_transformer_encoder_layer(self): |
| 2010 | # Test norm_first=True (default) |
| 2011 | layer = nn.TransformerEncoderLayer(dims=32, num_heads=4) |
| 2012 | x = mx.random.normal(shape=(2, 5, 32)) |
| 2013 | out = layer(x, mask=None) |
| 2014 | self.assertEqual(out.shape, x.shape) |
| 2015 | |
| 2016 | # Test norm_first=False |
| 2017 | layer = nn.TransformerEncoderLayer(dims=32, num_heads=4, norm_first=False) |
| 2018 | out = layer(x, mask=None) |
| 2019 | self.assertEqual(out.shape, x.shape) |
| 2020 | |
| 2021 | # Test with causal mask |
| 2022 | mask = nn.MultiHeadAttention.create_additive_causal_mask(5) |
| 2023 | out = layer(x, mask=mask) |
| 2024 | self.assertEqual(out.shape, x.shape) |
| 2025 | |
| 2026 | # Test with custom mlp_dims |
| 2027 | layer = nn.TransformerEncoderLayer(dims=32, num_heads=4, mlp_dims=64) |
| 2028 | out = layer(x, mask=None) |
| 2029 | self.assertEqual(out.shape, x.shape) |
| 2030 | |
| 2031 | def test_transformer_decoder_layer(self): |
| 2032 | dims = 32 |
nothing calls this directly
no test coverage detected