gqa: reorganize benches into gqa_helpers/ subpackage; drop legacy headline

Splits the GQA helpers into a dedicated subpackage to make room for the
prefill 4-cases study (next commit) and a single umbrella bench
(milestone-1h-gqa, after that).

Layout:
  benches/gqa_helpers/
    long_ctx/    — decode 4-cases kernels + sweep runner
    short_ctx/   — prefill/decode short-context kernels
    shared/      — _gqa_panel_helpers + decode_opt2 (context-agnostic)

The registry audit now skips subpackages so gqa_helpers/ (without a
leading underscore) doesn't get audited for @bench decorators.

Also drops the legacy milestone-gqa-headline bench, its
_gqa_attention_prefill_long kernel, 6 dependent prefill tests, the
paper_gqa_latency.py report harness, and the 3 stale headline-derived
PNGs the §6 wire-up referenced (paper will re-pull from the new
1H_milestone_output/gqa/long_ctx/ once §6 is updated).

The _ccl_cfg and _summarize_op_log helpers used to live in the
headline bench; extracted them to gqa_helpers/shared/_gqa_panel_helpers.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
2026-06-16 13:05:41 -07:00
parent 359a0eaa44
commit e45626c036
32 changed files with 178 additions and 1674 deletions
@@ -68,7 +68,7 @@ def _run_case4_smoke(*, S_kv: int):
``S_kv=128K`` runs come from ``kernbench run --bench
milestone-gqa-decode-long-ctx-4cases``, not pytest.
"""
from kernbench.benches.milestone_gqa_decode_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_run_decode_panel_long_ctx_cube_sp_pe_sp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
@@ -99,13 +99,13 @@ def test_case4_panel_registered():
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)
S_kv = 8_192 (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; it is not a panel parameter.
"""
from kernbench.benches.milestone_gqa_decode_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
@@ -120,7 +120,7 @@ def test_case4_panel_registered():
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("S_kv") == 8_192
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
@@ -197,7 +197,7 @@ def _run_case2_smoke(*, S_kv: int):
slide-11 memory waste); for B=1 only one rank does the work; no
inter-rank comm.
"""
from kernbench.benches.milestone_gqa_decode_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_run_decode_panel_long_ctx_cube_repl_pe_tp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
@@ -229,7 +229,7 @@ def test_case2_panel_registered():
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_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
@@ -242,7 +242,7 @@ def test_case2_panel_registered():
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("S_kv") == 8_192
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
@@ -303,7 +303,7 @@ def _run_case3_smoke(*, S_kv: int):
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_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_run_decode_panel_long_ctx_cube_repl_pe_sp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
@@ -333,7 +333,7 @@ def test_case3_panel_registered():
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_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
@@ -346,7 +346,7 @@ def test_case3_panel_registered():
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("S_kv") == 8_192
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
@@ -423,7 +423,7 @@ def _run_case1_smoke(*, S_kv: int):
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_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_run_decode_panel_long_ctx_cube_sp_pe_tp,
)
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
@@ -454,7 +454,7 @@ def test_case1_panel_registered():
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_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_PANEL_DISPATCH,
_PANELS,
)
@@ -467,7 +467,7 @@ def test_case1_panel_registered():
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("S_kv") == 8_192
assert params.get("d_head") == 128
assert params.get("h_q") == 8
assert params.get("h_kv") == 1
@@ -548,7 +548,7 @@ def test_panel_metrics_helpers_present_and_correct():
Helpers exercised via the Case 2 smoke runner's op_log to avoid
paying the 128K-config ``_run_panel`` runtime.
"""
from kernbench.benches.milestone_gqa_decode_long_ctx_4cases import (
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
_end_to_end_ns,
_engine_occupancy_ns,
)
@@ -579,7 +579,7 @@ def test_run_panel_returns_latency_and_engine_occupancy(monkeypatch):
Uses ``monkeypatch.setitem`` on ``_PANEL_DISPATCH`` to lower S_kv to
``_SMOKE_S_KV`` just for this test — avoids the 131K runtime.
"""
import kernbench.benches.milestone_gqa_decode_long_ctx_4cases as mod
import kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases as mod
orig_kind, orig_params = mod._PANEL_DISPATCH[_CASE2_PANEL]
fast_params = {**orig_params, "S_kv": _SMOKE_S_KV}