From 223725e52afb145e6a5256a4dab6f43af882601f Mon Sep 17 00:00:00 2001 From: shiaho <222622008+shiaho777@users.noreply.github.com> Date: Fri, 9 Oct 2026 15:45:18 +0800 Subject: [PATCH] fix: bound QSA keys by the key length under context parallelism The sparse kernel used the query count as the key bound (offs_k < T) and built the block bitmap with ceil(T / block_size). Context parallelism holds a 1/cp_size query shard and the full key sequence. The CP path therefore copied the local queries into a buffer as long as the keys and filled the rest with -1. Those rows launch tiles and allocate a bitmap for the whole sequence, then the kernel skips them. The key bound and the bitmap width now follow k.shape[0]. The query grid still follows the query count. The CP path passes the local queries and the gathered keys directly. When the two lengths are equal, the bounds match the previous kernel. Checked with python3 -m pytest tests/test_qsa_key_length.py. Two queries and eight keys, block size 4, flag key index 6 in block 1. flake8 is clean. The Triton kernels were not executed here; triton is not installed. --- .../modules/kernels/qsa_block_sparse_attn.py | 30 ++++++----- .../model/modules/kernels/qsa_kernels.py | 52 +++---------------- tests/test_qsa_key_length.py | 27 ++++++++++ 3 files changed, 53 insertions(+), 56 deletions(-) create mode 100644 tests/test_qsa_key_length.py diff --git a/src/mcore_bridge/model/modules/kernels/qsa_block_sparse_attn.py b/src/mcore_bridge/model/modules/kernels/qsa_block_sparse_attn.py index 6231c0b..6689af4 100644 --- a/src/mcore_bridge/model/modules/kernels/qsa_block_sparse_attn.py +++ b/src/mcore_bridge/model/modules/kernels/qsa_block_sparse_attn.py @@ -70,6 +70,7 @@ def _qsa_bs_fwd_kernel( stride_ot, stride_oh, T, + S, NB, scale, GROUP: tl.constexpr, @@ -108,7 +109,7 @@ def _qsa_bs_fwd_kernel( for i in range(0, n_tiles): kt = tl.load(KLIST + pid_t * stride_kl + i).to(tl.int64) offs_k = kt * BK + tl.arange(0, BK) - k_in = offs_k < T + k_in = offs_k < S # per-sequence block grid: after the packed-indexer fix a sequence's blocks start # at its own first token, which is not a multiple of BLK in a packed batch. @@ -186,6 +187,7 @@ def _qsa_bs_dq_kernel( stride_ot, stride_oh, T, + S, NB, scale, GROUP: tl.constexpr, @@ -220,7 +222,7 @@ def _qsa_bs_dq_kernel( for i in range(0, n_tiles): kt = tl.load(KLIST + pid_t * stride_kl + i).to(tl.int64) offs_k = kt * BK + tl.arange(0, BK) - k_in = offs_k < T + k_in = offs_k < S # per-sequence block grid: after the packed-indexer fix a sequence's blocks start # at its own first token, which is not a multiple of BLK in a packed batch. @@ -284,6 +286,7 @@ def _qsa_bs_dkdv_kernel( stride_ot, stride_oh, T, + S, NB, scale, GROUP: tl.constexpr, @@ -304,7 +307,7 @@ def _qsa_bs_dkdv_kernel( offs_k = pid_k * BK + tl.arange(0, BK) offs_d = tl.arange(0, D) - k_in = offs_k < T + k_in = offs_k < S k_tile = tl.load( K + offs_k[:, None] * stride_kt + kv_head * stride_kh + offs_d[None, :], mask=k_in[:, None], other=0.0) @@ -366,14 +369,14 @@ def _qsa_bs_dkdv_kernel( tl.store(DV + offs_k[:, None] * stride_vt + kv_head * stride_vh + offs_d[None, :], dv, mask=k_in[:, None]) -def selection_to_block_bitmap(indices: Tensor, num_tokens: int, block_size: int) -> Tensor: - """``[T, K]`` token indices (``-1`` pad) -> ``[T, ceil(T / block_size)]`` uint8 flags. +def selection_to_block_bitmap(indices: Tensor, num_keys: int, block_size: int) -> Tensor: + """``[T, K]`` token indices (``-1`` pad) -> ``[T, ceil(num_keys / block_size)]`` uint8 flags. - A block is flagged when any of its tokens appears in the row. Tokens that the caller - clamped away inside an otherwise selected block are re-excluded by the ``lo``/``hi`` - range test in the kernel, so this stays exact while being ``block_size``x smaller. + ``num_keys`` is the key length. Under context parallelism the query shard is shorter + than the gathered keys, and sizing the bitmap from the query count drops every key + past that count. """ - num_blocks = -(-num_tokens // block_size) + num_blocks = -(-num_keys // block_size) flags = torch.zeros(indices.shape[0], num_blocks, dtype=torch.uint8, device=indices.device) valid = indices >= 0 rows = torch.arange(indices.shape[0], device=indices.device).unsqueeze(1).expand_as(indices) @@ -472,6 +475,7 @@ def forward(ctx, q, k, v, sel, lo, hi, blk_base, tok_base, scale, block_size): o.stride(0), o.stride(1), T, + S, selc.shape[1], scale, GROUP=group, @@ -531,6 +535,7 @@ def backward(ctx, grad_out): do.stride(0), do.stride(1), T, + kc.shape[0], selc.shape[1], ctx.scale, GROUP=ctx.group, @@ -541,7 +546,7 @@ def backward(ctx, grad_out): num_warps=8, num_stages=1, ) - _qsa_bs_dkdv_kernel[(triton.cdiv(T, BK), kc.shape[1])]( + _qsa_bs_dkdv_kernel[(triton.cdiv(kc.shape[0], BK), kc.shape[1])]( qc, kc, vc, @@ -562,6 +567,7 @@ def backward(ctx, grad_out): do.stride(0), do.stride(1), T, + kc.shape[0], selc.shape[1], ctx.scale, GROUP=ctx.group, @@ -602,8 +608,8 @@ def qsa_sparse_attention_from_indices(q: Tensor, scale: float, block_size: int = 4) -> Tensor: """Drop-in for the gather kernel: derives the bitmap and range from ``indices``.""" - T = q.shape[0] - sel = selection_to_block_bitmap(indices, T, block_size) + key_len = k.shape[0] + sel = selection_to_block_bitmap(indices, key_len, block_size) valid = indices >= 0 big = torch.iinfo(torch.int32).max lo = torch.where(valid, indices, torch.full_like(indices, big)).min(dim=1).values.to(torch.int32) diff --git a/src/mcore_bridge/model/modules/kernels/qsa_kernels.py b/src/mcore_bridge/model/modules/kernels/qsa_kernels.py index fa57d75..0ef0711 100644 --- a/src/mcore_bridge/model/modules/kernels/qsa_kernels.py +++ b/src/mcore_bridge/model/modules/kernels/qsa_kernels.py @@ -126,12 +126,8 @@ def qsa_sparse_attention_thd(q, k, v, indices, scale, block_size): 'selected for packing (thd) or CP>1, where no dense fallback is correct.') if q.shape[-1] & (q.shape[-1] - 1): raise RuntimeError(f'QSA sparse attention needs a power-of-two head dim, got {q.shape[-1]}.') - if q.shape[0] != k.shape[0]: - # The kernel takes its key bound from the query count (T, Hq, D = q.shape, - # then `offs_k < T`), so unequal lengths would silently drop every key past - # len(q). Callers must equalise first -- _forward_cp does this by scattering - # the local query shard into a full-length buffer. - raise ValueError(f'QSA sparse attention needs len(q) == len(k), got {q.shape[0]} vs {k.shape[0]}.') + if q.shape[0] > k.shape[0]: + raise ValueError(f'QSA sparse attention got more queries than keys, {q.shape[0]} vs {k.shape[0]}.') return qsa_sparse_attention_from_indices(q, k, v, indices.contiguous(), scale, block_size) @@ -254,7 +250,7 @@ def _forward_cp(self, query, key, value, indices, scale, packed_seq_params): gathered_pos = torch.cat([_cp_query_global_positions_thd(cu_q, cp_size, r, device) for r in range(cp_size)]) kv_reorder = torch.argsort(gathered_pos) else: - sq, b = query.shape[0], query.shape[1] + sq = query.shape[0] q_pos = _cp_query_global_positions(sq * cp_size, cp_size, cp_rank, device) kv_reorder = _cp_gathered_to_logical_order(sq * cp_size, cp_size, device) # gather k/v across CP with a DIFFERENTIABLE all-gather (backward is a @@ -270,41 +266,9 @@ def _gather_full(t): key_full = _gather_full(key) value_full = _gather_full(value) - # The kernel derives the key bound from the query count (T, Hq, D = q.shape, - # then `offs_k < T`), so it structurally requires len(q) == len(k). Under CP - # the queries are a 1/cp_size shard while k/v are now full length, so scatter - # the local queries back into a full-length buffer, run, and take our rows - # out again. The padding rows carry an all-`-1` selection, which the kernel - # skips, so they cost tile launches but produce nothing. + # Query stays the local CP shard. Key indices are in the gathered sequence. + # The kernel bounds keys by k.shape[0], so the empty rows of a full-length + # query buffer are not required. if thd: - local_idx = indices[q_pos] - out_full = qsa_sparse_attention(*self._scatter_q_to_full(query, key_full, value_full, local_idx, q_pos), - scale, self.block_size) - return out_full.index_select(0, q_pos) - # sbhd: token-space kernel on the batch-major flattening (t = r*sk + p) - local_idx = indices[:, q_pos] - sk = key_full.shape[0] - k_f = key_full.permute(1, 0, 2, 3).reshape(b * sk, key_full.shape[2], key_full.shape[3]) - v_f = value_full.permute(1, 0, 2, 3).reshape(b * sk, value_full.shape[2], value_full.shape[3]) - off = torch.arange(b, device=device).view(b, 1, 1) * sk - idx_f = torch.where(local_idx >= 0, local_idx + off, local_idx.new_full((), -1)).reshape(sq * b, -1) - q_f = query.permute(1, 0, 2, 3).reshape(sq * b, query.shape[2], query.shape[3]) - # batch-major token ids of this rank's rows: sample r contributes q_pos + r*sk - rows = (q_pos[None, :] + torch.arange(b, device=device).view(b, 1) * sk).reshape(-1) - out_f = qsa_sparse_attention(*self._scatter_q_to_full(q_f, k_f, v_f, idx_f, rows), scale, self.block_size) - out_f = out_f.index_select(0, rows) - return out_f.view(b, sq, query.shape[2], query.shape[3]).permute(1, 0, 2, 3) - - @staticmethod - def _scatter_q_to_full(q, k, v, indices, rows): - """Place ``q``/``indices`` rows at ``rows`` inside a len(k)-row buffer. - - index_copy keeps this differentiable: backward gathers the same rows, so the - padded positions contribute no gradient. - """ - n = k.shape[0] - q_full = q.new_zeros((n, *q.shape[1:])) - q_full = q_full.index_copy(0, rows, q) - idx_full = indices.new_full((n, indices.shape[1]), -1) - idx_full = idx_full.index_copy(0, rows, indices) - return q_full, k, v, idx_full + return qsa_sparse_attention(query, key_full, value_full, indices[q_pos], scale, self.block_size) + return qsa_sparse_attention(query, key_full, value_full, indices[:, q_pos], scale, self.block_size) diff --git a/tests/test_qsa_key_length.py b/tests/test_qsa_key_length.py new file mode 100644 index 0000000..990ccc9 --- /dev/null +++ b/tests/test_qsa_key_length.py @@ -0,0 +1,27 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The QSA block bitmap is sized from the key length, not the query count.""" +import ast +import torch +from pathlib import Path + + +def _bitmap(): + path = Path(__file__).resolve().parents[1] / 'src/mcore_bridge/model/modules/kernels/qsa_block_sparse_attn.py' + tree = ast.parse(path.read_text()) + fn = next( + node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == 'selection_to_block_bitmap') + module = ast.Module(body=[fn], type_ignores=[]) + ast.fix_missing_locations(module) + namespace = {'Tensor': torch.Tensor, 'torch': torch} + exec(compile(module, str(path), 'exec'), namespace) + return namespace['selection_to_block_bitmap'] + + +def test_short_query_can_flag_a_later_key_block(): + bitmap = _bitmap() + # 2 queries, keys of length 8, block 4. Index 6 is in block 1, past the query count. + indices = torch.tensor([[6, -1], [-1, -1]]) + flags = bitmap(indices, 8, 4) + assert flags.shape == (2, 2) + assert int(flags[0, 1]) == 1 + assert int(flags[0, 0]) == 0