Skip to content

Commit 25860ce

Browse files
committed
fixing sliding window attn (with caveat)
1 parent 26403fa commit 25860ce

3 files changed

Lines changed: 310 additions & 9 deletions

File tree

src/fairseq2/models/transformer/_block_mask.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -63,14 +63,14 @@ def mask_fn(b: Tensor, h: Tensor, q_idx: Tensor, kv_idx: Tensor) -> Tensor:
6363
# Calculate diagonal offset
6464
d = kv_len - q_len
6565

66-
# For window_size=1, only allow the exact diagonal position
67-
if window_size == 1:
68-
return q_idx == kv_idx - d
69-
else:
70-
# For larger windows, use the range logic
71-
causal_mask = q_idx >= kv_idx - d
72-
window_mask = q_idx >= kv_idx - d - window_size + 1
73-
return causal_mask & window_mask
66+
# Apply both causal and window constraint
67+
# NOTE: There is a incompatibility here with our _create_causal_bias_tensor
68+
# function, since that requires data-dependent control flow via the q_len,
69+
# which is not currently supported by torch.compile. This is a simplified
70+
# version of that logic, which won't match exactly in all cases.
71+
causal_mask = q_idx >= kv_idx - d
72+
window_mask = kv_idx - d >= q_idx - window_size + 1
73+
return causal_mask & window_mask
7474

7575
return mask_fn
7676

