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
|