(sample_data)
| 35 | |
| 36 | |
| 37 | def test_pad_input(sample_data): |
| 38 | inputs, attention_mask, _, labels = sample_data |
| 39 | unpadded_inputs, indices, _, _, _, unpadded_labels = unpad_input(inputs, attention_mask, labels=labels) |
| 40 | |
| 41 | padded_inputs, padded_labels = pad_input(unpadded_inputs, indices, batch=2, seqlen=4, labels=unpadded_labels) |
| 42 | |
| 43 | assert padded_inputs.shape == (2, 4, 3) |
| 44 | assert torch.allclose(padded_inputs[attention_mask.bool()], unpadded_inputs) |
| 45 | assert torch.all(padded_inputs[~attention_mask.bool()] == 0) |
| 46 | assert torch.all(padded_labels[attention_mask.bool()] == unpadded_labels) |
| 47 | assert torch.all(padded_labels[~attention_mask.bool()] == -100) |
| 48 | |
| 49 | |
| 50 | def test_roundtrip(sample_data): |
nothing calls this directly
no test coverage detected