File size: 9,738 Bytes
677e207
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
#
# Copyright (c) 2026 BEL ESPRIT D ACCORD TRUST HOLDINGS INC
# All rights reserved.

# SPDX-License-Identifier: Apache-2.0
# Copyright 2026 X.AI Corp.
"""
Mosaic GPU warp-specialized forward kernel for ranker attention.

Uses 3 warp groups:
  - WG0, WG1: Compute (WGMMA, online softmax, cap)
  - WG2: Memory (TMA prefetch pipeline)

Register budget: 232 (compute) / 40 (memory)
Pipeline depth: min(num_stages, 4) for TMA overlap.
"""

import math

import jax
import jax.numpy as jnp
from jax import lax
from jax.experimental import pallas as pl
from jax.experimental.pallas import mosaic_gpu as plgpu

from .cap_functions import cap_forward, CapMethod, CapParams
from .segment_bounds import SegmentBounds
from .kernel_config import KernelConfig


def ranker_mask_mosaic(q_seq_base, block_q, kv_seq_base, block_kv, bounds):
    q_ids = plgpu.broadcasted_iota(jnp.int32, (block_q, block_kv), 0,
                                    layout=plgpu.Layout.WGMMA) + q_seq_base
    kv_ids = plgpu.broadcasted_iota(jnp.int32, (block_q, block_kv), 1,
                                     layout=plgpu.Layout.WGMMA) + kv_seq_base

    q_hist = (q_ids >= bounds.history_lower) & (q_ids < bounds.history_upper)
    q_cand = (q_ids >= bounds.candidate_lower) & (q_ids < bounds.candidate_upper)
    kv_hist = (kv_ids >= bounds.history_lower) & (kv_ids < bounds.history_upper)
    kv_cand = (kv_ids >= bounds.candidate_lower) & (kv_ids < bounds.candidate_upper)

    hist_mask = kv_hist & (q_hist | q_cand)
    cand_self = q_cand & kv_cand & (q_ids == kv_ids)
    return hist_mask | cand_self


