()
| 75 | |
| 76 | |
| 77 | def test_2d_input(): |
| 78 | batch, seqlen = 2, 4 |
| 79 | inputs = torch.randn(batch, seqlen) |
| 80 | attention_mask = torch.tensor([[1, 1, 1, 0], [1, 1, 0, 0]], dtype=torch.int32) |
| 81 | |
| 82 | unpadded_inputs, indices, cu_seqlens, max_seqlen, _, _ = unpad_input(inputs, attention_mask) |
| 83 | |
| 84 | assert unpadded_inputs.shape == (5,) # 5 valid tokens |
| 85 | assert indices.tolist() == [0, 1, 2, 4, 5] |
| 86 | assert cu_seqlens.tolist() == [0, 3, 5] |
| 87 | assert max_seqlen == 3 |
| 88 | |
| 89 | padded_inputs, _ = pad_input(unpadded_inputs, indices, batch=2, seqlen=4) |
| 90 | |
| 91 | assert padded_inputs.shape == (2, 4) |
| 92 | assert torch.allclose(padded_inputs[attention_mask.bool()], unpadded_inputs) |
| 93 | assert torch.all(padded_inputs[~attention_mask.bool()] == 0) |
nothing calls this directly
no test coverage detected