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

163 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 P3c: tile-granular Ring KV in prefill_long
(ADR-0060 §5.5 amendment + ADR-0063 §A.2).
Today's prefill_long ring sends and receives full ``(d_head, S_local)``
KV slices per step. Step 0's local-attention intermediates also live
outside any ``scratch_scope`` (they're loaded as full-slice ``Kc``,
``Vc`` and feed ``scores``, ``exp_scores`` as persistent allocations).
At larger ``S_local`` both the step-0 leak and the ring step's
in-scope intermediates grow linearly with ``S_local``, and at
``S_local = 32K`` (S_kv=128K, C=4, T_q=4) the peak overflows the
1 MiB pool.
P3c converts the ring to **tile-granular**: a nested loop
``for k in range(C): for t in range(n_tiles): ...`` where each
iteration sends/recvs one ``(d_head, TILE_S_KV)`` tile (and its V
counterpart). The persistent state shrinks to ``(m, , O)`` only
(~1 KB); per-tile in-scope scratch is bounded by
``TILE_S_KV`` regardless of ``S_local``. Ceiling lifted.
Trade-off: the per-CUBE send count grows from ``2·(C-1)`` to
``2·n_tiles·(C-1)``. At ``n_tiles=1`` (small ``S_local``) the count is
unchanged, so the existing ``test_prefill_ring_c_*`` tests at
``S_kv ∈ {16, 32}`` still pass.
Phase 1 (this commit): tests only — production code lands in Phase 2.
T1 fails today with a ``TLContext scratch overflow``; T2 fails today
with the slice-granular ipcq_copy count; T5 passes today and is a
regression guard for the per-CUBE output write count.
"""
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.ccl.install import load_ccl_config, resolve_algorithm_config
from kernbench.ccl.sfr_config import configure_sfr_intercube_ring
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 _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 _count(op_log, name: str) -> int:
return sum(1 for r in op_log if r.op_name == name)
def _run_prefill_ring(*, T_q: int, S_kv: int, C: int):
"""Head-parallel prefill with Ring KV across C CUBEs."""
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
configure_sfr_intercube_ring(
ctx.engine, ctx.spec, _ccl_cfg(), ring_size=C,
)
dp_q = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=1)
dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate",
pe="replicate", num_cubes=C, num_pes=1)
dp_o = DPPolicy(cube="row_wise" if C > 1 else "replicate",
pe="replicate", num_cubes=C, num_pes=1)
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q,
name=f"q_tlr_t{T_q}_c{C}_s{S_kv}")
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
name=f"k_tlr_t{T_q}_c{C}_s{S_kv}")
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
name=f"v_tlr_t{T_q}_c{C}_s{S_kv}")
o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o,
name=f"o_tlr_t{T_q}_c{C}_s{S_kv}")
ctx.launch(
f"gqa_prefill_long_tile_ring_t{T_q}_c{C}_s{S_kv}",
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,
)
# ── T1: 128K ceiling lift ────────────────────────────────────────────
def test_prefill_long_context_128k_completes():
"""ADR-0063 §A.2 headline ceiling lift. At S_kv=128K with C=4,
S_local=32K. Today: step-0 score-stack (~768 KB persistent) + ring
scope (~768 KB) → 1.5 MB peak → TLContext scratch overflow. After
P3c the persistent state is just ``(m, , O)`` (≈ 1 KB) and the
per-tile in-scope scratch is bounded by TILE_S_KV.
"""
result = _run_prefill_ring(T_q=4, S_kv=131_072, C=4)
assert result.completion.ok, (
f"prefill_long at S_kv=128K must complete after tile-granular "
f"ring lands; got {result.completion}"
)
# ── T2: tile-granular ipcq_copy count ────────────────────────────────
def test_prefill_long_tile_granular_ipcq_count():
"""ADR-0060 §5.5 amendment: with tile-granular sends, the per-CUBE
send count grows from ``2·(C-1)`` to ``2·n_tiles·(C-1)``. Aggregated
across all C CUBEs the total ipcq_copy becomes
``2·n_tiles·(C-1)·C``.
Config: T_q=4, S_kv=4096, C=2 → S_local=2048, n_tiles=2 (with
TILE_S_KV=1024). Today: slice-granular total =
``(C-1)·2·C = 4``. After P3c: tile-granular total =
``(C-1)·n_tiles·2·C = 8``.
"""
C = 2
n_tiles = 2 # S_local=2048 / TILE_S_KV=1024
result = _run_prefill_ring(T_q=4, S_kv=4096, C=C)
assert result.completion.ok, (
f"prefill_long multi-tile ring must complete; got {result.completion}"
)
n_copy = _count(result.engine.op_log, "ipcq_copy")
expected = (C - 1) * n_tiles * 2 * C
assert n_copy == expected, (
f"tile-granular ring: expected {expected} ipcq_copy "
f"((C-1)·n_tiles·2·C = {C - 1}·{n_tiles}·2·{C}); got {n_copy}"
)
# ── T5: per-CUBE distributed output unchanged (regression guard) ─────
def test_prefill_long_tile_ring_dma_write_count():
"""ADR-0060 §5.5: per-CUBE distributed output must hold under the
tile-granular ring rewrite. Each CUBE still writes its own head's
output (no inter-CUBE reduce); dma_write_count == C.
Today passes; must continue to pass after P3c.
"""
C = 4
result = _run_prefill_ring(T_q=4, S_kv=4096, C=C)
assert result.completion.ok
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == C, (
f"per-CUBE distributed output: expected {C} dma_writes (one per "
f"CUBE); got {n_writes}"
)