def make_mosaic_forward_kernel(config: KernelConfig, q_heads_per_kv_head: int, head_dim: int):
    block_q = config.block_q
    block_kv = config.block_kv
    max_concurrent = min(config.num_stages, 4)

    def kernel(q_ref, k_ref, v_ref, bound_ref, out_ref, lse_ref, scoped):
        smem_buffers, buffer_barriers, consumed_barriers, schedule_barrier = scoped
        wg_idx = lax.axis_index("wg")
        batch = lax.axis_index("batch")
        q_head = lax.axis_index("heads")
        q_seq = lax.axis_index("q_seq")

        qo_smem2, k_smem, v_smem, lse_smem2 = smem_buffers
        k_barriers, v_barriers, q_barriers = buffer_barriers
        k_consumed, v_consumed = consumed_barriers

        hl = plgpu.load(bound_ref, (batch, 0))
        hu = plgpu.load(bound_ref, (batch, 1))
        cl = plgpu.load(bound_ref, (batch, 2))
        cu = plgpu.load(bound_ref, (batch, 3))
        bounds = SegmentBounds(hl, hu, cl, cu)

        q_tile_base = q_seq * (2 * block_q)
        q_tile_end = q_tile_base + (2 * block_q)

        def tile_has_tokens(lo, hi):
            return (q_tile_base < hi) & (q_tile_end > lo)

        valid = tile_has_tokens(hl, hu) | tile_has_tokens(cl, cu)

        hist_k_start = lax.div(hl, block_kv)
        hist_k_end = pl.cdiv(hu, block_kv)
        hist_steps = jnp.maximum(hist_k_end - hist_k_start, 0)

        cand_start = jnp.maximum(cl, q_tile_base)
        cand_end = jnp.minimum(cu, q_tile_end)
        cand_has = cand_start < cand_end
        cand_k_start = lax.div(cand_start, block_kv)
        cand_k_end = pl.cdiv(cand_end, block_kv)
        cand_steps = jnp.where(cand_has, cand_k_end - cand_k_start, 0)

        total_steps = hist_steps + cand_steps

        @pl.when((wg_idx < 2) & (~valid))
        def _zero():
            qo_smem = qo_smem2.at[wg_idx]
            zero = plgpu.layout_cast(
                jnp.zeros((block_q, head_dim), jnp.float32), plgpu.Layout.WGMMA)
            qo_smem[...] = zero.astype(q_ref.dtype)
            plgpu.commit_smem()
            q_seq_base = q_seq * (2 * block_q) + wg_idx * block_q
            plgpu.copy_smem_to_gmem(qo_smem, out_ref.at[batch, pl.ds(q_seq_base, block_q), q_head])
            plgpu.wait_smem_to_gmem(0)

        @pl.when((wg_idx < 2) & valid)
        def _compute():
            plgpu.set_max_registers(232, action="increase")
            qo_smem = qo_smem2.at[wg_idx]
            lse_smem = lse_smem2.at[wg_idx] if lse_smem2 is not None else None
            q_seq_base = q_seq * (2 * block_q) + wg_idx * block_q
            kv_head = lax.div(q_head, jnp.array(q_heads_per_kv_head, q_head.dtype))

            plgpu.copy_gmem_to_smem(
                q_ref.at[batch, pl.ds(q_seq_base, block_q), q_head],
                qo_smem, q_barriers.at[wg_idx]
            )
            plgpu.barrier_wait(q_barriers.at[wg_idx])

            m_i = plgpu.layout_cast(
                jnp.full((block_q,), -jnp.inf, jnp.float32), plgpu.Layout.WGMMA_ROW)
            l_i = plgpu.layout_cast(
                jnp.zeros((block_q,), jnp.float32), plgpu.Layout.WGMMA_ROW)
            acc = plgpu.layout_cast(
                jnp.zeros((block_q, head_dim), jnp.float32), plgpu.Layout.WGMMA)

            @pl.when(total_steps > 0)
            def _wait_first():
                plgpu.barrier_wait(k_barriers.at[0])

            def kv_loop(kv_step, carry):
                acc, m_i, l_i = carry
                slot = lax.rem(kv_step, jnp.array(max_concurrent, kv_step.dtype))

                kv_block_idx = jnp.where(
                    kv_step < hist_steps,
                    hist_k_start + kv_step,
                    cand_k_start + (kv_step - hist_steps)
                )

                def compute_qk(acc_ref):
                    plgpu.wgmma(acc_ref, qo_smem,
                                plgpu.transpose_ref(k_smem.at[slot], (1, 0)))
                    return acc_ref[...]
                qk = pl.run_scoped(compute_qk,
                                   plgpu.ACC((block_q, block_kv), jnp.float32))
                plgpu.barrier_arrive(k_consumed.at[slot])

                if config.sm_scale != 1.0:
                    qk *= config.sm_scale
                qk_capped = cap_forward(qk, config.cap_method, config.cap_params)

                kv_seq_base = kv_block_idx * block_kv
                mask = ranker_mask_mosaic(q_seq_base, block_q, kv_seq_base, block_kv, bounds)
                qk_capped = jnp.where(mask, qk_capped, -jnp.inf)

                log2e = math.log2(math.e)
                m_ij = jnp.maximum(m_i, qk_capped.max(axis=1) * log2e)
                alpha = jnp.exp2(m_i - m_ij)
                m_i = m_ij
                p = jnp.exp2(qk_capped * log2e -
                             lax.broadcast_in_dim(m_ij, qk_capped.shape, [0]))
                acc *= lax.broadcast_in_dim(alpha, acc.shape, [0])
                l_i *= alpha
                p16 = p.astype(q_ref.dtype)

                plgpu.barrier_arrive(schedule_barrier)
                plgpu.barrier_wait(v_barriers.at[slot])
                plgpu.barrier_wait(schedule_barrier)

                l_i += p.sum(axis=1)

                def compute_pv(acc_ref):
                    plgpu.wgmma(acc_ref, p16, v_smem.at[slot])
                    wait_step = kv_step + 1
                    wait_slot = lax.rem(wait_step, jnp.array(max_concurrent, kv_step.dtype))
                    @pl.when(wait_step < total_steps)
                    def _wait_next():
                        plgpu.barrier_wait(k_barriers.at[wait_slot])
                acc = pl.run_state(compute_pv)(plgpu.ACC.init(acc))
                plgpu.barrier_arrive(v_consumed.at[slot])

                return acc, m_i, l_i

            acc, m_i, l_i = lax.fori_loop(0, total_steps, kv_loop, (acc, m_i, l_i))

            acc /= lax.broadcast_in_dim(l_i, (block_q, head_dim), [0])
            qo_smem[...] = acc.astype(q_ref.dtype)
            if lse_smem is not None:
                RCP_LN2 = 1.4426950408889634
                lse_smem[...] = m_i + jnp.log2(l_i) * RCP_LN2

            plgpu.commit_smem()
            plgpu.copy_smem_to_gmem(qo_smem,
                                     out_ref.at[batch, pl.ds(q_seq_base, block_q), q_head])
            if lse_smem is not None:
                plgpu.copy_smem_to_gmem(lse_smem,
                                         lse_ref.at[batch, q_head, pl.ds(q_seq_base, block_q)])
            plgpu.wait_smem_to_gmem(0)

        @pl.when((wg_idx == 2) & valid)
        def _memory():
            plgpu.set_max_registers(40, action="decrease")
            kv_head = lax.div(q_head, jnp.array(q_heads_per_kv_head, q_head.dtype))

            for i in range(max_concurrent):
                @pl.when(i < total_steps)
                def _prefetch(i=i):
                    kv_block_idx = jnp.where(
                        i < hist_steps,
                        hist_k_start + i,
                        cand_k_start + (i - hist_steps)
                    )
                    s = (batch, pl.ds(kv_block_idx * block_kv, block_kv), kv_head)
                    plgpu.copy_gmem_to_smem(k_ref.at[s], k_smem.at[i], k_barriers.at[i])
                    plgpu.copy_gmem_to_smem(v_ref.at[s], v_smem.at[i], v_barriers.at[i])

            @pl.loop(0, jnp.maximum(total_steps - max_concurrent, 0))
            def _pipe(kv_step):
                tma_step = kv_step + max_concurrent
                tma_slot = lax.rem(kv_step, jnp.array(max_concurrent, kv_step.dtype))
                kv_block_idx = jnp.where(
                    tma_step < hist_steps,
                    hist_k_start + tma_step,
                    cand_k_start + (tma_step - hist_steps)
                )
                s = (batch, pl.ds(kv_block_idx * block_kv, block_kv), kv_head)
                plgpu.barrier_wait(k_consumed.at[tma_slot])
                plgpu.copy_gmem_to_smem(k_ref.at[s], k_smem.at[tma_slot],
                                         k_barriers.at[tma_slot])
                plgpu.barrier_wait(v_consumed.at[tma_slot])
                plgpu.copy_gmem_to_smem(v_ref.at[s], v_smem.at[tma_slot],
                                         v_barriers.at[tma_slot])

    return kernel