Files
kernbench2/tests/attention/test_gqa_prefill.py
T
mukesh 7fad0371c5 gqa: tile-granular Ring KV (P3c) + rename to gqa_attention_* + ADR-0060/62/63/64 → Accepted
Three logically distinct changes, bundled for atomic test green:

1. **P3c — prefill_long tile-granular Ring KV** (ADR-0060 §5.5.1 amendment).
   Convert the ring from slice-granular (one full ``(d_head, S_local)``
   KV slice per step) to tile-granular (``n_tiles`` tiles of
   ``TILE_S_KV`` per step). Nested loop with outer tile, inner ring step:
   each tile propagates through all C ring positions before the next
   tile starts, so IPCQ in-flight depth stays at 1 per direction.
   Bootstrap at ``(t=0, k=0)`` outside the scratch_scope establishes the
   persistent ``(m, ℓ, O)``; every other iteration scope-wraps + persists
   via ``copy_to``. Per-rank persistent scratch shrinks to ~1 KB; per-tile
   scope bounded by TILE_S_KV regardless of S_local. Headline:
   prefill_long now completes at S_kv=128K (previously overflowed).
   New: ``tests/attention/test_gqa_prefill_long_tile_ring.py``
   (3 tests — ceiling-lift + tile-granular ipcq_copy count +
   per-CUBE distributed output regression guard).

2. **Rename ``gqa_*`` → ``gqa_attention_*``** across kernel files,
   function names, and importers. The "attention" name makes the role
   explicit (GQA is grouped-query attention) and matches upstream Triton
   FlashAttention naming conventions. Renames:
     _gqa_decode_long.py        -> _gqa_attention_decode_long.py
     _gqa_decode_short.py       -> _gqa_attention_decode_short.py
     _gqa_prefill_long.py       -> _gqa_attention_prefill_long.py
     _gqa_prefill_short.py      -> _gqa_attention_prefill_short.py
   And function names ``gqa_<phase>_<context>_kernel`` →
   ``gqa_attention_<phase>_<context>_kernel``. Updated 1 bench file
   (milestone_gqa_headline.py) and 10 test files.

3. **ADR-0060 / 0062 / 0063 / 0064: Proposed → Accepted**.
   All four are reflected in production code and covered by tests:
   - ADR-0060 (GQA fused attention): 4 kernels deployed; §5.5.1
     amendment added for the tile-granular Ring KV introduced by P3c
     (EN + KO mirror).
   - ADR-0062 (lazy tl.load): LoadFuture + _await_pending live in
     tl_context.py.
   - ADR-0063 (tl.scratch_scope + tl.copy_to): used in every chain
     reduce + tile sweep + ring step. EN-only previously; KO
     translation authored as part of this commit (CLAUDE.md
     bidirectional rule).
   - ADR-0064 (per-op-type CPU issue cost): cpu_issue_cost.py +
     issue_cost_table wiring in tl_context.py (Phase E).
   Files git mv'd from docs/adr-proposed/ to docs/adr/ (EN) and
   docs/adr-ko/ (KO). ADR-0061 (tl.broadcast) stays Proposed — no
   implementation; documented as optional convenience primitive in
   the ADR itself.

Tests: 88/88 focused regression green
(tests/attention/ + Phase E + TL discipline).
ADR pair verification: ``python tools/verify_adr_lang_pairs.py`` OK.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-10 16:20:00 -07:00

119 lines
4.4 KiB
Python

