"""GQA fused-attention prefill kernel — short context (ADR-0060 §B.split.2). Prefill analogue of ``_gqa_decode_short.py`` — same CUBE/PE layout (``kv_per_cube`` heads per CUBE, group-PE-SP, within-group chain reduce, no inter-CUBE reduce). The only structural difference from short decode: ``T_q`` may be > 1 (prefill processes multiple query tokens) and Q is shaped ``(T_q, h_kv·d_head)`` — one Q head per KV head, no GQA M-fold. The local attention uses an S_kv-axis tile sweep (ADR-0063 §A.2) so per-rank scratch is bounded by ``TILE_S_KV``. No Ring KV here — each owned head is fully resident at its CUBE. """ from __future__ import annotations TILE_S_KV = 1024 # ADR-0063 §A.2 S_kv-axis tile sweep (per-tile width). def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl): """Online-softmax merge of two partial ``(m, ℓ, O)`` triples.""" m_new = tl.maximum(m_local, m_other) scale_old = tl.exp(m_local - m_new) scale_new = tl.exp(m_other - m_new) l_new = l_local * scale_old + l_other * scale_new O_new = O_local * scale_old + O_other * scale_new return m_new, l_new, O_new def gqa_attention_prefill_short_kernel( q_ptr: int, k_ptr: int, v_ptr: int, o_ptr: int, T_q: int, S_kv: int, h_kv: int, d_head: int, C: int, P: int, kv_per_cube: int, *, tl, ) -> None: """Short-context prefill with PE-parallel heads + intra-group PE-SP.""" group_size = P // kv_per_cube pe_id = tl.program_id(axis=0) pe_in_group = pe_id % group_size S_local = S_kv // group_size # ── Local attention (S_kv-axis tile sweep, ADR-0063 §A.2) ── Q = tl.load(q_ptr, shape=(h_kv * T_q, d_head), dtype="f16") n_tiles = (S_local + TILE_S_KV - 1) // TILE_S_KV KV_ROW_BYTES = d_head * 2 # f16 # Tile 0: establishes persistent (m_local, l_local, O_local). # # Cannot be folded into the Tiles 1..N loop (kernbench-only limitation): # - persistent (m, ℓ, O) must live OUTSIDE ``tl.scratch_scope``, # otherwise scope teardown discards them before the next tile's # merge can read them; # - kernbench has no scratch-backed initializer — ``tl.zeros`` / # ``tl.full`` return addr=0 handles with no backing storage, so # they cannot be overwritten via ``tl.copy_to`` to seed (-inf, 0, 0). # So Tile 0 computes the initial running state directly; Tiles 1..N # fold into it. Triton port: limitation does not apply (SSA tensors # stay live across iterations) — a single unified loop suffices. tile_s0 = min(TILE_S_KV, S_local) K_T = tl.load(k_ptr, shape=(d_head, tile_s0), dtype="f16") V = tl.load(v_ptr, shape=(tile_s0, d_head), dtype="f16") scores = tl.dot(Q, K_T) m_local = tl.max(scores, axis=-1) centered = scores - m_local exp_scores = tl.exp(centered) l_local = tl.sum(exp_scores, axis=-1) O_local = tl.dot(exp_scores, V) # Tiles 1..n_tiles-1: fold into running state via online-softmax merge. # Triton port: drop the ``with tl.scratch_scope():`` line and replace # each ``copy_to`` with a Python rebind. for tile_idx in range(1, n_tiles): tile_start = tile_idx * TILE_S_KV tile_s = min(TILE_S_KV, S_local - tile_start) with tl.scratch_scope(): K_T_t = tl.load(k_ptr + tile_start * KV_ROW_BYTES, shape=(d_head, tile_s), dtype="f16") V_t = tl.load(v_ptr + tile_start * KV_ROW_BYTES, shape=(tile_s, d_head), dtype="f16") scores_t = tl.dot(Q, K_T_t) m_tile = tl.max(scores_t, axis=-1) centered_t = scores_t - m_tile exp_scores_t = tl.exp(centered_t) l_tile = tl.sum(exp_scores_t, axis=-1) O_tile = tl.dot(exp_scores_t, V_t) m_new, l_new, O_new = _merge_running( m_local, l_local, O_local, m_tile, l_tile, O_tile, tl=tl, ) tl.copy_to(m_local, m_new) tl.copy_to(l_local, l_new) tl.copy_to(O_local, O_new) # ── Communication: within-group chain reduce-to-root (Level-2 only) ── group_cols = min(4, group_size) group_rows = (group_size + group_cols - 1) // group_cols pe_col_in_group = pe_in_group % group_cols pe_row_in_group = pe_in_group // group_cols # Row chain (within group's row, along intra_W, leftward). if group_cols > 1: if pe_col_in_group < group_cols - 1: with tl.scratch_scope(): m_other = tl.recv(dir="intra_E", shape=m_local.shape, dtype="f16") l_other = tl.recv(dir="intra_E", shape=l_local.shape, dtype="f16") O_other = tl.recv(dir="intra_E", shape=O_local.shape, dtype="f16") m_new, l_new, O_new = _merge_running( m_local, l_local, O_local, m_other, l_other, O_other, tl=tl, ) tl.copy_to(m_local, m_new) tl.copy_to(l_local, l_new) tl.copy_to(O_local, O_new) if pe_col_in_group > 0: tl.send(dir="intra_W", src=m_local) tl.send(dir="intra_W", src=l_local) tl.send(dir="intra_W", src=O_local) # Col bridge (within group, along intra_N, row-1 → row-0). if pe_col_in_group == 0 and group_rows > 1: if pe_row_in_group < group_rows - 1: with tl.scratch_scope(): m_other = tl.recv(dir="intra_S", shape=m_local.shape, dtype="f16") l_other = tl.recv(dir="intra_S", shape=l_local.shape, dtype="f16") O_other = tl.recv(dir="intra_S", shape=O_local.shape, dtype="f16") m_new, l_new, O_new = _merge_running( m_local, l_local, O_local, m_other, l_other, O_other, tl=tl, ) tl.copy_to(m_local, m_new) tl.copy_to(l_local, l_new) tl.copy_to(O_local, O_new) if pe_row_in_group > 0: tl.send(dir="intra_N", src=m_local) tl.send(dir="intra_N", src=l_local) tl.send(dir="intra_N", src=O_local) # ── Final normalise + store (group root only) ── if pe_in_group == 0: O_final = O_local / l_local tl.store(o_ptr, O_final)