tests/unit/models/transformer/test_attention.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -201,7 +201,6 @@ class TestFlexScaledDotProductAttention:
201201
(False, True, False, True, None),
202202
(False, True, True, True, None),
203203
(False, False, True, True, 1),
204-
(False, False, True, True, 2),
205204
(True, False, True, True, 1),
206205
(False, True, True, True, 1),
207206
(False, False, False, False, None),
Lines changed: 302 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,302 @@
1+
import pytest
2+
import torch
3+
from unittest.mock import Mock, patch
4+
5+
from fairseq2.models.transformer._attention_bias import IdentityBias
6+
from fairseq2.device import Device
7+
8+
from fairseq2.models.transformer._block_mask import (
9+
_causal_mask_fn,
10+
_sliding_window_causal_mask_fn,
11+
_offsets_to_doc_ids_tensor,
12+
_create_packed_mask_fn,
13+
_create_padding_mask_fn,
14+
_create_composed_mask,
15+
BlockMaskCacheKey,
16+
BlockMaskCache,
17+
)
18+
19+
20+
class TestMaskFunctions:
21+
"""Test individual mask functions."""
22+
23+
def test_causal_mask_fn(self):
24+
"""Test causal mask function behavior."""
25+
q_lens = torch.tensor([3, 2])
26+
kv_lens = torch.tensor([3, 2])
27+
mask_fn = _causal_mask_fn(q_lens, kv_lens)
28+
29+
# Test for batch 0
30+
b = torch.tensor(0)
31+
h = torch.tensor(0)
32+
33+
# Test diagonal and upper triangular positions
34+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(0)) == True
35+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(0)) == True
36+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(1)) == True
37+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(1)) == False
38+
assert mask_fn(b, h, torch.tensor(2), torch.tensor(1)) == True
39+
40+
def test_sliding_window_causal_mask_fn(self):
41+
"""Test sliding window causal mask function."""
42+
q_lens = torch.tensor([4])
43+
kv_lens = torch.tensor([4])
44+
window_size = 2
45+
mask_fn = _sliding_window_causal_mask_fn(window_size, q_lens, kv_lens)
46+
47+
b = torch.tensor(0)
48+
h = torch.tensor(0)
49+
50+
# Test window behavior
51+
assert mask_fn(b, h, torch.tensor(2), torch.tensor(1)) == True # Within window
52+
assert mask_fn(b, h, torch.tensor(2), torch.tensor(2)) == True # Diagonal
53+
assert mask_fn(b, h, torch.tensor(3), torch.tensor(1)) == False # Outside window
54+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(2)) == False # Future token
55+
56+
def test_sliding_window_size_one(self):
57+
"""Test sliding window with size 1 (diagonal only)."""
58+
q_lens = torch.tensor([3])
59+
kv_lens = torch.tensor([3])
60+
mask_fn = _sliding_window_causal_mask_fn(1, q_lens, kv_lens)
61+
62+
b = torch.tensor(0)
63+
h = torch.tensor(0)
64+
65+
# Only diagonal should be True
66+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(0)) == True
67+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(1)) == True
68+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(0)) == False
69+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(1)) == False
70+
71+
def test_offsets_to_doc_ids_tensor(self):
72+
"""Test conversion of offsets to document IDs."""
73+
offsets = torch.tensor([0, 3, 5, 8])
74+
doc_ids = _offsets_to_doc_ids_tensor(offsets)
75+
expected = torch.tensor([0, 0, 0, 1, 1, 2, 2, 2], dtype=torch.int32)
76+
assert torch.equal(doc_ids, expected)
77+
78+
def test_padding_mask_fn(self):
79+
"""Test padding mask function."""
80+
q_lens = torch.tensor([2, 3])
81+
kv_lens = torch.tensor([3, 2])
82+
mask_fn = _create_padding_mask_fn(q_lens, kv_lens)
83+
84+
b = torch.tensor(0)
85+
h = torch.tensor(0)
86+
87+
# Valid positions
88+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(0)) == True
89+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(2)) == True
90+
# Invalid positions (beyond sequence length)
91+
assert mask_fn(b, h, torch.tensor(2), torch.tensor(0)) == False
92+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(3)) == False
93+
94+
95+
class TestPackedMaskFunction:
96+
"""Test packed sequence mask function."""
97+
98+
def test_create_packed_mask_fn_basic(self):
99+
"""Test basic packed mask functionality."""
100+
seq_begin_indices = torch.tensor([0, 3, 5])
101+
keys_begin_indices = torch.tensor([0, 3, 5])
102+
103+
mask_fn = _create_packed_mask_fn(seq_begin_indices, keys_begin_indices)
104+
105+
b = torch.tensor(0)
106+
h = torch.tensor(0)
107+
108+
# Same document
109+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(1)) == True
110+
assert mask_fn(b, h, torch.tensor(3), torch.tensor(4)) == True
111+
# Different documents
112+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(3)) == False
113+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(4)) == False
114+
115+
def test_create_packed_mask_fn_with_base_mask(self):
116+
"""Test packed mask with base causal mask."""
117+
seq_begin_indices = torch.tensor([0, 2, 4])
118+
keys_begin_indices = torch.tensor([0, 2, 4])
119+
q_lens = torch.tensor([2, 2])
120+
kv_lens = torch.tensor([2, 2])
121+
122+
base_mask_fn = _causal_mask_fn(q_lens, kv_lens)
123+
mask_fn = _create_packed_mask_fn(
124+
seq_begin_indices, keys_begin_indices, base_mask_fn
125+
)
126+
127+
b = torch.tensor(0)
128+
h = torch.tensor(0)
129+
130+
# Same document, causal valid
131+
assert mask_fn(b, h, torch.tensor(1), torch.tensor(0)) == True
132+
# Same document, causal invalid
133+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(1)) == False
134+
# Different documents
135+
assert mask_fn(b, h, torch.tensor(0), torch.tensor(2)) == False
136+
137+
138+
class TestBlockMaskCache:
139+
"""Test block mask caching functionality."""
140+
141+
def test_cache_key_creation(self):
142+
"""Test cache key creation for different layouts."""
143+
cache = BlockMaskCache()
144+
145+
# Mock BatchLayout for non-packed sequences
146+
seqs_layout = Mock()
147+
seqs_layout.packed = False
148+
seqs_layout.seq_lens = [3, 4, 2]
149+
seqs_layout.max_seq_len = 4
150+
151+
keys_layout = Mock()
152+
keys_layout.packed = False
153+
keys_layout.seq_lens = [3, 4, 2]
154+
keys_layout.max_seq_len = 4
155+
156+
key = cache._create_cache_key(seqs_layout, keys_layout)
157+
assert key.batch_size == 3
158+
assert key.seqs_len == 4
159+
assert key.keys_len == 4
160+
161+
def test_cache_key_creation_packed(self):
162+
"""Test cache key creation for packed sequences."""
163+
cache = BlockMaskCache()
164+
165+
# Mock BatchLayout for packed sequences
166+
seqs_layout = Mock()
167+
seqs_layout.packed = True
168+
seqs_layout.seq_begin_indices = [0, 3, 7]
169+
170+
keys_layout = Mock()
171+
keys_layout.packed = True
172+
keys_layout.seq_begin_indices = [0, 3, 7]
173+
174+
key = cache._create_cache_key(seqs_layout, keys_layout)
175+
assert key.batch_size == 1
176+
assert key.seqs_len == 7
177+
assert key.keys_len == 7
178+
179+
def test_cache_key_hash(self):
180+
"""Test that cache keys are hashable."""
181+
key1 = BlockMaskCacheKey(batch_size=2, seqs_len=10, keys_len=10)
182+
key2 = BlockMaskCacheKey(batch_size=2, seqs_len=10, keys_len=10)
183+
key3 = BlockMaskCacheKey(batch_size=3, seqs_len=10, keys_len=10)
184+
185+
assert hash(key1) == hash(key2)
186+
assert hash(key1) != hash(key3)
187+
assert key1 == key2
188+
assert key1 != key3
189+
190+
@patch('fairseq2.models.transformer._block_mask._create_composed_mask')
191+
def test_cache_hit_and_miss(self, mock_create_mask):
192+
"""Test cache hit and miss behavior."""
193+
cache = BlockMaskCache()
194+
mock_mask = Mock()
195+
mock_create_mask.return_value = mock_mask
196+
197+
# Mock inputs
198+
bias = Mock(spec=IdentityBias)
199+
seqs_layout = Mock()
200+
seqs_layout.packed = False
201+
seqs_layout.seq_lens = [3, 4]
202+
seqs_layout.max_seq_len = 4
203+
204+
keys_layout = Mock()
205+
keys_layout.packed = False
206+
keys_layout.seq_lens = [3, 4]
207+
keys_layout.max_seq_len = 4
208+
209+
device = Mock(spec=Device)
210+
211+
# First call - cache miss
212+
result1 = cache.get_or_create_mask(bias, seqs_layout, keys_layout, device)
213+
assert result1 == mock_mask
214+
assert mock_create_mask.call_count == 1
215+
216+
# Second call - cache hit
217+
result2 = cache.get_or_create_mask(bias, seqs_layout, keys_layout, device)
218+
assert result2 == mock_mask
219+
assert mock_create_mask.call_count == 1 # Should not increase
220+
221+
def test_cache_clear(self):
222+
"""Test cache clearing."""
223+
cache = BlockMaskCache()
224+
cache._cache["test"] = "value"
225+
assert len(cache._cache) == 1
226+
227+
cache.clear()
228+
assert len(cache._cache) == 0
229+
230+
231+
class TestCreateComposedMask:
232+
"""Test the main composed mask creation function."""
233+
234+
@patch('fairseq2.models.transformer._block_mask.create_block_mask')
235+
def test_create_composed_mask_identity_bias(self, mock_create_block_mask):
236+
"""Test composed mask creation with identity bias."""
237+
mock_block_mask = Mock()
238+
mock_create_block_mask.return_value = mock_block_mask
239+
240+
bias = Mock(spec=IdentityBias)
241+
242+
# Mock BatchLayout
243+
seqs_layout = Mock()
244+
seqs_layout.packed = False
245+
seqs_layout.padded = True
246+
seqs_layout.seq_lens = [3, 4]
247+
seqs_layout.max_seq_len = 4
248+
seqs_layout.seq_lens_pt = torch.tensor([3, 4])
249+
250+
keys_layout = Mock()
251+
keys_layout.packed = False
252+
keys_layout.padded = True
253+
keys_layout.seq_lens = [3, 4]
254+
keys_layout.max_seq_len = 4
255+
keys_layout.seq_lens_pt = torch.tensor([3, 4])
256+
257+
device = Mock(spec=Device)
258+
259+
result = _create_composed_mask(bias, seqs_layout, keys_layout, device)
260+
261+
# Should create block mask with padding mask only
262+
mock_create_block_mask.assert_called_once()
263+
assert result == mock_block_mask
264+
265+
@patch('fairseq2.models.transformer._block_mask.create_block_mask')
266+
def test_create_composed_mask_no_masks_needed(self, mock_create_block_mask):
267+
"""Test when no masks are needed."""
268+
bias = Mock(spec=IdentityBias)
269+
270+
# Mock BatchLayout with no padding
271+
seqs_layout = Mock()
272+
seqs_layout.packed = False
273+
seqs_layout.padded = False
274+
275+
keys_layout = Mock()
276+
keys_layout.packed = False
277+
keys_layout.padded = False
278+
279+
device = Mock(spec=Device)
280+
281+
result = _create_composed_mask(bias, seqs_layout, keys_layout, device)
282+
283+
# Should return None when no masks are needed
284+
assert result is None
285+
mock_create_block_mask.assert_not_called()
286+
287+
def test_unsupported_bias_type(self):
288+
"""Test that unsupported bias types raise an error."""
289+
bias = Mock() # Unknown bias type
290+
291+
seqs_layout = Mock()
292+
seqs_layout.packed = False
293+
seqs_layout.padded = False
294+
295+
keys_layout = Mock()
296+
keys_layout.packed = False
297+
keys_layout.padded = False
298+
299+
device = Mock(spec=Device)
300+
301+
with pytest.raises(Exception): # Should raise NotSupportedError
302+
_create_composed_mask(bias, seqs_layout, keys_layout, device)

0 commit comments

Comments
 (0)