Kernels:
Trusted publisher
File size: 7,282 Bytes
e19323e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 | # Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
import torch
import triton
import triton.language as tl
from ...ops.utils.index import prepare_chunk_offsets
@triton.heuristics({
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
'USE_BLOCK_COUNTS': lambda args: isinstance(args['block_counts'], torch.Tensor),
})
@triton.jit(do_not_specialize=['T', 'N'])
def prepare_block_csr_kernel(
block_indices,
block_counts,
cu_seqlens,
chunk_offsets,
cursor,
csr_indices,
csr_offsets,
N,
T,
H: tl.constexpr,
S: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
TC: tl.constexpr,
COUNT_ONLY: tl.constexpr,
USE_BLOCK_COUNTS: tl.constexpr,
IS_VARLEN: tl.constexpr,
):
i_t, i_bh = tl.program_id(0), tl.program_id(1)
i_b, i_h = i_bh // H, i_bh % H
o_t = i_t * BT + tl.arange(0, BT)
o_s = tl.arange(0, S)
m_t = o_t < T
# [BT] flattened (batch, query, kv-head) index; int64 to keep address arithmetic safe at large T
i_qh = ((i_b * T).to(tl.int64) + o_t) * H + i_h
# [BT, S] selected blocks, masked to each query's valid causal in-range slots
b_i = tl.load(block_indices + i_qh[:, None] * S + o_s[None, :], mask=m_t[:, None], other=-1).to(tl.int64)
if USE_BLOCK_COUNTS:
b_m = m_t[:, None] & (o_s[None, :] < tl.load(block_counts + i_qh, mask=m_t, other=0)[:, None])
else:
b_m = m_t[:, None] & (o_s[None, :] < block_counts)
b_m = b_m & (b_i >= 0) & (b_i < TC) & (b_i * BS <= o_t[:, None])
if IS_VARLEN:
# vectorized binary search for the sequence holding each query (32 steps cover any num_seq)
lo, hi = tl.zeros([BT], dtype=tl.int32), tl.full([BT], N, dtype=tl.int32)
for _ in range(32):
mid = (lo + hi) // 2
go = tl.load(cu_seqlens + mid + 1, mask=m_t, other=0) <= o_t
lo, hi = tl.where(go, mid + 1, lo), tl.where(go, hi, mid)
block_base = tl.load(chunk_offsets + lo, mask=m_t, other=0).to(tl.int64)
block_id = (block_base[:, None] + b_i) * H + i_h
else:
block_id = (i_b * H + i_h) * TC + b_i
if COUNT_ONLY:
tl.atomic_add(csr_offsets + block_id + 1, 1, mask=b_m)
else:
dst = tl.load(csr_offsets + block_id, mask=b_m, other=0).to(tl.int64) + tl.atomic_add(cursor + block_id, 1, mask=b_m)
b_q = tl.broadcast_to((i_b * T + o_t)[:, None], (BT, S))
tl.store(csr_indices + dst, b_q.to(csr_indices.dtype.element_ty), mask=b_m)
def prepare_block_csr(
block_indices: torch.LongTensor,
block_counts: torch.LongTensor | int,
cu_seqlens: torch.LongTensor | None,
chunk_indices: torch.LongTensor | None,
num_blocks: int,
block_size: int,
) -> tuple[torch.Tensor, torch.Tensor]:
r"""
Invert a per-query block selection into CSR (compressed sparse row) form.
`block_indices[b, t, h, :]` lists the blocks query `t` (kv-head `h`) selects.
The inverse maps each block to the queries that selected it,
which a block-parallel backward (e.g. NSA `bwd_dkv`) needs.
The result is CSR over a `[block, query]` matrix: `csr_indices` holds the selecting query positions grouped by block,
and `csr_offsets` holds the per-block row offsets,
so block `i` owns `csr_indices[csr_offsets[i]:csr_offsets[i + 1]]`.
Block ids follow the launching kernel's program-id layout:
`(b * H + h) * num_blocks + s` for dense, `(global_block + s) * H + h` for varlen,
where the varlen block base comes from an in-kernel binary search over `cu_seqlens` into `chunk_offsets`.
Short problems use a counting-then-scatter sort (`csr_indices` over-allocated to its upper bound, so its length
need not be read back to the host); long ones bucket the pairs with a radix sort. Both yield the same CSR.
Example (dense, B = H = 1, block_size = 1 so block b covers token b; causal needs b <= t):
# input: each query lists the blocks it selects, -1 is padding
block_indices[0, :, 0, :] =
[[ 0, -1], # query 0 selects block 0
[ 0, 1], # query 1 selects blocks 0, 1
[ 1, 2], # query 2 selects blocks 1, 2
[ 0, 3]] # query 3 selects blocks 0, 3
# invert -> which queries selected each block:
# block 0: queries 0, 1, 3
# block 1: queries 1, 2
# block 2: query 2
# block 3: query 3
csr_indices = [0, 1, 3, 1, 2, 2, 3] # 7 selections, grouped by block (order within a block is arbitrary)
csr_offsets = [0, 3, 5, 6, 7] # block i's queries = csr_indices[csr_offsets[i]:csr_offsets[i+1]]
Args:
block_indices (torch.LongTensor):
Selected block ids of shape `[B, T, H, S]`, padded with `-1`.
block_counts (torch.LongTensor or int):
Number of valid slots per query, a `[B, T, H]` tensor or an int.
cu_seqlens (torch.LongTensor, Optional):
Cumulative sequence lengths for variable-length packing. Default: `None` (dense).
chunk_indices (torch.LongTensor):
Per-chunk `(sequence, local-block)` index pairs; read only to size the varlen block-id space.
num_blocks (int):
Number of blocks per `(batch, head)`, i.e. the dense kernel's `TC`.
block_size (int):
Selected block size, used for the causal check and varlen block ids.
Returns:
csr_indices (torch.Tensor):
`int32` selecting query positions, grouped by block; absolute (`b * T + t`).
csr_offsets (torch.Tensor):
`int32` CSR row offsets of shape `[NB + 1]`, one per block plus a final end offset.
"""
B, T, H, S = block_indices.shape
N = 0 if cu_seqlens is None else cu_seqlens.numel() - 1
NB = B * H * num_blocks if cu_seqlens is None else chunk_indices.shape[0] * H
chunk_offsets = prepare_chunk_offsets(cu_seqlens, block_size) if cu_seqlens is not None else None
cursor = block_indices.new_zeros(NB, dtype=torch.int32)
csr_offsets = block_indices.new_zeros(NB + 1, dtype=torch.int32)
csr_indices = block_indices.new_empty(B * T * H * S, dtype=torch.int32)
BT = max(1, min(128, triton.next_power_of_2(max(1, 2048 // S))))
grid = (triton.cdiv(T, BT), B * H)
# counting sort: tally per-block counts, prefix-sum them into start offsets, then scatter.
# the two kernel passes can't merge -- the scatter position needs csr_offsets ready (global prefix sum).
kwargs = dict(
block_indices=block_indices,
block_counts=block_counts,
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
cursor=cursor,
csr_indices=csr_indices,
csr_offsets=csr_offsets,
N=N,
T=T,
H=H,
S=S,
BT=BT,
BS=block_size,
TC=num_blocks,
)
prepare_block_csr_kernel[grid](**kwargs, COUNT_ONLY=True)
csr_offsets.cumsum_(0)
prepare_block_csr_kernel[grid](**kwargs, COUNT_ONLY=False)
return csr_indices, csr_offsets
|