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