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)