(bitnet_model, batch_size, seq_len)
| 226 | ], |
| 227 | ) |
| 228 | def test_bitnet_transformer_for_various_input_shapes(bitnet_model, batch_size, seq_len): |
| 229 | tokens = torch.randint(0, 20000, (batch_size, seq_len)) |
| 230 | logits = bitnet_model(tokens) |
| 231 | assert logits.shape == (batch_size, 20000) |
| 232 | |
| 233 | |
| 234 | def test_rotary_embedding(bitnet_model, random_tensor): |
nothing calls this directly
no test coverage detected