colipri / tests /test_pooling.py
fepegar's picture
Mask padding tokens during text pooling (#7)
25f9f37
Raw
History Blame Contribute Delete
1.58 kB
import pytest
import torch
from colipri.pooling import AttentionPool1D
from colipri.pooling import MultiLearnedQueryAttentionPool1D
@pytest.mark.parametrize(
"pool_type",
[AttentionPool1D, MultiLearnedQueryAttentionPool1D],
)
def test_ignores_masked_tokens(pool_type):
"""Ignore padded positions when pooling language embeddings."""
torch.manual_seed(0)
pool = pool_type(embed_dim=8, num_heads=2)
valid = torch.randn(2, 3, 8)
padded = torch.cat([valid, torch.full((2, 5, 8), 1e6)], dim=1)
valid_mask = torch.ones(2, 3, dtype=torch.long)
padded_mask = torch.cat(
[valid_mask, torch.zeros(2, 5, dtype=torch.long)],
dim=1,
)
expected = pool(valid)
result = pool(padded, attention_mask=padded_mask)
torch.testing.assert_close(result, expected)
def test_rejects_incorrect_attention_mask_shape():
"""Reject an attention mask that does not match the input sequence."""
pool = AttentionPool1D(embed_dim=8, num_heads=2)
sequence = torch.randn(2, 3, 8)
attention_mask = torch.ones(2, 4, dtype=torch.long)
with pytest.raises(ValueError, match="attention_mask shape"):
pool(sequence, attention_mask=attention_mask)
def test_rejects_fully_masked_sequence():
"""Reject a sequence without any valid tokens."""
pool = AttentionPool1D(embed_dim=8, num_heads=2)
sequence = torch.randn(2, 3, 8)
attention_mask = torch.tensor([[1, 1, 1], [0, 0, 0]])
with pytest.raises(ValueError, match="at least one valid token"):
pool(sequence, attention_mask=attention_mask)