"""Phase 1 spec test for P6a GQA prefill kernel (head-parallel, C=1 baseline).
P6a introduces ``_gqa_attention_prefill_long.py`` with the head-parallel structure (one
Q head per CUBE, per-CUBE distributed output, no reduce). C=1 is the
degenerate case — no Ring KV, no IPCQ traffic. Validates kernel
structure and T_q > 1 attention.
P6b (deferred) adds the Ring KV rotation for C > 1, which needs either
a new SFR install function (intra_* + wrapped E/W at CUBE level) or a
topology-specific config — separate design call.
The prefill kernel differs from decode (P1a/P2a/P2b) in three ways
(ADR-0060 §5.5 / TL;DR):
1. Q has T_q > 1 rows (not just decode's single timestep).
2. Head-parallel placement: each CUBE owns ONE Q head — no M-fold.
3. Each CUBE writes its own head's output — NO reduce.
Phase 1 (this commit): tests only — production code lands in Phase 2.
"""
from __future__ import annotations
from pathlib import Path
from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401
from kernbench.policy.placement.dp import DPPolicy
from kernbench.runtime_api.bench_runner import run_bench
from kernbench.runtime_api.types import resolve_device
from kernbench.sim_engine.engine import GraphEngine
from kernbench.topology.builder import resolve_topology
TOPOLOGY_DEFAULT = Path(__file__).resolve().parents[2] / "topology.yaml"
D_HEAD = 64
DTYPE = "f16"
def _engine_factory(t, d):
return GraphEngine(getattr(t, "topology_obj", t), enable_data=True)
def _run_prefill(*, T_q: int, S_kv: int, C: int = 1):
"""C=1 head-parallel prefill: single CUBE owns the one head + full KV."""
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
dp = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=1)
# Q: (T_q, d_head) — one head per CUBE (head-parallel; for C=1
# only one head total). 2D layout matches what the kernel loads.
# P6b will use a 3D (h_q, T_q, d_head) Q with cube_row_wise
# sharding so each CUBE owns its head.
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp,
name=f"q_t{T_q}_c{C}")
# K, V: full local for C=1 (no ring). Kernel loads K as
# (d_head, S_kv) via byte-conserving reshape.
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp,
name=f"k_t{T_q}_c{C}")
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp,
name=f"v_t{T_q}_c{C}")
# O: (T_q, d_head) — per-CUBE distributed output.
o = ctx.empty((T_q, D_HEAD), dtype=DTYPE, dp=dp,
name=f"o_t{T_q}_c{C}")
ctx.launch(
f"gqa_prefill_p6a_t{T_q}_s{S_kv}_c{C}",
gqa_attention_prefill_long_kernel,
q, k, v, o,
T_q, S_kv, D_HEAD, C,
_auto_dim_remap=False,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None),
engine_factory=_engine_factory,
)
def _count(op_log, name: str) -> int:
return sum(1 for r in op_log if r.op_name == name)
def test_prefill_c_one_t_q_one_completes():
"""C=1, T_q=1: smallest workload (decode-like)."""
result = _run_prefill(T_q=1, S_kv=16, C=1)
assert result.completion.ok, (
f"prefill C=1 T_q=1 failed: {result.completion}"
)
def test_prefill_c_one_t_q_four_completes():
"""C=1, T_q=4: real prefill (Q has multiple rows) — distinguishes
prefill from decode (T_q=1)."""
result = _run_prefill(T_q=4, S_kv=16, C=1)
assert result.completion.ok, (
f"prefill C=1 T_q=4 failed: {result.completion}"
)
def test_prefill_c_one_no_ipcq_traffic():
"""C=1: no ring step, no IPCQ traffic."""
result = _run_prefill(T_q=4, S_kv=16, C=1)
assert result.completion.ok
n_copy = _count(result.engine.op_log, "ipcq_copy")
assert n_copy == 0, (
f"C=1 must have no IPCQ traffic (no ring); got {n_copy}"
)
def test_prefill_c_one_one_dma_write():
"""C=1, one head: exactly one dma_write (per-CUBE distributed output;
no reduce). For C > 1 in P6b this becomes dma_write_count == C."""
result = _run_prefill(T_q=4, S_kv=16, C=1)
assert result.completion.ok
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 1, (
f"C=1 prefill: expected 1 dma_write (one head per cube); "
f"got {n_writes}"
)