MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / test_2d_input

Function test_2d_input

tests/test_padding.py:77–93  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

75
76
77def 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)

Callers

nothing calls this directly

Calls 2

unpad_inputFunction · 0.90
pad_inputFunction · 0.90

Tested by

no test coverage detected