Fix dual-chunk sparse fallback index overflow (#27361)

This commit is contained in:
Han Yu
2026-06-07 23:15:37 -07:00
committed by GitHub
parent 1aa5040c74
commit 3d2165a286
5 changed files with 597 additions and 205 deletions
@@ -1134,22 +1134,21 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
- prev_chunk_end_pos - prev_chunk_end_pos
) )
if intra_vertical_indices.nelement() == 0: if intra_vertical_indices.nelement() == 0:
intra_vertical_indices = torch.cat( intra_vertical_indices = _sparse_fallback_indices(
[
intra_vertical_indices,
torch.arange(
0,
k_states_intra.size(0), k_states_intra.size(0),
max(1, k_states_intra.size(0) / 5), heads_vertical_size[head_i],
dtype=torch.int32,
device=intra_vertical_indices.device, device=intra_vertical_indices.device,
),
]
) )
slash_topk = slash_topk_buffer[head_i, : heads_slash_size[head_i]] slash_topk = slash_topk_buffer[head_i, : heads_slash_size[head_i]]
intra_slash_indices = (qk.size(-1) - 1) - slash_topk[ intra_slash_indices = (qk.size(-1) - 1) - slash_topk[
slash_topk >= prev_chunk_end_pos slash_topk >= prev_chunk_end_pos
] ]
if intra_slash_indices.nelement() == 0:
intra_slash_indices = _sparse_fallback_indices(
k_states_intra.size(0),
heads_slash_size[head_i],
device=intra_vertical_indices.device,
)
# fill buffer # fill buffer
v_count = intra_vertical_indices.nelement() v_count = intra_vertical_indices.nelement()
s_count = intra_slash_indices.nelement() s_count = intra_slash_indices.nelement()
@@ -1165,17 +1164,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
] - (prev_chunk_end_pos - chunk_len) ] - (prev_chunk_end_pos - chunk_len)
# TODO: support no vertical # TODO: support no vertical
if succ_vertical_indices.nelement() == 0: if succ_vertical_indices.nelement() == 0:
succ_vertical_indices = torch.cat( succ_vertical_indices = _sparse_fallback_indices(
[
succ_vertical_indices,
torch.arange(
0,
k_states_succ.size(0), k_states_succ.size(0),
max(1, k_states_succ.size(0) / 5), heads_vertical_size[head_i],
dtype=torch.int32,
device=intra_vertical_indices.device, device=intra_vertical_indices.device,
),
]
) )
succ_slash_indices = ( succ_slash_indices = (
prev_chunk_end_pos + (qend - qbegin) - 1 prev_chunk_end_pos + (qend - qbegin) - 1
@@ -1186,17 +1178,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
) )
] ]
if succ_slash_indices.nelement() == 0: if succ_slash_indices.nelement() == 0:
succ_slash_indices = torch.cat( succ_slash_indices = _sparse_fallback_indices(
[
succ_slash_indices,
torch.arange(
0,
k_states_succ.size(0), k_states_succ.size(0),
max(1, k_states_succ.size(0) / 5), heads_slash_size[head_i],
dtype=torch.int32,
device=intra_vertical_indices.device, device=intra_vertical_indices.device,
),
]
) )
# fill buffer # fill buffer
v_count = succ_vertical_indices.nelement() v_count = succ_vertical_indices.nelement()
@@ -1214,17 +1199,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
] ]
if inter_vertical_indices.nelement() == 0: if inter_vertical_indices.nelement() == 0:
inter_vertical_indices = torch.cat( inter_vertical_indices = _sparse_fallback_indices(
[
inter_vertical_indices,
torch.arange(
0,
k_states_inter.size(0), k_states_inter.size(0),
max(1, k_states_inter.size(0) / 5), heads_vertical_size[head_i],
dtype=torch.int32,
device=intra_vertical_indices.device, device=intra_vertical_indices.device,
),
]
) )
inter_slash_indices = ( inter_slash_indices = (
prev_chunk_end_pos - chunk_len + (qend - qbegin) - 1 prev_chunk_end_pos - chunk_len + (qend - qbegin) - 1
@@ -1233,17 +1211,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
< (prev_chunk_end_pos - chunk_len + (qend - qbegin)) < (prev_chunk_end_pos - chunk_len + (qend - qbegin))
] ]
if inter_slash_indices.nelement() == 0: if inter_slash_indices.nelement() == 0:
inter_slash_indices = torch.cat( inter_slash_indices = _sparse_fallback_indices(
[
inter_slash_indices,
torch.arange(
0,
k_states_inter.size(0), k_states_inter.size(0),
max(1, k_states_inter.size(0) / 5), heads_slash_size[head_i],
dtype=torch.int32,
device=intra_vertical_indices.device, device=intra_vertical_indices.device,
),
]
) )
# fill buffer # fill buffer
v_count = inter_vertical_indices.nelement() v_count = inter_vertical_indices.nelement()
@@ -1612,6 +1583,16 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
return out, softmax_lse return out, softmax_lse
def _sparse_fallback_indices(
seq_len: int, max_count: int, device: torch.device
) -> torch.Tensor:
count = min(int(max_count), seq_len)
if count <= 0:
return torch.empty(0, dtype=torch.int64, device=device)
step = max(1, math.ceil(seq_len / count))
return torch.arange(0, seq_len, step, dtype=torch.int64, device=device)[:count]
def _vertical_slash_sparse_attention( def _vertical_slash_sparse_attention(
query: torch.Tensor, # [BATCH, N_HEADS, N_CTX, D_HEAD] query: torch.Tensor, # [BATCH, N_HEADS, N_CTX, D_HEAD]
key: torch.Tensor, # [BATCH, N_HEADS, N_KV_CTX, D_HEAD] key: torch.Tensor, # [BATCH, N_HEADS, N_KV_CTX, D_HEAD]
@@ -57,27 +57,19 @@ DUAL_CHUNK_SPARSE_THRESHOLD_GATED_CONFIG = {
} }
# Sub-context-window sparse: vertical_size + slash_size < intra K count, so # Sub-context-window sparse: vertical_size + slash_size < intra K count, so
# the kernel's vertical+slash topk genuinely prunes (the union of selected # the kernel's vertical+slash topk genuinely prunes (the union of selected
# columns + slashes does NOT cover every K column). We can't predict the # columns + slashes does NOT cover every K column). The selection is
# exact selection because it's content-aware top-k by softmax-summed scores, # deterministic but content-aware; this case compares the sgl-kernel sparse
# but we can verify the sparse path runs, produces finite output, and # attention output against a torch implementation that consumes the same
# differs from the dense reference (proving pruning happened, not silent # block/column metadata.
# fallback). See dual_chunk/README.md for the engineering paths to a strict
# correctness reference.
DUAL_CHUNK_SPARSE_SUB_WINDOW_CONFIG = { DUAL_CHUNK_SPARSE_SUB_WINDOW_CONFIG = {
**DUAL_CHUNK_CONFIG, **DUAL_CHUNK_CONFIG,
"sparse_attention_enabled": True, "sparse_attention_enabled": True,
"sparse_attention_threshold": 0, "sparse_attention_threshold": 0,
"sparse_attention_last_q": 8, "sparse_attention_last_q": 8,
"sparse_attention_config": { "sparse_attention_config": {
0: {str(head_id): ("vertical_and_slash", 8, 8, None) for head_id in range(4)} 0: {str(head_id): ("vertical_and_slash", 4, 4, None) for head_id in range(4)}
}, },
} }
# `vertical_size=8` (not 4): the production fallback at
# dual_chunk_flashattention_backend.py:1110-1122 appends
# `torch.arange(0, k_states_intra.size(0), max(1, k_states_intra.size(0)/5))`
# when a chunk gets zero vertical indices, which can produce 5 elements
# into a `vertical_size`-slot buffer. vertical_size >= 8 avoids that
# overflow path. This is a known production edge case, not a test bug.
# Unit tests run without distributed initialization. Sparse dual-chunk config # Unit tests run without distributed initialization. Sparse dual-chunk config
# lookup should see the single-rank default. # lookup should see the single-rank default.
@@ -239,6 +231,22 @@ def make_dual_chunk_sparse_threshold_gated_cases(
) )
def make_dual_chunk_sparse_sub_window_cases(
backend: str,
) -> tuple[DualChunkAttentionCase, ...]:
common = dict(backend=backend, num_heads=4, num_kv_heads=4)
return (
DualChunkAttentionCase(
name="dual_chunk_sparse_prefill_sub_window_seq128",
forward_mode=ForwardMode.EXTEND,
page_size=16,
prefix_lens=(0,),
extend_lens=(128,),
**common,
),
)
class TinyDualChunkModelConfig: class TinyDualChunkModelConfig:
def __init__( def __init__(
self, self,
@@ -555,6 +563,15 @@ class DualChunkAttentionFixture:
input_hidden: torch.Tensor input_hidden: torch.Tensor
@dataclass(frozen=True)
class _DualChunkSparseStageSelection:
stage: str
q_len: int
kv_len: int
vertical_indices: tuple[tuple[int, ...], ...]
slash_indices: tuple[tuple[int, ...], ...]
def _set_orig_seq_lens(batch: ForwardBatch, case: DualChunkAttentionCase) -> None: def _set_orig_seq_lens(batch: ForwardBatch, case: DualChunkAttentionCase) -> None:
batch.orig_seq_lens = torch.tensor( batch.orig_seq_lens = torch.tensor(
case.seq_lens, case.seq_lens,
@@ -795,6 +812,508 @@ def _dual_chunk_attention_reference(
return module.reconstruct_output(attn_output) return module.reconstruct_output(attn_output)
def _dual_chunk_sparse_fallback_indices_reference(
seq_len: int, max_count: int, device: torch.device
) -> torch.Tensor:
count = min(int(max_count), seq_len)
if count <= 0:
return torch.empty(0, dtype=torch.int64, device=device)
step = max(1, (seq_len + count - 1) // count)
return torch.arange(0, seq_len, step, dtype=torch.int64, device=device)[:count]
def _dual_chunk_sum_all_diagonal_matrix(mat: torch.Tensor) -> torch.Tensor:
h, n, m = mat.shape
zero_mat = torch.zeros((h, n, n), dtype=mat.dtype, device=mat.device)
mat_padded = torch.cat((zero_mat, mat, zero_mat), -1)
mat_strided = mat_padded.as_strided(
(1, n, n + m), (n * (2 * n + m), 2 * n + m + 1, 1)
)
return torch.sum(mat_strided, 1)[:, 1:]
def _normalise_sparse_stage_indices(
indices: torch.Tensor,
counts: torch.Tensor,
*,
descending: bool,
) -> tuple[tuple[int, ...], ...]:
indices = indices.detach().reshape(1, counts.numel(), -1)
counts = counts.detach().to(torch.int64).cpu().tolist()
stage_indices = []
for head_i, count in enumerate(counts):
head_indices = indices[0, head_i, :count].to(torch.int64)
head_indices = head_indices.sort(descending=descending).values
stage_indices.append(tuple(int(idx) for idx in head_indices.cpu().tolist()))
return tuple(stage_indices)
def _make_sparse_stage_selection(
stage: str,
q_len: int,
kv_len: int,
vertical_indices: list[torch.Tensor],
slash_indices: list[torch.Tensor],
) -> _DualChunkSparseStageSelection:
return _DualChunkSparseStageSelection(
stage=stage,
q_len=q_len,
kv_len=kv_len,
vertical_indices=tuple(
tuple(int(idx) for idx in head_indices.sort().values.cpu().tolist())
for head_indices in vertical_indices
),
slash_indices=tuple(
tuple(
int(idx)
for idx in head_indices.sort(descending=True).values.cpu().tolist()
)
for head_indices in slash_indices
),
)
def _capture_sparse_stage_selection(
stage: str,
query: torch.Tensor,
key: torch.Tensor,
vertical_indices: torch.Tensor,
slash_indices: torch.Tensor,
vertical_counts: torch.Tensor | None,
slash_counts: torch.Tensor | None,
) -> _DualChunkSparseStageSelection:
assert vertical_counts is not None
assert slash_counts is not None
return _DualChunkSparseStageSelection(
stage=stage,
q_len=query.shape[2],
kv_len=key.shape[2],
vertical_indices=_normalise_sparse_stage_indices(
vertical_indices, vertical_counts, descending=False
),
slash_indices=_normalise_sparse_stage_indices(
slash_indices, slash_counts, descending=True
),
)
def _dual_chunk_sparse_selection_reference(
fixture: DualChunkAttentionFixture,
) -> list[_DualChunkSparseStageSelection]:
"""Reference DCA's content-aware top-k split and fallback indices.
This covers the DCA-specific part of sparse prefill: vertical/slash top-k
selection, intra/succ/inter splitting, and empty-stage fallback. The lower
level 64x64 block conversion and sparse kernel math are covered separately.
"""
case = fixture.case
assert case.batch_size == 1
assert case.prefix_lens == (0,)
assert case.num_heads == case.num_kv_heads
module = fixture.reference_module
(
q,
q_succ,
q_inter,
q_succ_critical,
q_inter_critical,
k,
_,
) = module.project_dual_qkv(fixture.input_hidden)
q = q.view(-1, case.num_heads, module.head_dim)
q_succ = q_succ.view(-1, case.num_heads, module.head_dim)
q_inter = q_inter.view(-1, case.num_heads, module.head_dim)
q_succ_critical = q_succ_critical.view(-1, case.num_heads, module.head_dim)
q_inter_critical = q_inter_critical.view(-1, case.num_heads, module.head_dim)
k = k.view(-1, case.num_kv_heads, module.head_dim)
config = DUAL_CHUNK_SPARSE_SUB_WINDOW_CONFIG
chunk_len = config["chunk_size"] - config["local_size"]
softmax_scale = module.scaling
scaling_factor = (
0.1
* torch.log(
torch.tensor(case.seq_lens[0] / config["original_max_position_embeddings"])
)
+ 1.0
).clamp(min=1)
softmax_scale *= float(scaling_factor.item())
head_config = config["sparse_attention_config"][0]
heads_vertical_size = []
heads_slash_size = []
for head_i in range(case.num_heads):
ty, vertical_size, slash_size, _ = head_config[str(head_i)]
assert ty == "vertical_and_slash"
if vertical_size == 30:
vertical_size += 100
heads_vertical_size.append(vertical_size)
heads_slash_size.append(slash_size)
selections = []
k_length = k.shape[0]
begin = k_length - q.shape[0]
while begin < k_length:
prev_chunk_end_pos = (begin // chunk_len) * chunk_len
next_chunk_end_pos = prev_chunk_end_pos + chunk_len
end = min(next_chunk_end_pos, k_length)
qbegin = begin - (k_length - q.shape[0])
qend = end - (k_length - q.shape[0])
chunk_q_len = qend - qbegin
last_q_size = min(chunk_q_len, config["sparse_attention_last_q"])
q_states_intra = q[qbegin:qend]
k_states_intra = k[prev_chunk_end_pos:end]
qk_chunks = [
(q_states_intra.transpose(0, 1)[:, -last_q_size:] * softmax_scale)
@ k_states_intra.permute(1, 2, 0)
]
stage_kv_lens = {"intra": k_states_intra.size(0)}
if prev_chunk_end_pos - chunk_len >= 0:
q_states_succ_critical = q_succ_critical[qbegin:qend]
k_states_succ = k[prev_chunk_end_pos - chunk_len : prev_chunk_end_pos]
qk_chunks.append(
(
q_states_succ_critical.transpose(0, 1)[:, -last_q_size:]
* softmax_scale
)
@ k_states_succ.permute(1, 2, 0)
)
stage_kv_lens["succ"] = k_states_succ.size(0)
if prev_chunk_end_pos - chunk_len * 2 >= 0:
q_states_inter_critical = q_inter_critical[qbegin:qend]
k_states_inter = k[: prev_chunk_end_pos - chunk_len]
qk_chunks.append(
(
q_states_inter_critical.transpose(0, 1)[:, -last_q_size:]
* softmax_scale
)
@ k_states_inter.permute(1, 2, 0)
)
stage_kv_lens["inter"] = k_states_inter.size(0)
qk = torch.cat(qk_chunks[::-1], dim=-1)
arange = torch.arange(last_q_size, device=q.device)
last_q_mask = arange[:, None] >= arange[None, :]
qk[:, :, -last_q_size:] = torch.where(
last_q_mask.unsqueeze(0),
qk[:, :, -last_q_size:],
-torch.inf,
)
qk = torch.softmax(qk, dim=-1, dtype=torch.float32)
vertical = qk.sum(-2, keepdim=True)
vertical[..., :30] = torch.inf
vertical = vertical.reshape(case.num_heads, -1)
max_vertical_topk = min(vertical.shape[-1], max(heads_vertical_size))
max_slash_topk = max(heads_slash_size)
vertical_topk_buffer = torch.topk(vertical, max_vertical_topk, -1).indices
slash_topk_buffer = torch.empty(
(case.num_heads, max_slash_topk), dtype=torch.int64, device=q.device
)
current_vertical_size = [
min(head_vertical_size, max_vertical_topk)
for head_vertical_size in heads_vertical_size
]
current_slash_size = []
for head_i in range(case.num_heads):
head_score = qk[head_i : head_i + 1, :, :]
slash_scores = _dual_chunk_sum_all_diagonal_matrix(head_score)
if head_score.size(1) != 1:
slash_scores = slash_scores[..., : -last_q_size + 1]
slash_scores[..., -100:] = torch.inf
head_slash_size = min(heads_slash_size[head_i], vertical.size(-1))
current_slash_size.append(head_slash_size)
slash_topk = torch.topk(slash_scores, head_slash_size, -1).indices
slash_topk_buffer[head_i, :head_slash_size] = slash_topk.reshape(-1)
stage_vertical_indices = {stage: [] for stage in stage_kv_lens}
stage_slash_indices = {stage: [] for stage in stage_kv_lens}
for head_i in range(case.num_heads):
vertical_topk = vertical_topk_buffer[
head_i, : current_vertical_size[head_i]
]
slash_topk = slash_topk_buffer[head_i, : current_slash_size[head_i]]
intra_vertical_indices = (
vertical_topk[vertical_topk >= prev_chunk_end_pos] - prev_chunk_end_pos
)
if intra_vertical_indices.nelement() == 0:
intra_vertical_indices = _dual_chunk_sparse_fallback_indices_reference(
stage_kv_lens["intra"], current_vertical_size[head_i], q.device
)
intra_slash_indices = (qk.size(-1) - 1) - slash_topk[
slash_topk >= prev_chunk_end_pos
]
if intra_slash_indices.nelement() == 0:
intra_slash_indices = _dual_chunk_sparse_fallback_indices_reference(
stage_kv_lens["intra"], current_slash_size[head_i], q.device
)
stage_vertical_indices["intra"].append(intra_vertical_indices)
stage_slash_indices["intra"].append(intra_slash_indices)
if "succ" in stage_kv_lens:
succ_vertical_indices = vertical_topk[
(vertical_topk < prev_chunk_end_pos)
& (vertical_topk >= prev_chunk_end_pos - chunk_len)
] - (prev_chunk_end_pos - chunk_len)
if succ_vertical_indices.nelement() == 0:
succ_vertical_indices = (
_dual_chunk_sparse_fallback_indices_reference(
stage_kv_lens["succ"],
current_vertical_size[head_i],
q.device,
)
)
succ_slash_indices = (
prev_chunk_end_pos + chunk_q_len - 1
) - slash_topk[
(slash_topk >= (prev_chunk_end_pos - chunk_len))
& (slash_topk < (prev_chunk_end_pos + chunk_q_len))
]
if succ_slash_indices.nelement() == 0:
succ_slash_indices = _dual_chunk_sparse_fallback_indices_reference(
stage_kv_lens["succ"],
current_slash_size[head_i],
q.device,
)
stage_vertical_indices["succ"].append(succ_vertical_indices)
stage_slash_indices["succ"].append(succ_slash_indices)
if "inter" in stage_kv_lens:
inter_vertical_indices = vertical_topk[
vertical_topk < prev_chunk_end_pos - chunk_len
]
if inter_vertical_indices.nelement() == 0:
inter_vertical_indices = (
_dual_chunk_sparse_fallback_indices_reference(
stage_kv_lens["inter"],
current_vertical_size[head_i],
q.device,
)
)
inter_slash_indices = (
prev_chunk_end_pos - chunk_len + chunk_q_len - 1
) - slash_topk[
slash_topk < (prev_chunk_end_pos - chunk_len + chunk_q_len)
]
if inter_slash_indices.nelement() == 0:
inter_slash_indices = _dual_chunk_sparse_fallback_indices_reference(
stage_kv_lens["inter"],
current_slash_size[head_i],
q.device,
)
stage_vertical_indices["inter"].append(inter_vertical_indices)
stage_slash_indices["inter"].append(inter_slash_indices)
for stage in ("intra", "succ", "inter"):
if stage not in stage_kv_lens:
continue
selections.append(
_make_sparse_stage_selection(
stage,
chunk_q_len,
stage_kv_lens[stage],
stage_vertical_indices[stage],
stage_slash_indices[stage],
)
)
begin = end
return selections
def _assert_sparse_stage_selections_match(
testcase,
actual: list[_DualChunkSparseStageSelection],
expected: list[_DualChunkSparseStageSelection],
) -> None:
testcase.assertEqual(
len(actual),
len(expected),
f"expected {len(expected)} sparse stage calls, got {len(actual)}",
)
for index, (actual_stage, expected_stage) in enumerate(zip(actual, expected)):
testcase.assertEqual(
actual_stage,
expected_stage,
f"sparse stage selection mismatch at call {index}",
)
def _run_dual_chunk_fixture_with_sparse_selection_capture(
fixture: DualChunkAttentionFixture,
) -> tuple[torch.Tensor, list[_DualChunkSparseStageSelection]]:
original_sparse_attention = _dual_chunk_backend._vertical_slash_sparse_attention
selections = []
def capture_sparse_attention(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
vertical_indices: torch.Tensor,
slash_indices: torch.Tensor,
softmax_scale: float,
causal: bool = True,
stage: str = "intra",
block_size_M: int = 64,
block_size_N: int = 64,
vertical_indices_count: torch.Tensor | None = None,
slash_indices_count: torch.Tensor | None = None,
):
selections.append(
_capture_sparse_stage_selection(
stage,
query,
key,
vertical_indices,
slash_indices,
vertical_indices_count,
slash_indices_count,
)
)
return original_sparse_attention(
query,
key,
value,
vertical_indices,
slash_indices,
softmax_scale,
causal=causal,
stage=stage,
block_size_M=block_size_M,
block_size_N=block_size_N,
vertical_indices_count=vertical_indices_count,
slash_indices_count=slash_indices_count,
)
try:
_dual_chunk_backend._vertical_slash_sparse_attention = capture_sparse_attention
output = run_dual_chunk_fixture_eager(fixture)
finally:
_dual_chunk_backend._vertical_slash_sparse_attention = original_sparse_attention
return output, selections
def _torch_sparse_attn_metadata_mask(
block_count: torch.Tensor,
block_offset: torch.Tensor,
column_count: torch.Tensor,
column_index: torch.Tensor,
q_len: int,
kv_len: int,
*,
block_size_m: int = 64,
block_size_n: int = 64,
) -> torch.Tensor:
batch_size, num_heads, num_rows = block_count.shape
mask = torch.zeros(
(batch_size, num_heads, q_len, kv_len),
dtype=torch.bool,
device=block_count.device,
)
for batch_i in range(batch_size):
for head_i in range(num_heads):
for row_i in range(num_rows):
row_start = row_i * block_size_m
row_end = min(row_start + block_size_m, q_len)
if row_start >= row_end:
continue
for block_i in range(int(block_count[batch_i, head_i, row_i].item())):
col_start = int(
block_offset[batch_i, head_i, row_i, block_i].item()
)
col_end = min(col_start + block_size_n, kv_len)
if 0 <= col_start < col_end:
mask[batch_i, head_i, row_start:row_end, col_start:col_end] = (
True
)
col_count = int(column_count[batch_i, head_i, row_i].item())
cols = column_index[batch_i, head_i, row_i, :col_count].to(torch.long)
cols = cols[(cols >= 0) & (cols < kv_len)]
if cols.numel() > 0:
mask[batch_i, head_i, row_start:row_end, cols] = True
return mask
def _torch_sparse_attn_func(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_count: torch.Tensor,
block_offset: torch.Tensor,
column_count: torch.Tensor,
column_index: torch.Tensor,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
causal: bool = False,
softcap: float = 0.0,
alibi_slopes: torch.Tensor | None = None,
deterministic: bool = False,
return_attn_probs: bool = False,
*,
return_softmax_lse: bool = False,
out: torch.Tensor | None = None,
):
assert dropout_p == 0.0
assert softcap == 0.0
assert alibi_slopes is None
assert not deterministic
assert not return_attn_probs
assert out is None
if softmax_scale is None:
softmax_scale = q.shape[-1] ** -0.5
dtype = q.dtype
_, q_len, num_heads, _ = q.shape
kv_len = k.shape[1]
if k.shape[2] != num_heads:
group_size = num_heads // k.shape[2]
k = torch.repeat_interleave(k, group_size, dim=2)
v = torch.repeat_interleave(v, group_size, dim=2)
sparse_mask = _torch_sparse_attn_metadata_mask(
block_count, block_offset, column_count, column_index, q_len, kv_len
)
if causal:
q_pos = torch.arange(q_len, device=q.device) + (kv_len - q_len)
k_pos = torch.arange(kv_len, device=q.device)
sparse_mask &= k_pos.view(1, 1, 1, kv_len) <= q_pos.view(1, 1, q_len, 1)
scores = torch.einsum("bqhd,bkhd->bhqk", q.float(), k.float()) * softmax_scale
scores = scores.masked_fill(~sparse_mask, -torch.inf)
softmax_lse = torch.logsumexp(scores, dim=-1)
valid_rows = sparse_mask.any(dim=-1)
probs = torch.softmax(scores, dim=-1)
probs = torch.where(valid_rows.unsqueeze(-1), probs, torch.zeros_like(probs))
output = torch.einsum("bhqk,bkhd->bqhd", probs, v.float()).to(dtype)
if return_softmax_lse:
return output, softmax_lse
return output
def _run_dual_chunk_fixture_with_torch_sparse_kernel(
fixture: DualChunkAttentionFixture,
) -> torch.Tensor:
original_sparse_attn_func = _dual_chunk_backend.sparse_attn_func
try:
_dual_chunk_backend.sparse_attn_func = _torch_sparse_attn_func
return run_dual_chunk_fixture_eager(fixture)
finally:
_dual_chunk_backend.sparse_attn_func = original_sparse_attn_func
def run_dual_chunk_attention_case( def run_dual_chunk_attention_case(
testcase, testcase,
case: DualChunkAttentionCase, case: DualChunkAttentionCase,
@@ -876,22 +1395,7 @@ def run_dual_chunk_sparse_sub_window_case(
dtype: torch.dtype = DEFAULT_DTYPE, dtype: torch.dtype = DEFAULT_DTYPE,
device: str = DEFAULT_DEVICE, device: str = DEFAULT_DEVICE,
) -> None: ) -> None:
"""Smoke test for genuine sub-context-window sparse pruning. """Correctness test for genuine sub-context-window sparse pruning."""
The vertical+slash topk in `_dual_chunk_flash_attn_prefill_func` is
content-aware (per-head top-k by softmax-summed attention scores), so we
can't predict the exact v_idx/s_idx and thus can't build a strict
PyTorch reference without re-implementing ~300 lines of inline production
logic (see `dual_chunk/README.md` for the engineering paths). This case
instead verifies:
1. The sparse path runs without crash on a sub-window config (4 vertical
+ 4 slash, intra K count > 8).
2. The output is finite (no NaN/inf).
3. The output shape matches the dense reference.
4. The output **differs** from the dense reference — proving the kernel
genuinely pruned rather than silently falling back to dense.
"""
fixture = build_dual_chunk_attention_fixture( fixture = build_dual_chunk_attention_fixture(
testcase, testcase,
case, case,
@@ -902,36 +1406,15 @@ def run_dual_chunk_sparse_sub_window_case(
device=device, device=device,
dual_chunk_attention_config=DUAL_CHUNK_SPARSE_SUB_WINDOW_CONFIG, dual_chunk_attention_config=DUAL_CHUNK_SPARSE_SUB_WINDOW_CONFIG,
) )
actual = run_dual_chunk_fixture_eager(fixture) actual, actual_selections = _run_dual_chunk_fixture_with_sparse_selection_capture(
expected = expected_dual_chunk_fixture_output(fixture) fixture
testcase.assertEqual(actual.shape, expected.shape)
testcase.assertTrue(
torch.isfinite(actual).all(),
f"sparse sub-window output has non-finite values: {actual}",
) )
# Bound the absolute magnitude to catch runaway softmax/scaling bugs. expected_selections = _dual_chunk_sparse_selection_reference(fixture)
max_abs = actual.abs().max().item() _assert_sparse_stage_selections_match(
testcase.assertLess( testcase, actual_selections, expected_selections
max_abs,
1e3,
f"sparse sub-window output magnitude {max_abs} suggests a numerical bug",
)
# The sparse path must differ from dense for at least one element by
# more than bf16 FP-noise (~1e-3 at typical accumulation depth). A
# silent fallback to dense produces diff ~0 on identical inputs, so the
# 5e-4 floor cleanly distinguishes "kernel pruned something" from
# "fallback to dense + FP noise".
abs_diff = (actual.float() - expected.float()).abs()
max_diff = abs_diff.max().item()
testcase.assertGreater(
max_diff,
5e-4,
"sparse sub-window output is too close to dense — the sparse path "
"may not have actually pruned. Check `sparse_attn_enabled` gate and "
"config (vertical_size + slash_size should be < intra K count, and "
"seq_len should exceed the production-hardcoded vertical[:30]=inf "
"and slash[-100:]=inf always-include heuristics).",
) )
expected = _run_dual_chunk_fixture_with_torch_sparse_kernel(fixture)
torch.testing.assert_close(actual, expected, atol=DENSE_ATOL, rtol=DENSE_RTOL)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -6,7 +6,7 @@ unit-test suite, **organized by the action needed to address it**.
Anything failing that is not listed here should be treated as a regression. Anything failing that is not listed here should be treated as a regression.
Last updated: 2026-05-29 Last updated: 2026-06-05
## Reference runs ## Reference runs
@@ -182,8 +182,6 @@ moment production is fixed; no test method invokes them today.
| Citation | Symptom | Trigger | Status | | Citation | Symptom | Trigger | Status |
|---|---|---|---| |---|---|---|---|
| `dual_chunk_flashattention_backend.py:1110-1132` | `RuntimeError: The size of tensor a (4) must match the size of tensor b (5)` at `vertical_buffer.copy_()` | `vertical_size ≤ 5`: fallback `torch.arange(0, intra_K_size, max(1, intra_K_size/5))` returns up to 5 elements into `vertical_size=4` buffer when `intra_vertical_indices.nelement() == 0` | `[no test]` (`dual_chunk/README.md`); smoke helper `run_dual_chunk_sparse_sub_window_case` wired but not invoked |
| `_vertical_slash_sparse_attention` (`convert_vertical_slash_indexes` block math) | `cudaErrorIllegalAddress` deep inside the kernel | `vertical_size=8` with `seq_len ≥ 128`: unstated invariant that `vertical_size + slash_size >= chunk_len_blocks` | `[no test]` (same smoke helper) |
| Triton dense `DRAFT_EXTEND` (non-V2) | Eager fixture/reference mismatch on narrow accepted-token layouts | Test omitted | `[no test]` (`dense/README.md`) | | Triton dense `DRAFT_EXTEND` (non-V2) | Eager fixture/reference mismatch on narrow accepted-token layouts | Test omitted | `[no test]` (`dense/README.md`) |
## C.6. DSA-specific structural gaps ## C.6. DSA-specific structural gaps
@@ -19,6 +19,7 @@ Columns are runner modes; rows are kernel-path modes of the single
|---|---|---|---|---|---|---|---|---|---|---|---|---| |---|---|---|---|---|---|---|---|---|---|---|---|---|
| Non-sparse | ✓ first-window, successor-chunk, inter-chunk extend/decode layouts + GQA decode | deferred: graph metadata for dual-chunk not scoped | deferred | deferred | blocked: `init_forward_metadata` asserts `is_prefill() or is_decode()` (`dual_chunk_flashattention_backend.py:179`); `TARGET_VERIFY` falls under `is_prefill()` but the wrapper hasn't been wired through | deferred | deferred | deferred | blocked: `DRAFT_EXTEND_V2` excluded from `is_prefill()` alias (see Production-Unsupported below) | deferred | deferred | — | | Non-sparse | ✓ first-window, successor-chunk, inter-chunk extend/decode layouts + GQA decode | deferred: graph metadata for dual-chunk not scoped | deferred | deferred | blocked: `init_forward_metadata` asserts `is_prefill() or is_decode()` (`dual_chunk_flashattention_backend.py:179`); `TARGET_VERIFY` falls under `is_prefill()` but the wrapper hasn't been wired through | deferred | deferred | deferred | blocked: `DRAFT_EXTEND_V2` excluded from `is_prefill()` alias (see Production-Unsupported below) | deferred | deferred | — |
| Sparse all-column (`vertical_size`/`slash_size` chosen so every key in the first chunk is selected) | ✓ single-request first-chunk, multi-request first-chunk, page-boundary first-chunk | — | — | — | blocked: same `is_prefill` assertion | — | — | — | blocked: same | — | — | — | | Sparse all-column (`vertical_size`/`slash_size` chosen so every key in the first chunk is selected) | ✓ single-request first-chunk, multi-request first-chunk, page-boundary first-chunk | — | — | — | blocked: same `is_prefill` assertion | — | — | — | blocked: same | — | — | — |
| Sparse sub-window (`vertical_size=4`, `slash_size=4`, `seq_len=128`) | ✓ independent DCA top-k/split/fallback reference + torch sparse-kernel reference | — | — | — | — | — | — | — | — | — | — | — |
| Threshold-gated sparse (`sparse_attention_threshold=100`, seq_len=16 → gate disables sparse, falls back to dense) | ✓ verifies `current_orig_seq_len > threshold` gate semantics | — | — | — | — | — | — | — | — | — | — | — | | Threshold-gated sparse (`sparse_attention_threshold=100`, seq_len=16 → gate disables sparse, falls back to dense) | ✓ verifies `current_orig_seq_len > threshold` gate semantics | — | — | — | — | — | — | — | — | — | — | — |
## Input And Config Coverage ## Input And Config Coverage
@@ -33,6 +34,11 @@ Columns are runner modes; rows are kernel-path modes of the single
(≤16 tokens) so the dense reference remains valid. (≤16 tokens) so the dense reference remains valid.
- Multi-request sparse and page-boundary sparse variants exercise per-request - Multi-request sparse and page-boundary sparse variants exercise per-request
`cu_seqlens_*` slicing inside `_dual_chunk_flash_attn_prefill_func`. `cu_seqlens_*` slicing inside `_dual_chunk_flash_attn_prefill_func`.
- Sub-window sparse prefill uses `vertical_size=4`, `slash_size=4`, and
`seq_len=128` to verify the DCA-specific content-aware top-k split and
empty-stage fallback against an independent reference, then verifies the
sparse output against a torch sparse-kernel reference that consumes the
production block/column metadata.
- Threshold-gated sparse uses `sparse_attention_threshold=100` so a 16-token - Threshold-gated sparse uses `sparse_attention_threshold=100` so a 16-token
prompt bypasses the sparse kernel and falls through to the dense chunk prompt bypasses the sparse kernel and falls through to the dense chunk
flash path, exercising the gate semantics in the wrapper. flash path, exercising the gate semantics in the wrapper.
@@ -77,65 +83,8 @@ See `KNOWN_FAILURES.md` §1 for the full root cause + fix.
- Populate CUDA graph and PCG/BCG runner metadata after eager non-sparse - Populate CUDA graph and PCG/BCG runner metadata after eager non-sparse
coverage is stable across more chunk layouts. coverage is stable across more chunk layouts.
- **Sub-context-window sparse pruning reference (genuine follow-up)** — - **Broaden sub-window sparse coverage** — the current regression case covers
The current "all-column" sparse cases match the dense reference exactly `prefix_lens=(0,)`, `extend_lens=(128,)`, and no GQA. Add
because the chosen `vertical_size=16` + `slash_size=16` + `last_q=16` multi-request batches, nonzero prefixes, GQA, and more sparse config variants
configuration covers every column in the first chunk for `seq_len <= 16`. once those paths need explicit sparse pruning coverage. The 64x64
A truly pruning case needs `seq_len >> vertical_size + slash_size` and a vertical/slash converter remains covered at the sgl-kernel layer.
reference that applies the same mask the kernel applies.
The blocker is that the production sparse-attention config
`("vertical_and_slash", v_size, s_size, threshold)` is **content-aware**:
per-head `v_idx` and `s_idx` are picked by top-k attention scores over
the last `last_q` queries, not from a fixed schedule
(`dual_chunk_flashattention_backend.py:_dual_chunk_flash_attn_prefill`).
An independent reference therefore has three paths:
1. **Mock the sparse-config lookup** — patch
`get_sparse_attention_config` or the per-layer top-k selection so the
fixture supplies known `v_idx` / `s_idx` tensors. Then write a
token-level reference that masks `attn_scores[q, k] = -inf` unless
`k in v_idx` or `(q - k) in s_idx` (with causal `k <= q`). This is the
cleanest path but needs a hook in `_dual_chunk_flash_attn_prefill_func`
that doesn't exist today.
2. **Replicate `convert_vertical_slash_indexes`** at block granularity in
pure-PyTorch, then iterate `(block_count, block_offset, column_count,
column_index)` to build a per-(query_block, key_block) mask matching
the kernel's selection. Faithful but tedious — the block math (M=64,
N=64) needs to be mirrored exactly.
3. **Statistical recovery check** — compute dense attention scores
`softmax(Q @ K^T)` per head, identify the top-k columns by score, and
verify the sparse kernel output approximates the dense output modulo
the dropped probability mass. Not strict `assert_close`; rejects only
gross divergences.
Option 1 is recommended. It requires either: (a) a new
`sparse_attention_config_override` kwarg threaded through
`DualChunkFlashAttentionBackend.__init__` that bypasses the content-aware
selection, or (b) monkeypatching `get_sparse_attention_config` on the
fixture's backend instance. Until that lands, the all-column sparse +
threshold-gated cases keep the kernel/wrapper integration covered but
the per-column sparse math is unverified.
**Production-side bugs surfaced while attempting Option 3
(smoke-test "sparse output != dense output"):** two issues block even a
smoke-only sub-window test today.
- `dual_chunk_flashattention_backend.py:1110-1122`: when a chunk's
`intra_vertical_indices.nelement() == 0`, the fallback appends
`torch.arange(0, intra_K_size, max(1, intra_K_size/5))`. With
`intra_K_size=48` this is `arange(0, 48, 9.6)` → 5 elements, but the
`vertical_buffer` is sized to `vertical_size` (=4 in a sub-window
config). The copy at line 1132 then raises
`RuntimeError: The size of tensor a (4) must match the size of
tensor b (5)`. The fallback should clip to `vertical_size` slots.
- With `vertical_size=8` to clear the overflow, the sparse kernel
crashes with `cudaErrorIllegalAddress` deep inside
`_vertical_slash_sparse_attention`, suggesting the
`convert_vertical_slash_indexes` block math has an unstated
invariant that `vertical_size + slash_size >= chunk_len_blocks` or
similar. Needs a kernel-side audit.
The smoke-test helper `run_dual_chunk_sparse_sub_window_case` is wired
through `common/attention_methods/dual_chunk_attention.py` for when
those production bugs are fixed; no test method invokes it today.
@@ -13,9 +13,11 @@ from sglang.test.kits.attention_unittest.attention_methods.dual_chunk_attention
DualChunkAttentionCase, DualChunkAttentionCase,
make_dual_chunk_cases, make_dual_chunk_cases,
make_dual_chunk_sparse_cases, make_dual_chunk_sparse_cases,
make_dual_chunk_sparse_sub_window_cases,
make_dual_chunk_sparse_threshold_gated_cases, make_dual_chunk_sparse_threshold_gated_cases,
run_dual_chunk_attention_case, run_dual_chunk_attention_case,
run_dual_chunk_sparse_attention_case, run_dual_chunk_sparse_attention_case,
run_dual_chunk_sparse_sub_window_case,
run_dual_chunk_sparse_threshold_gated_case, run_dual_chunk_sparse_threshold_gated_case,
) )
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import ( from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
@@ -75,6 +77,9 @@ class TestDualChunkFlashAttentionBackendCorrectness(CustomTestCase):
SPARSE_THRESHOLD_GATED_CASES = make_dual_chunk_sparse_threshold_gated_cases( SPARSE_THRESHOLD_GATED_CASES = make_dual_chunk_sparse_threshold_gated_cases(
"dual_chunk_flash_attn" "dual_chunk_flash_attn"
) )
SPARSE_SUB_WINDOW_CASES = make_dual_chunk_sparse_sub_window_cases(
"dual_chunk_flash_attn"
)
# Replay prefix_lens must each be >= capture_prefix_len (= fill-value - 1). # Replay prefix_lens must each be >= capture_prefix_len (= fill-value - 1).
# Dual-chunk's `get_cuda_graph_seq_len_fill_value()` returns 1, so capture # Dual-chunk's `get_cuda_graph_seq_len_fill_value()` returns 1, so capture
# uses prefix=0. We pick a 3-request batch with varied lengths to exercise # uses prefix=0. We pick a 3-request batch with varied lengths to exercise
@@ -106,34 +111,10 @@ class TestDualChunkFlashAttentionBackendCorrectness(CustomTestCase):
with self.subTest(case=case.name, backend=case.backend): with self.subTest(case=case.name, backend=case.backend):
run_dual_chunk_sparse_threshold_gated_case(self, case) run_dual_chunk_sparse_threshold_gated_case(self, case)
# Sub-context-window sparse pruning: BLOCKED on production-side def test_sparse_dual_chunk_sub_window_cases(self):
# edge cases. for case in self.SPARSE_SUB_WINDOW_CASES:
# with self.subTest(case=case.name, backend=case.backend):
# The `run_dual_chunk_sparse_sub_window_case` helper in run_dual_chunk_sparse_sub_window_case(self, case)
# `common/attention_methods/dual_chunk_attention.py` is left in
# place for when those production gaps are fixed, but no test
# method invokes it today. See `dual_chunk/README.md` →
# "Sub-context-window sparse pruning" for the engineering paths
# and the two production bugs surfaced while attempting to land
# this coverage:
#
# - `dual_chunk_flashattention_backend.py:1110-1122`: when a chunk's
# `intra_vertical_indices.nelement() == 0`, the fallback appends
# `torch.arange(0, intra_K_size, max(1, intra_K_size/5))` which
# can produce more elements than the `vertical_size`-slot buffer
# allows, raising `RuntimeError: The size of tensor a (4) must
# match the size of tensor b (5)`. Triggered by
# `vertical_size in [4, 5]` with `seq_len=128`.
# - With `vertical_size=8` to avoid the overflow above, the sparse
# kernel raises a `cudaErrorIllegalAddress` deep inside
# `_vertical_slash_sparse_attention`, suggesting the
# `convert_vertical_slash_indexes` block math expects different
# invariants than what a `vertical_size + slash_size < chunk_len`
# config supplies.
#
# The all-column + threshold-gated cases above keep the integration
# path covered; sub-window correctness needs production hardening
# before unit-test coverage is safe.
def test_runner_mode_cuda_graph_decode_cases(self): def test_runner_mode_cuda_graph_decode_cases(self):
for case in self.CUDA_GRAPH_DECODE_CASES: for case in self.CUDA_GRAPH_DECODE_CASES: