Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 18 additions & 12 deletions src/mcore_bridge/model/modules/kernels/qsa_block_sparse_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@ def _qsa_bs_fwd_kernel(
stride_ot,
stride_oh,
T,
S,
NB,
scale,
GROUP: tl.constexpr,
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -186,6 +187,7 @@ def _qsa_bs_dq_kernel(
stride_ot,
stride_oh,
T,
S,
NB,
scale,
GROUP: tl.constexpr,
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -284,6 +286,7 @@ def _qsa_bs_dkdv_kernel(
stride_ot,
stride_oh,
T,
S,
NB,
scale,
GROUP: tl.constexpr,
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
52 changes: 8 additions & 44 deletions src/mcore_bridge/model/modules/kernels/qsa_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down Expand Up @@ -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
Expand All @@ -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)
27 changes: 27 additions & 0 deletions tests/test_qsa_key_length.py
Original file line number Diff line number Diff line change
@@ -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