Download python/xrex_unified/mosaic_kernel.py from Snapkitty/ironic-mirror: direct link, hf CLI and curl.
- Browser
- Download file 9.74 kB
-
https://huggingface.co/Snapkitty/ironic-mirror/resolve/main/python/xrex_unified/mosaic_kernel.py
- Command line
-
hf download hf://Snapkitty/ironic-mirror/python/xrex_unified/mosaic_kernel.py
-
curl -L -o mosaic_kernel.py https://huggingface.co/Snapkitty/ironic-mirror/resolve/main/python/xrex_unified/mosaic_kernel.py
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 | |
| 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) | |
| 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) | |
| 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)) | |
| 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) | |
| 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): | |
| 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]) | |
| 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 | |