Files
kernbench2/tests/attention/test_gqa_decode_sp.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

154 lines
6.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Phase 1 spec test for P2a GQA decode SP (chain reduce-to-root, Level-2 only).
P2a is the first half of DDD-0060 P2: the kernel becomes multi-PE within
one CUBE and reduces to root (PE 0) using a chain over the 1D intra-cube
ring (W direction). This **replaces the baseline's bidirectional O(N)
fan-out** where every rank ends with O — ADR-0060 §A.2's headline.
Deviation from DDD-0060 §7 P2 gate: the gate text asks for
``⌈log₂ P⌉`` reduce rounds. The intra-cube SFR install
(``configure_sfr_intracube_pe_ring``) wires only a 1D E/W ring, so a
true tree would require either multi-hop forwarding or a new SFR install
(future ADR). P2a uses a **chain reduce-to-root**: ``P-1`` rounds along
W. The architectural property the ADR cares about
(root-only output vs every-rank-has-O) is preserved; the logarithmic
collective is deferred.
P2b (deferred) covers Level-1 inter-CUBE center-mesh reduce (C>1).
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_decode_long import gqa_attention_decode_long_kernel # noqa: F401
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
from kernbench.ccl.sfr_config import configure_sfr_intercube_multisip
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"
T_Q = 1
D_HEAD = 64
DTYPE = "f16"
def _ccl_cfg():
return resolve_algorithm_config(
load_ccl_config(), name="lrab_hierarchical_allreduce",
)
def _engine_factory(t, d):
return GraphEngine(getattr(t, "topology_obj", t), enable_data=True)
def _run_decode_sp(*, h_q: int, h_kv: int, P: int, S_kv: int):
"""Single-CUBE SP decode: P PEs share the work along the intra-cube ring."""
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=1, num_pes=P)
dp_kv = DPPolicy(cube="replicate", pe="row_wise",
num_cubes=1, num_pes=P)
q = ctx.zeros((T_Q, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"q_h{h_q}_kv{h_kv}_p{P}")
# KV: total S_kv split across P PEs along axis 0 (row_wise sharding).
# Each PE sees (S_kv/P, h_kv·D_HEAD).
k = ctx.zeros((S_kv, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"k_h{h_q}_kv{h_kv}_p{P}")
v = ctx.zeros((S_kv, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"v_h{h_q}_kv{h_kv}_p{P}")
o = ctx.empty((T_Q, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"o_h{h_q}_kv{h_kv}_p{P}")
ctx.launch(
f"gqa_decode_sp_h{h_q}_kv{h_kv}_p{P}",
gqa_attention_decode_long_kernel,
q, k, v, o,
T_Q, S_kv, h_q, h_kv, D_HEAD,
1, P, # C=1, P=P (single-CUBE SP)
_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)
# ── Root-only write ────────────────────────────────────────────────────
def test_sp_chain_reduce_root_only_writes_o():
"""ADR-0060 §A.2: only the root rank (PE 0) writes O. Baseline today
has every rank write the full final O (bidirectional fan-out)."""
result = _run_decode_sp(h_q=1, h_kv=1, P=8, S_kv=64)
assert result.completion.ok, f"P=8 chain reduce failed: {result.completion}"
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 1, (
f"reduce-to-root must produce exactly 1 dma_write (PE 0); "
f"got {n_writes}"
)
# ── Chain step count ───────────────────────────────────────────────────
def test_sp_chain_reduce_p_minus_one_ipcq_pairs():
"""Chain reduce-to-root has P-1 send→recv pairs along the W chain;
each pair logs one ``ipcq_copy`` (inbound DMA, per
``milestone_gqa_llama70b._summarize_op_log``). Each chain step ships
the triplet (m, , O) → 3 handles per step → 7 steps × 3 = 21."""
result = _run_decode_sp(h_q=1, h_kv=1, P=8, S_kv=64)
assert result.completion.ok, f"P=8 chain reduce failed: {result.completion}"
n_copy = _count(result.engine.op_log, "ipcq_copy")
expected = (8 - 1) * 3
assert n_copy == expected, (
f"chain reduce: expected {expected} ipcq_copy (P-1=7 steps × "
f"3 handles m//O); got {n_copy}"
)
# ── Real GQA × SP combined ─────────────────────────────────────────────
def test_sp_real_gqa_h_q_eight_h_kv_one_p_eight():
"""The combined unlock: real GQA (h_q=G·h_kv with G=8) AND SP
(P=8) together — neither expressible by the baseline."""
result = _run_decode_sp(h_q=8, h_kv=1, P=8, S_kv=64)
assert result.completion.ok, (
f"real GQA + SP combined run failed: {result.completion}"
)
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 1, (
f"root-only write must hold under M-fold too; got {n_writes}"
)
# ── Degenerate P=1 ────────────────────────────────────────────────────
def test_sp_p_one_degenerate_no_ipcq_traffic():
"""P=1: SP degenerates to a single rank. No IPCQ traffic; one dma_write."""
result = _run_decode_sp(h_q=8, h_kv=1, P=1, S_kv=16)
assert result.completion.ok, f"P=1 degenerate failed: {result.completion}"
n_send = _count(result.engine.op_log, "ipcq_send")
n_recv = _count(result.engine.op_log, "ipcq_recv")
assert n_send == 0, f"P=1 must have no ipcq_send; got {n_send}"
assert n_recv == 0, f"P=1 must have no ipcq_recv; got {n_recv}"
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 1, f"P=1: one dma_write; got {n_writes}"