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

261 lines
10 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 Phase D: short-context GQA kernels
(ADR-0060 §B "Items from the long/short context split", item B.split.2).
The long-context kernels (§5.2 decode / §5.5 prefill) shard each KV
head row-wise across all CUBEs. At short context (S_kv < 256K), that
shard is too thin to feed the engines and the cube-level collective
overhead dominates. The short-context kernels drop cube-SP entirely:
each CUBE owns ``kv_per_cube`` *whole* KV heads, no S_kv sharding across
CUBEs.
Design (per AskUserQuestion answers in this session):
- PE-parallel heads: P PEs split into ``kv_per_cube`` groups, each
group does PE-SP across (P/kv_per_cube) PEs for one owned head.
- scratch_scope + tl.copy_to discipline mirrors the long kernels.
- One short kernel handles kv_per_cube ∈ {1, 2, 4, 8} via a parameter.
Group layout on the 2×4 PE grid:
kv_per_cube=1, group=8 PEs (full 2×4): row chain + col bridge.
kv_per_cube=2, group=4 PEs (one row): row chain only.
kv_per_cube=4, group=2 PEs (adj cols): 1-step chain.
kv_per_cube=8, group=1 PE: no chain.
After chain reduce-to-group-root, the group's root PE writes its
owned head's output. No inter-CUBE reduce.
Phase 1 (this commit): tests only — production code lands in Phase 2.
All tests fail because the short kernels (``_gqa_attention_decode_short.py`` and
``_gqa_attention_prefill_short.py``) do not exist yet.
"""
from __future__ import annotations
from pathlib import Path
import pytest
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"
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)
# ── Decode short-context kernel ──────────────────────────────────────
def _run_decode_short(*, h_q: int, h_kv: int, kv_per_cube: int,
C: int, P: int, S_kv: int):
"""Run the short-context decode kernel with PE-parallel heads.
Layout (after design iteration — see Phase D failure-recovery):
Q: (T_q, h_q·D_HEAD) replicated; kernel reshapes byte-conservingly.
K, V: (h_kv·S_kv, D_HEAD) head-stacked, ``cube=row_wise, pe=row_wise``
so each PE gets a contiguous (S_local, D_HEAD) chunk at its own
addressable shard (no partial reads needed).
O: replicated; each group root writes the full byte-conserving
(h_q·T_q, D_HEAD) result.
"""
from kernbench.benches._gqa_attention_decode_short import gqa_attention_decode_short_kernel # Phase 2
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=C, num_pes=P)
# Head-stacked KV with row_wise sharding so each PE's chunk is
# contiguous and exactly (S_local, D_HEAD) addressable.
dp_kv = DPPolicy(cube="row_wise", pe="row_wise",
num_cubes=C, num_pes=P)
q = ctx.zeros((1, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"q_short_{kv_per_cube}_{C}")
k = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"k_short_{kv_per_cube}_{C}")
v = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"v_short_{kv_per_cube}_{C}")
o = ctx.empty((1, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"o_short_{kv_per_cube}_{C}")
ctx.launch(
f"gqa_decode_short_{kv_per_cube}_{C}",
gqa_attention_decode_short_kernel,
q, k, v, o,
1, S_kv, h_q, h_kv, D_HEAD, C, P, kv_per_cube,
_auto_dim_remap=False,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None), engine_factory=_engine_factory,
)
def test_short_decode_smoke_kv_per_cube_2_C_4():
"""ADR-0060 §B.split.2 headline: kv_per_cube=2, C=4 — each CUBE
owns 2 heads (8 heads / 4 CUBEs). PE-parallel heads splits the 8
PEs into 2 groups of 4, each group does PE-SP for one owned head.
Smoke: kernel completes."""
result = _run_decode_short(
h_q=8, h_kv=8, kv_per_cube=2, C=4, P=8, S_kv=64,
)
assert result.completion.ok, (
f"short decode kv_per_cube=2 must complete; got {result.completion}"
)
def test_short_decode_smoke_kv_per_cube_4_C_2():
"""kv_per_cube=4, C=2 — half the CUBEs participate, each owns 4
heads, 4 PE groups of 2 PEs each."""
result = _run_decode_short(
h_q=8, h_kv=8, kv_per_cube=4, C=2, P=8, S_kv=64,
)
assert result.completion.ok, (
f"short decode kv_per_cube=4 must complete; got {result.completion}"
)
def test_short_decode_smoke_kv_per_cube_8_C_1():
"""kv_per_cube=8, C=1 — all heads on one CUBE. 8 PE groups of 1 PE
each → no chain reduce, each PE writes its head's output."""
result = _run_decode_short(
h_q=8, h_kv=8, kv_per_cube=8, C=1, P=8, S_kv=64,
)
assert result.completion.ok, (
f"short decode kv_per_cube=8 must complete; got {result.completion}"
)
def test_short_decode_dma_write_count_equals_h_kv():
"""ADR-0060 §B.split.2: each owned head produces exactly one
output (no inter-CUBE reduce). Total dma_writes = h_kv across all
participating CUBEs and groups.
For h_kv=8, kv_per_cube=2, C=4: each CUBE writes 2 outputs →
4 × 2 = 8 dma_writes total.
"""
result = _run_decode_short(
h_q=8, h_kv=8, kv_per_cube=2, C=4, P=8, S_kv=64,
)
assert result.completion.ok
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 8, (
f"short decode: expected 8 dma_writes (h_kv); got {n_writes}"
)
def test_short_decode_no_inter_cube_traffic():
"""ADR-0060 §B.split.2: each head is fully owned by one CUBE → no
inter-CUBE reduce. The kernel must not invoke CUBE-level E/W IPCQ.
Today's long-context kernel at C=4 emits ~12 inter-CUBE ipcq_copy
via direction "E"/"W". The short kernel must emit zero of those,
keeping all IPCQ traffic on the ``intra_*`` directions.
"""
result = _run_decode_short(
h_q=8, h_kv=8, kv_per_cube=2, C=4, P=8, S_kv=64,
)
assert result.completion.ok
# Count IPCQ traffic that targeted the CUBE-level "E"/"W" directions.
inter_cube_ipcq = sum(
1 for r in result.engine.op_log
if r.op_name == "ipcq_copy"
and r.params.get("direction") in ("E", "W")
)
assert inter_cube_ipcq == 0, (
f"short decode must have no inter-CUBE E/W IPCQ; got {inter_cube_ipcq}"
)
# ── Prefill short-context kernel ─────────────────────────────────────
def _run_prefill_short(*, h_kv: int, kv_per_cube: int,
C: int, P: int, T_q: int, S_kv: int):
"""Run the short-context prefill kernel with PE-parallel heads.
Layout: same head-stacked K/V scheme as decode short, with Q/O
replicated and byte-conserving reshape inside the kernel.
"""
from kernbench.benches._gqa_attention_prefill_short import gqa_attention_prefill_short_kernel # Phase 2
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_q = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=P)
dp_kv = DPPolicy(cube="row_wise", pe="row_wise",
num_cubes=C, num_pes=P)
dp_o = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=P)
q = ctx.zeros((T_q, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_q, name=f"q_pre_short_{kv_per_cube}_{C}")
k = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"k_pre_short_{kv_per_cube}_{C}")
v = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"v_pre_short_{kv_per_cube}_{C}")
o = ctx.empty((T_q, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_o, name=f"o_pre_short_{kv_per_cube}_{C}")
ctx.launch(
f"gqa_prefill_short_{kv_per_cube}_{C}",
gqa_attention_prefill_short_kernel,
q, k, v, o,
T_q, S_kv, h_kv, D_HEAD, C, P, kv_per_cube,
_auto_dim_remap=False,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None), engine_factory=_engine_factory,
)
def test_short_prefill_smoke_kv_per_cube_2_C_4():
"""Short prefill kv_per_cube=2, C=4. Smoke: completes."""
result = _run_prefill_short(
h_kv=8, kv_per_cube=2, C=4, P=8, T_q=4, S_kv=64,
)
assert result.completion.ok, (
f"short prefill kv_per_cube=2 must complete; got {result.completion}"
)
def test_short_prefill_no_ring_KV_traffic():
"""ADR-0060 §B.split.2: short prefill DOES NOT use Ring KV — each
CUBE owns its KV heads fully, no rotation. The kernel must not
emit any KV-rotation IPCQ traffic at the CUBE level.
"""
result = _run_prefill_short(
h_kv=8, kv_per_cube=2, C=4, P=8, T_q=4, S_kv=64,
)
assert result.completion.ok
inter_cube_ipcq = sum(
1 for r in result.engine.op_log
if r.op_name == "ipcq_copy"
and r.params.get("direction") in ("E", "W")
)
assert inter_cube_ipcq == 0, (
f"short prefill must have no Ring KV (no inter-CUBE E/W IPCQ); "
f"got {inter_cube_ipcq}"
)