ironic-mirror / python /xrex_unified /mosaic_kernel.py
SNAPKITTYWEST's picture
push from SNAPKITTYWEST/ironic-mirror
677e207 verified
Raw
History Blame Contribute Delete
9.74 kB
#
# 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