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