Files
kernbench2/tests/attention/test_milestone_gqa_decode_4cases.py
T
mukesh 3c155be8e6 gqa(decode-4cases): give Case 4 its own thin kernel file for naming symmetry
Cases 1-3 each have a dedicated _gqa_attention_decode_<case>.py
kernel file; Case 4 previously reached into _gqa_attention_decode_long.py
via the headline bench's _run_decode_panel helper, breaking the
one-file-per-case convention. Adds _gqa_attention_decode_cube_sp_pe_sp.py
as a 20-line wrapper that bakes in sub_w=4 (the C=8 lrab geometry)
and gives Case 4 its own kind ("decode_cube_sp_pe_sp") and helper
(_run_decode_panel_cube_sp_pe_sp). decode_long.py is unchanged
(still serves the legacy decode_long tests). 16 tests pass.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-15 15:58:11 -07:00

540 lines
19 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.
"""Tests for the decode 4-cases comparative-study bench.
Per ``GQA_full_deck.pptx`` slides 11-17, the 4 cases differ in how KV
cache is sharded across the 8 cubes and 8 PEs of a single KV-head
group on LLaMA-3.1-70B GQA:
Case 1 Cube-SP / PE-TP → KV split by S_kv across cubes; PEs TP on batch
Case 2 Cube-Repl / PE-TP → full KV per cube; PEs TP on batch
Case 3 Cube-Repl / PE-SP → full KV per cube; PEs SP on S_kv (intra-cube AR)
Case 4 Cube-SP / PE-SP → KV split 64-way; 2-phase AR on (m,,O) ★ optimal
This file grows as each case lands. Phase 1 of 5C.D adds Case 4 first,
which is structurally the existing ``_gqa_attention_decode_long.py`` at
``sub_w=4`` (Increment 2's lrab-adapted center-root reduce). The Phase 2
production change introduces a new bench file
``src/kernbench/benches/milestone_gqa_decode_4cases.py`` housing the 4
case panels under a single ``milestone-gqa-decode-4cases`` entry, and
extends ``_run_decode_panel`` in ``milestone_gqa_headline`` to accept
``sub_w``/``d_head``/``h_q``/``h_kv`` overrides.
Deviation from slide 13: slide prescribes AllReduce on (m,,O); the
kernel does reduce-to-root (only the lrab center cube has the answer)
per ADR-0060 §4. Treated as the kernbench Case-4 baseline.
Phase 1: tests only. T1, T2, T3, T4 fail today (new bench file does
not yet exist; ``_run_decode_panel`` does not yet accept ``sub_w``).
"""
from __future__ import annotations
import re
from pathlib import Path
from kernbench.benches.milestone_gqa_headline import _run_decode_panel
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"
_CASE4_PANEL = "single_kv_group_decode_gqa_cube_sp_pe_sp"
_CASE3_PANEL = "single_kv_group_decode_gqa_cube_repl_pe_sp"
_CASE2_PANEL = "single_kv_group_decode_gqa_cube_repl_pe_tp"
_CASE1_PANEL = "single_kv_group_decode_gqa_cube_sp_pe_tp"
_CUBE_RE = re.compile(r"\bcube(\d+)\b")
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 _dma_write_cubes(op_log) -> list[int]:
cubes: list[int] = []
for r in op_log:
if r.op_name != "dma_write":
continue
m = _CUBE_RE.search(r.component_id)
if m is not None:
cubes.append(int(m.group(1)))
return cubes
def _run_case4_smoke(*, S_kv: int):
"""Drive the Case 4 decode panel via ``_run_decode_panel``.
Uses ``S_kv=8192`` (smoke) to keep test time bounded; the headline
``S_kv=128K`` runs come from ``kernbench run --bench
milestone-gqa-decode-4cases``, not pytest.
"""
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
_run_decode_panel(
ctx, panel=_CASE4_PANEL,
C=8, P=8, sub_w=4,
S_kv=S_kv,
d_head=128, h_q=8, h_kv=1,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None),
engine_factory=_engine_factory,
)
# ── T1: Case 4 panel is registered in the new bench ──────────────────
def test_case4_panel_registered():
"""The Case 4 panel must be in the new bench's ``_PANELS`` +
``_PANEL_DISPATCH`` with the expected LLaMA-3.1-70B target dims.
Headline config:
C = 8 (head-parallel CUBE Group)
P = 8 (intra-CUBE PE-SP)
T_q = 1 (decode: one new token per pass)
S_kv = 131_072 (LLaMA long-context decode target)
d_head = 128, h_q = 8, h_kv = 1
The lrab sub_w=4 / sub_h=2 geometry is baked into the Case 4
wrapper kernel (``_gqa_attention_decode_cube_sp_pe_sp.py``); it
is not a panel parameter.
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
assert _CASE4_PANEL in _PANELS, (
f"{_CASE4_PANEL!r} not in _PANELS; got {_PANELS}"
)
assert _CASE4_PANEL in _PANEL_DISPATCH
kind, params = _PANEL_DISPATCH[_CASE4_PANEL]
assert kind == "decode_cube_sp_pe_sp", (
f"kind={kind!r}, expected 'decode_cube_sp_pe_sp'"
)
assert params.get("C") == 8
assert params.get("P") == 8
assert params.get("T_q") == 1
assert params.get("S_kv") == 131_072
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
# ── T2: Case 4 runner drives the kernel to completion ───────────────
def test_case4_runner_smoke():
"""``_run_decode_panel`` must accept the new ``sub_w``,
``d_head``, ``h_q``, ``h_kv`` kwargs and launch the kernel at
``(C, P, sub_w) = (8, 8, 4)`` (the Case 4 / lrab path).
Smoke uses ``S_kv=8192`` so the simulation completes quickly. The
headline 128K dims run via ``kernbench run --bench
milestone-gqa-decode-4cases``.
"""
result = _run_case4_smoke(S_kv=8192)
assert result.completion.ok, (
f"Case 4 decode smoke at C=8 P=8 sub_w=4 must complete; "
f"got {result.completion}"
)
# ── T3: reduce-to-root lands at the lrab center cube (cube 6) ───────
def test_case4_root_at_center_cube_6():
"""For ``sub_w=4, sub_h=2``: root_col=2, root_row=1, root_cube=6.
The decode kernel writes the final O exclusively from PE 0 of cube
6 (ADR-0060 §4 reduce-to-root variant of the Case-4 AR pattern).
"""
result = _run_case4_smoke(S_kv=8192)
assert result.completion.ok
cubes = _dma_write_cubes(result.engine.op_log)
assert cubes, "expected at least one dma_write for the final O store"
distinct = set(cubes)
assert distinct == {6}, (
f"Case 4 root must be the lrab center cube 6; "
f"got cubes={sorted(distinct)}"
)
# ── T4: 2-phase AR ipcq pattern matches the predicted Case-4 traffic ─
def _run_case2_smoke(*, S_kv: int):
"""Drive the Case 2 decode panel via the case-specific runner.
Case 2 = Cube-Repl × PE-TP. K, V are replicated everywhere (the
slide-11 memory waste); for B=1 only one rank does the work; no
inter-rank comm.
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_run_decode_panel_cube_repl_pe_tp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
_run_decode_panel_cube_repl_pe_tp(
ctx, panel=_CASE2_PANEL,
C=8, P=8,
T_q=1, S_kv=S_kv,
d_head=128, h_q=8, h_kv=1,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None),
engine_factory=_engine_factory,
)
# ── Case 2 — T1: panel registered ───────────────────────────────────
def test_case2_panel_registered():
"""The Case 2 panel must be in the bench's ``_PANELS`` +
``_PANEL_DISPATCH`` with the expected single-KV-group dims.
Case 2: Cube-Repl × PE-TP. K, V replicated everywhere
(8 KB/tok/PE — slide-11 memory waste); no inter-rank comm.
For B=1 only one rank works (PEs 1-7 idle — slide-11 calls
out this PE-TP waste).
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
assert _CASE2_PANEL in _PANELS, (
f"{_CASE2_PANEL!r} not in _PANELS; got {_PANELS}"
)
assert _CASE2_PANEL in _PANEL_DISPATCH
kind, params = _PANEL_DISPATCH[_CASE2_PANEL]
assert kind == "decode_cube_repl_pe_tp"
assert params.get("C") == 8
assert params.get("P") == 8
assert params.get("T_q") == 1
assert params.get("S_kv") == 131_072
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
# ── Case 2 — T2: smoke runner completes ─────────────────────────────
def test_case2_runner_smoke():
"""Case 2 runner drives the new kernel to completion at smoke S_kv."""
result = _run_case2_smoke(S_kv=8192)
assert result.completion.ok, (
f"Case 2 decode smoke at C=8 P=8 must complete; "
f"got {result.completion}"
)
# ── Case 2 — T3: zero inter-rank comm by design ─────────────────────
def test_case2_zero_ipcq_copy_no_comm():
"""Case 2's defining property: full KV per rank ⇒ NO inter-rank
communication. Slide 11 lists comm cost as 'none'.
"""
result = _run_case2_smoke(S_kv=8192)
assert result.completion.ok
n_copy = _count(result.engine.op_log, "ipcq_copy")
assert n_copy == 0, (
f"Case 2 must have zero inter-rank comm; got ipcq_copy={n_copy}"
)
# ── Case 2 — T4: single dma_write from cube 0 (B=1 single-rank work) ─
def test_case2_single_dma_write_at_cube_0():
"""For B=1, only PE 0 of CUBE 0 does the work (the inherent PE-TP
waste at B=1). Exactly 1 dma_write, from cube 0.
"""
result = _run_case2_smoke(S_kv=8192)
assert result.completion.ok
cubes = _dma_write_cubes(result.engine.op_log)
assert cubes, "expected at least one dma_write for the final O store"
distinct = set(cubes)
assert distinct == {0}, (
f"Case 2 B=1 single writer must be cube 0; "
f"got cubes={sorted(distinct)}"
)
# ── Case 3 (Cube-Repl × PE-SP) ──────────────────────────────────────
def _run_case3_smoke(*, S_kv: int):
"""Drive the Case 3 decode panel via the case-specific runner.
Case 3 = Cube-Repl × PE-SP. K, V replicated per cube; S_kv split
8-way across PEs within each cube. Intra-CUBE 8-way reduce on
(m, , O); no inter-CUBE comm (every cube ends with full answer).
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_run_decode_panel_cube_repl_pe_sp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
_run_decode_panel_cube_repl_pe_sp(
ctx, panel=_CASE3_PANEL,
C=8, P=8,
T_q=1, S_kv=S_kv,
d_head=128, h_q=8, h_kv=1,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None),
engine_factory=_engine_factory,
)
# ── Case 3 — T1: panel registered ───────────────────────────────────
def test_case3_panel_registered():
"""The Case 3 panel must be in the bench's ``_PANELS`` +
``_PANEL_DISPATCH`` with the expected single-KV-group dims.
Case 3: Cube-Repl × PE-SP. K, V replicated per cube; PEs SP on
S_kv. Intra-CUBE 8-way AR; no inter-CUBE comm.
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
assert _CASE3_PANEL in _PANELS, (
f"{_CASE3_PANEL!r} not in _PANELS; got {_PANELS}"
)
assert _CASE3_PANEL in _PANEL_DISPATCH
kind, params = _PANEL_DISPATCH[_CASE3_PANEL]
assert kind == "decode_cube_repl_pe_sp"
assert params.get("C") == 8
assert params.get("P") == 8
assert params.get("T_q") == 1
assert params.get("S_kv") == 131_072
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
# ── Case 3 — T2: smoke runner completes ─────────────────────────────
def test_case3_runner_smoke():
"""Case 3 runner drives the new kernel to completion at smoke S_kv."""
result = _run_case3_smoke(S_kv=8192)
assert result.completion.ok, (
f"Case 3 decode smoke at C=8 P=8 must complete; "
f"got {result.completion}"
)
# ── Case 3 — T3: intra-CUBE-only AR (168 ipcq, no inter-CUBE) ───────
def test_case3_intra_cube_ar_only_ipcq():
"""Case 3 reduce pattern: per-CUBE 8-way PE-SP AR (same structural
cost as Case 4's intra-CUBE phase = 21 ipcq_copy per cube), and
NO inter-CUBE traffic (each cube has a full copy of KV).
per-CUBE intra (2×4 PE grid):
row chain along intra_W: cols 1,2,3 each row × 2 rows ×
3 tensors (m, , O) = 18
col bridge along intra_N: pe4 only × 3 tensors = 3
per-CUBE intra total = 21
× 8 CUBEs = 168
inter-CUBE: 0 (replicated KV ⇒ no AllReduce needed).
Grand total: 168.
"""
result = _run_case3_smoke(S_kv=8192)
assert result.completion.ok
n_copy = _count(result.engine.op_log, "ipcq_copy")
assert n_copy == 168, (
f"Case 3 expected 168 ipcq_copy (intra-CUBE only, no inter-CUBE); "
f"got {n_copy}"
)
# ── Case 3 — T4: single dma_write from cube 0 (designated writer) ───
def test_case3_single_dma_write_at_cube_0():
"""Every cube ends with the full answer after intra-CUBE AR; only
the designated writer (cube 0, PE 0) stores O to avoid 8 redundant
DMAs.
"""
result = _run_case3_smoke(S_kv=8192)
assert result.completion.ok
cubes = _dma_write_cubes(result.engine.op_log)
assert cubes, "expected at least one dma_write for the final O store"
distinct = set(cubes)
assert distinct == {0}, (
f"Case 3 designated writer must be cube 0; "
f"got cubes={sorted(distinct)}"
)
# ── Case 1 (Cube-SP × PE-TP) ────────────────────────────────────────
def _run_case1_smoke(*, S_kv: int):
"""Drive the Case 1 decode panel via the case-specific runner.
Case 1 = Cube-SP × PE-TP. K, V split S_kv-wise across the 8 cubes
(cube=row_wise), replicated within each cube (pe=replicate). PEs
nominally split on the batch dim (PE-TP); at B=1 only PE 0 of
each cube has work; PEs 1-7 idle. Inter-CUBE 8-way reduce via the
lrab-adapted center-root pattern (root cube 6); no intra-CUBE comm.
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_run_decode_panel_cube_sp_pe_tp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
_run_decode_panel_cube_sp_pe_tp(
ctx, panel=_CASE1_PANEL,
C=8, P=8,
T_q=1, S_kv=S_kv,
d_head=128, h_q=8, h_kv=1,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None),
engine_factory=_engine_factory,
)
# ── Case 1 — T1: panel registered ───────────────────────────────────
def test_case1_panel_registered():
"""The Case 1 panel must be in the bench's ``_PANELS`` +
``_PANEL_DISPATCH`` with the expected single-KV-group dims.
Case 1: Cube-SP × PE-TP. K, V split across cubes; PEs TP on
batch. At B=1 only PE 0 of each cube works (PE-TP waste).
Inter-CUBE lrab AR; no intra-CUBE comm.
"""
from kernbench.benches.milestone_gqa_decode_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
assert _CASE1_PANEL in _PANELS, (
f"{_CASE1_PANEL!r} not in _PANELS; got {_PANELS}"
)
assert _CASE1_PANEL in _PANEL_DISPATCH
kind, params = _PANEL_DISPATCH[_CASE1_PANEL]
assert kind == "decode_cube_sp_pe_tp"
assert params.get("C") == 8
assert params.get("P") == 8
assert params.get("T_q") == 1
assert params.get("S_kv") == 131_072
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
# ── Case 1 — T2: smoke runner completes ─────────────────────────────
def test_case1_runner_smoke():
"""Case 1 runner drives the new kernel to completion at smoke S_kv."""
result = _run_case1_smoke(S_kv=8192)
assert result.completion.ok, (
f"Case 1 decode smoke at C=8 P=8 must complete; "
f"got {result.completion}"
)
# ── Case 1 — T3: inter-CUBE lrab only (21 ipcq, no intra-CUBE) ──────
def test_case1_inter_cube_lrab_only_ipcq():
"""Case 1 reduce pattern: only PE 0 of each cube has work (PE-TP
at B=1), so NO intra-CUBE AR. Inter-CUBE 8-way reduce uses the
lrab-adapted center-root pattern (same structural cost as Case 4's
inter-CUBE phase).
Inter-CUBE lrab (sub_w=4, sub_h=2):
Phase 1 row reduce — 3 sends/row × 3 tensors × 2 rows = 18
Phase 2 col reduce — cube 2 → S × 3 tensors = 3
inter-CUBE total = 21
Intra-CUBE: 0 (PEs 1-7 idle, no partials to merge).
Grand total: 21.
"""
result = _run_case1_smoke(S_kv=8192)
assert result.completion.ok
n_copy = _count(result.engine.op_log, "ipcq_copy")
assert n_copy == 21, (
f"Case 1 expected 21 ipcq_copy (inter-CUBE lrab only, no intra-CUBE); "
f"got {n_copy}"
)
# ── Case 1 — T4: root at lrab center cube 6 ─────────────────────────
def test_case1_root_at_center_cube_6():
"""Case 1 uses the same lrab-adapted center-root reduce as Case 4's
inter-CUBE phase; the answer lands at the lrab center cube
(sub_w=4, sub_h=2 → root_col=2, root_row=1 → root_cube=6).
"""
result = _run_case1_smoke(S_kv=8192)
assert result.completion.ok
cubes = _dma_write_cubes(result.engine.op_log)
assert cubes, "expected at least one dma_write for the final O store"
distinct = set(cubes)
assert distinct == {6}, (
f"Case 1 root must be the lrab center cube 6; "
f"got cubes={sorted(distinct)}"
)
# ── Case 4 — T4 (existing, kept) ────────────────────────────────────
def test_case4_two_level_ar_ipcq_pattern():
"""Total ipcq_copy for the Case 4 reduce at (C, P, sub_w) =
(8, 8, 4):
Intra-CUBE (per CUBE = 8 PEs in a 2×4 grid):
row chain along intra_W: cols 1,2,3 each row × 2 rows ×
3 tensors = 18
col bridge along intra_N: pe4 only × 3 tensors = 3
per-CUBE intra total = 21
× 8 CUBEs = 168
Inter-CUBE lrab (sub_w=4, sub_h=2):
Phase 1 row reduce — 3 sends/row × 3 tensors × 2 rows = 18
Phase 2 col reduce — cube 2 → S × 3 tensors = 3
inter-CUBE total = 21
Grand total: 168 + 21 = 189
"""
result = _run_case4_smoke(S_kv=8192)
assert result.completion.ok
n_copy = _count(result.engine.op_log, "ipcq_copy")
assert n_copy == 189, (
f"Case 4 expected 189 ipcq_copy "
f"(168 intra-CUBE + 21 inter-CUBE lrab); got {n_copy}"
)