# # 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