79ddb12b42
Comment-only change. Adds a kernbench-only limitation note at the
Tile-0 / bootstrap block in all four GQA attention kernels:
- persistent (m, l, O) must live outside tl.scratch_scope so
subsequent tiles' merges can read them;
- kernbench has no scratch-backed initializer (tl.zeros / tl.full
return addr=0 handles), so we can't seed (-inf, 0, 0) and rely
on tl.copy_to;
- therefore Tile 0 must compute the initial running state directly.
prefill_long adds a third bullet noting the ring-step ordering reason
(k=0 must send W before any k>0 can recv E, so the (t=0, k=0) step has
to run outside the scoped loop where the send is conditional on C > 1).
Each block also notes that a real Triton port collapses Tile 0 into a
unified loop (SSA tensors stay live across iterations).
No behavior change; tests unchanged.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
148 lines
6.2 KiB
Python
148 lines
6.2 KiB
Python
"""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)
|