Skip to content

Commit 8214c85

Browse files
committed
Update DSV4 sparse attention
1 parent 0e4a280 commit 8214c85

5 files changed

Lines changed: 715 additions & 75 deletions

File tree

‎tests/unit_tests/cpu/test_deepseek_v4_dsa.py‎

Lines changed: 106 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,16 @@
1212
the previous index-based formulation for every compression ratio.
1313
"""
1414

15+
import inspect
1516
import unittest
1617

1718
import torch
1819
from torch.nn.attention.flex_attention import BlockMask
1920

20-
from torchtitan.models.deepseek_v4.attention import DSV4FlexInnerAttention
21+
from torchtitan.models.deepseek_v4.attention import (
22+
CompressedSparseAttention,
23+
DSV4FlexInnerAttention,
24+
)
2125
from torchtitan.models.deepseek_v4.compressor import Indexer
2226

2327

@@ -92,6 +96,107 @@ def build_dsa(ratio, window_size, block_size=128):
9296

9397

9498
class TestDSABlockMask(unittest.TestCase):
99+
def test_csa_gather_attention_reference_forward_and_backward(self):
100+
torch.manual_seed(0)
101+
seqlen, n_heads, head_dim = 16, 2, 8
102+
ratio, index_topk = 4, 3
103+
cfg = CompressedSparseAttention.Config(
104+
block_size=8,
105+
window_size=5,
106+
compress_ratio=ratio,
107+
softmax_scale=head_dim**-0.5,
108+
index_topk=index_topk,
109+
)
110+
attention = CompressedSparseAttention(cfg)
111+
inputs = [
112+
torch.randn(seqlen, n_heads, head_dim, requires_grad=True),
113+
torch.randn(seqlen, head_dim, requires_grad=True),
114+
torch.randn(seqlen // ratio, head_dim, requires_grad=True),
115+
torch.randn(seqlen, 3, 6, requires_grad=True),
116+
torch.randn(seqlen // ratio, 6, requires_grad=True),
117+
torch.randn(seqlen, 3, requires_grad=True),
118+
torch.randn(n_heads, requires_grad=True),
119+
]
120+
121+
output = attention(*inputs)
122+
self.assertEqual(output.shape, (seqlen, n_heads, head_dim))
123+
self.assertTrue(torch.isfinite(output).all())
124+
output.square().mean().backward()
125+
for tensor in (*inputs[:3], inputs[-1]):
126+
self.assertIsNotNone(tensor.grad)
127+
self.assertTrue(torch.isfinite(tensor.grad).all())
128+
129+
def test_static_mask_does_not_capture_token_pair_tensor(self):
130+
seqlen = 32768
131+
dsa = build_dsa(ratio=128, window_size=2048, block_size=128)
132+
block_mask = dsa.build_block_mask(seqlen=seqlen, device="cpu")
133+
134+
closure = inspect.getclosurevars(block_mask.mask_mod)
135+
captured_tensors = [
136+
value
137+
for value in closure.nonlocals.values()
138+
if isinstance(value, torch.Tensor)
139+
]
140+
self.assertEqual(captured_tensors, [])
141+
142+
metadata = [
143+
block_mask.kv_num_blocks,
144+
block_mask.kv_indices,
145+
block_mask.q_num_blocks,
146+
block_mask.q_indices,
147+
]
148+
metadata_bytes = sum(x.numel() * x.element_size() for x in metadata)
149+
self.assertLess(metadata_bytes, 1024 * 1024)
150+
151+
def test_static_masks_match_index_formulation(self):
152+
device = torch.device("cpu")
153+
cases = [
154+
(128, 17, 0, 32),
155+
(128, 63, 1, (32, 16)),
156+
(256, 64, 128, 32),
157+
(512, 128, 128, (64, 32)),
158+
]
159+
for seqlen, window_size, ratio, block_size in cases:
160+
with self.subTest(
161+
seqlen=seqlen,
162+
window_size=window_size,
163+
ratio=ratio,
164+
block_size=block_size,
165+
):
166+
n_cmp = seqlen // ratio if ratio > 1 else 0
167+
sink_idx = seqlen + n_cmp
168+
win = window_idxs(window_size, 1, seqlen, device)
169+
compress = (
170+
compress_idxs(ratio, 1, seqlen, device, seqlen)
171+
if ratio > 1
172+
else torch.empty((1, seqlen, 0), dtype=torch.int64)
173+
)
174+
selected = (
175+
torch.cat([win, compress], dim=-1) if compress.size(-1) else win
176+
)
177+
178+
dsa = build_dsa(ratio, window_size, block_size)
179+
block_mask = dsa.build_block_mask(seqlen=seqlen, device=device)
180+
expected = old_attended(selected, sink_idx)
181+
actual = new_attended(block_mask, seqlen, n_cmp)
182+
self.assertEqual(expected, actual)
183+
184+
bq, bk = (
185+
block_size
186+
if isinstance(block_size, tuple)
187+
else (block_size, block_size)
188+
)
189+
for q_block in range(seqlen // bq):
190+
q_slice = expected[q_block * bq : (q_block + 1) * bq]
191+
expected_blocks = {
192+
kv_idx // bk for attended in q_slice for kv_idx in attended
193+
}
194+
num_blocks = int(block_mask.kv_num_blocks[0, 0, q_block])
195+
actual_blocks = set(
196+
block_mask.kv_indices[0, 0, q_block, :num_blocks].tolist()
197+
)
198+
self.assertEqual(expected_blocks, actual_blocks)
199+
95200
def test_attended_sets_match_old_formulation(self):
96201
torch.manual_seed(0)
97202
device = torch.device("cpu")
Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
# Copyright (c) Meta Platforms, Inc. and affiliates.
2+
# All rights reserved.
3+
#
4+
# This source code is licensed under the BSD-style license found in the
5+
# LICENSE file in the root directory of this source tree.
6+
7+
"""Compare cached and freshly built DeepSeek V4 SWA/HCA attention masks."""
8+
9+
import unittest
10+
from unittest import mock
11+
12+
import torch
13+
14+
from torchtitan.models.deepseek_v4 import model_registry
15+
from torchtitan.models.deepseek_v4.attention import (
16+
dsv4_mask_key,
17+
DSV4FlexInnerAttention,
18+
)
19+
20+
21+
@unittest.skipUnless(torch.cuda.is_available(), "CUDA is unavailable")
22+
class TestDeepSeekV4AttentionCache(unittest.TestCase):
23+
def _make_model(self):
24+
# Only mask construction and parameter-free inner attention are used.
25+
with torch.device("meta"):
26+
model = model_registry("debugmodel", seq_len=512).model.build()
27+
return model
28+
29+
def _assert_outputs_and_grads_match(self, inner, cached_mask, seqlen):
30+
device = torch.device("cuda")
31+
num_heads, head_dim = 2, 64
32+
inputs = [
33+
torch.randn(seqlen, num_heads, head_dim, device=device, requires_grad=True),
34+
torch.randn(seqlen, head_dim, device=device, requires_grad=True),
35+
]
36+
if inner.compress_ratio > 1:
37+
inputs.append(
38+
torch.randn(
39+
seqlen // inner.compress_ratio,
40+
head_dim,
41+
device=device,
42+
requires_grad=True,
43+
)
44+
)
45+
inputs.append(torch.randn(num_heads, device=device, requires_grad=True))
46+
reference_inputs = [x.detach().clone().requires_grad_(True) for x in inputs]
47+
48+
fresh_mask = DSV4FlexInnerAttention.build_block_mask(
49+
inner, seqlen=seqlen, device=device
50+
)
51+
self.assertIsNot(cached_mask, fresh_mask)
52+
cached_output = inner(*inputs, attention_masks=cached_mask)
53+
reference_output = inner(*reference_inputs, attention_masks=fresh_mask)
54+
grad_output = torch.randn_like(cached_output)
55+
cached_grads = torch.autograd.grad(cached_output, inputs, grad_output)
56+
reference_grads = torch.autograd.grad(
57+
reference_output, reference_inputs, grad_output
58+
)
59+
60+
torch.testing.assert_close(
61+
cached_output, reference_output, atol=1e-5, rtol=1e-5
62+
)
63+
names = ["q", "swa_k"]
64+
if inner.compress_ratio > 1:
65+
names.append("cmp_k")
66+
names.append("attn_sink")
67+
for name, actual, expected in zip(
68+
names, cached_grads, reference_grads, strict=True
69+
):
70+
torch.testing.assert_close(
71+
actual, expected, atol=1e-5, rtol=1e-5, msg=f"grad[{name}] mismatch"
72+
)
73+
74+
def test_cached_masks_reused_across_sequence_lengths(self):
75+
model = self._make_model()
76+
masks_by_length = {}
77+
build_mask = DSV4FlexInnerAttention.build_block_mask
78+
79+
for seqlen, expected_builds in [(256, 2), (256, 0), (512, 2), (256, 0)]:
80+
with self.subTest(seqlen=seqlen, expected_builds=expected_builds):
81+
positions = torch.arange(seqlen, device="cuda")
82+
with mock.patch.object(
83+
DSV4FlexInnerAttention,
84+
"build_block_mask",
85+
autospec=True,
86+
side_effect=build_mask,
87+
) as builder:
88+
masks = model.get_attention_masks(positions)
89+
self.assertEqual(builder.call_count, expected_builds)
90+
self.assertEqual(set(masks), {"swa", "hca_128"})
91+
92+
if seqlen in masks_by_length:
93+
for key, mask in masks.items():
94+
self.assertIs(mask, masks_by_length[seqlen][key])
95+
else:
96+
for other_masks in masks_by_length.values():
97+
for key, mask in masks.items():
98+
self.assertIsNot(mask, other_masks[key])
99+
masks_by_length[seqlen] = masks
100+
self.assertEqual(len(model.mask_cache), 2 * len(masks_by_length))
101+
102+
def test_cached_masks_match_fresh_outputs_and_grads(self):
103+
torch.manual_seed(0)
104+
model = self._make_model()
105+
seqlen = 128
106+
masks = model.get_attention_masks(torch.arange(seqlen, device="cuda"))
107+
checked = set()
108+
for layer in model.layers.values():
109+
inner = layer.attention.inner_attention
110+
key = dsv4_mask_key(inner.compress_ratio)
111+
if key is not None and key not in checked:
112+
with self.subTest(mask=key):
113+
self._assert_outputs_and_grads_match(inner, masks[key], seqlen)
114+
checked.add(key)
115+
self.assertEqual(checked, {"swa", "hca_128"})
116+
117+
def test_rejects_incompatible_layers_before_deduplication(self):
118+
for field, value in [("window_size", 32), ("block_size", 64)]:
119+
for populate_cache in (False, True):
120+
with self.subTest(field=field, populate_cache=populate_cache):
121+
model = self._make_model()
122+
positions = torch.arange(256, device="cuda")
123+
if populate_cache:
124+
model.get_attention_masks(positions)
125+
inner = model.layers["1"].attention.inner_attention
126+
self.assertNotEqual(getattr(inner, field), value)
127+
setattr(inner, field, value)
128+
with self.assertRaises(AssertionError):
129+
model.get_attention_masks(positions)
130+
131+
132+
if __name__ == "__main__":
133+
unittest.main()

0 commit comments

Comments
 (0)