gqa(decode-4cases): wire latency + engine occupancy into sweep.json; add comparative plot script (5C.F)
Bench changes:
- new _end_to_end_ns(op_log) and _engine_occupancy_ns(op_log) helpers
in milestone_gqa_decode_long_ctx_4cases.py (mirror paper_gqa_latency.py)
- _run_panel return dict now carries latency_ns + engine_occupancy_ns
alongside op_log_summary, so sweep.json is the single source of truth
for the comparative figures
Plot script:
- new scripts/paper/paper_plot_gqa_decode_long_ctx_4cases.py reads
sweep.json and emits 3 PNGs to docs/report/1H-codesign-paper/figures/:
gqa_decode_long_ctx_4cases_latency.png (end-to-end latency / case)
gqa_decode_long_ctx_4cases_traffic.png (ipcq/dma op counts / case)
gqa_decode_long_ctx_4cases_memory.png (KV bytes per cube / case)
Test changes:
- 2 new tests verifying the helpers + _run_panel dict shape
- lower smoke S_kv from 8192 -> 2048 (4x faster Case 2; assertions
are S_kv-independent; one fold-loop iteration preserved)
18/18 tests pass in ~4 min.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -31,6 +31,13 @@ _CASE2_PANEL = "single_kv_group_decode_long_ctx_gqa_cube_repl_pe_tp"
|
||||
_CASE1_PANEL = "single_kv_group_decode_long_ctx_gqa_cube_sp_pe_tp"
|
||||
_CUBE_RE = re.compile(r"\bcube(\d+)\b")
|
||||
|
||||
# Smoke S_kv: small enough that Case 2 finishes quickly (Case 2 is single-PE,
|
||||
# so tile count drives its runtime), big enough that the per-tile fold loop
|
||||
# is exercised once (>= 2 * TILE_S_KV = 2048). Assertions (ipcq counts,
|
||||
# root cube ids, latency > 0) are S_kv-independent. The headline 128K runs
|
||||
# come from the bench, not pytest.
|
||||
_SMOKE_S_KV = 2048
|
||||
|
||||
|
||||
def _engine_factory(t, d):
|
||||
return GraphEngine(getattr(t, "topology_obj", t), enable_data=True)
|
||||
@@ -57,7 +64,7 @@ def _dma_write_cubes(op_log) -> list[int]:
|
||||
def _run_case4_smoke(*, S_kv: int):
|
||||
"""Drive the Case 4 decode panel via the case-specific runner.
|
||||
|
||||
Uses ``S_kv=8192`` (smoke) to keep test time bounded; the headline
|
||||
Uses ``S_kv=_SMOKE_S_KV`` (smoke) to keep test time bounded; the headline
|
||||
``S_kv=128K`` runs come from ``kernbench run --bench
|
||||
milestone-gqa-decode-long-ctx-4cases``, not pytest.
|
||||
"""
|
||||
@@ -124,7 +131,7 @@ def test_case4_panel_registered():
|
||||
|
||||
def test_case4_runner_smoke():
|
||||
"""Case 4 runner drives the new kernel to completion at smoke S_kv."""
|
||||
result = _run_case4_smoke(S_kv=8192)
|
||||
result = _run_case4_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok, (
|
||||
f"Case 4 decode smoke at C=8 P=8 must complete; "
|
||||
f"got {result.completion}"
|
||||
@@ -139,7 +146,7 @@ def test_case4_root_at_center_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)
|
||||
result = _run_case4_smoke(S_kv=_SMOKE_S_KV)
|
||||
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"
|
||||
@@ -171,7 +178,7 @@ def test_case4_two_level_ar_ipcq_pattern():
|
||||
|
||||
Grand total: 168 + 21 = 189
|
||||
"""
|
||||
result = _run_case4_smoke(S_kv=8192)
|
||||
result = _run_case4_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
assert n_copy == 189, (
|
||||
@@ -246,7 +253,7 @@ def test_case2_panel_registered():
|
||||
|
||||
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)
|
||||
result = _run_case2_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok, (
|
||||
f"Case 2 decode smoke at C=8 P=8 must complete; "
|
||||
f"got {result.completion}"
|
||||
@@ -260,7 +267,7 @@ 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)
|
||||
result = _run_case2_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
assert n_copy == 0, (
|
||||
@@ -275,7 +282,7 @@ 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)
|
||||
result = _run_case2_smoke(S_kv=_SMOKE_S_KV)
|
||||
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"
|
||||
@@ -350,7 +357,7 @@ def test_case3_panel_registered():
|
||||
|
||||
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)
|
||||
result = _run_case3_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok, (
|
||||
f"Case 3 decode smoke at C=8 P=8 must complete; "
|
||||
f"got {result.completion}"
|
||||
@@ -376,7 +383,7 @@ def test_case3_intra_cube_ar_only_ipcq():
|
||||
|
||||
Grand total: 168.
|
||||
"""
|
||||
result = _run_case3_smoke(S_kv=8192)
|
||||
result = _run_case3_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
assert n_copy == 168, (
|
||||
@@ -393,7 +400,7 @@ def test_case3_single_dma_write_at_cube_0():
|
||||
the designated writer (cube 0, PE 0) stores O to avoid 8 redundant
|
||||
DMAs.
|
||||
"""
|
||||
result = _run_case3_smoke(S_kv=8192)
|
||||
result = _run_case3_smoke(S_kv=_SMOKE_S_KV)
|
||||
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"
|
||||
@@ -471,7 +478,7 @@ def test_case1_panel_registered():
|
||||
|
||||
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)
|
||||
result = _run_case1_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok, (
|
||||
f"Case 1 decode smoke at C=8 P=8 must complete; "
|
||||
f"got {result.completion}"
|
||||
@@ -496,7 +503,7 @@ def test_case1_inter_cube_lrab_only_ipcq():
|
||||
|
||||
Grand total: 21.
|
||||
"""
|
||||
result = _run_case1_smoke(S_kv=8192)
|
||||
result = _run_case1_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
assert n_copy == 21, (
|
||||
@@ -513,7 +520,7 @@ def test_case1_root_at_center_cube_6():
|
||||
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)
|
||||
result = _run_case1_smoke(S_kv=_SMOKE_S_KV)
|
||||
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"
|
||||
@@ -522,3 +529,73 @@ def test_case1_root_at_center_cube_6():
|
||||
f"Case 1 root must be the lrab center cube 6; "
|
||||
f"got cubes={sorted(distinct)}"
|
||||
)
|
||||
|
||||
|
||||
# ── 5C.F panel-metrics helpers + _run_panel wiring ──────────────────
|
||||
|
||||
|
||||
_EXPECTED_ENGINES = {
|
||||
"pe_gemm", "pe_math", "pe_dma", "pe_fetch_store", "pe_ipcq", "pe_cpu",
|
||||
}
|
||||
|
||||
|
||||
def test_panel_metrics_helpers_present_and_correct():
|
||||
"""The bench file must expose ``_end_to_end_ns`` and
|
||||
``_engine_occupancy_ns`` so ``_run_panel`` can carry latency and
|
||||
per-engine occupancy into sweep.json (required for the 5C.F
|
||||
comparative plot script).
|
||||
|
||||
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 (
|
||||
_end_to_end_ns,
|
||||
_engine_occupancy_ns,
|
||||
)
|
||||
result = _run_case2_smoke(S_kv=_SMOKE_S_KV)
|
||||
assert result.completion.ok
|
||||
op_log = result.engine.op_log
|
||||
|
||||
lat = _end_to_end_ns(op_log)
|
||||
assert lat > 0, f"expected positive end-to-end latency; got {lat}"
|
||||
|
||||
occ = _engine_occupancy_ns(op_log)
|
||||
assert isinstance(occ, dict)
|
||||
assert set(occ.keys()) >= _EXPECTED_ENGINES, (
|
||||
f"engine_occupancy_ns missing required engines; "
|
||||
f"got {set(occ.keys())}"
|
||||
)
|
||||
# GEMM engine must have done some work (Case 2's local attention).
|
||||
assert occ["pe_gemm"] > 0, (
|
||||
f"expected pe_gemm occupancy > 0; got {occ['pe_gemm']}"
|
||||
)
|
||||
|
||||
|
||||
def test_run_panel_returns_latency_and_engine_occupancy(monkeypatch):
|
||||
"""``_run_panel`` must include ``latency_ns`` + ``engine_occupancy_ns``
|
||||
alongside ``op_log_summary`` in its returned row dict, so sweep.json
|
||||
carries them for the comparative plot script.
|
||||
|
||||
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
|
||||
|
||||
orig_kind, orig_params = mod._PANEL_DISPATCH[_CASE2_PANEL]
|
||||
fast_params = {**orig_params, "S_kv": _SMOKE_S_KV}
|
||||
monkeypatch.setitem(
|
||||
mod._PANEL_DISPATCH, _CASE2_PANEL, (orig_kind, fast_params),
|
||||
)
|
||||
|
||||
row = mod._run_panel(_CASE2_PANEL, str(TOPOLOGY_DEFAULT))
|
||||
assert "op_log_summary" in row # existing key — sanity guard
|
||||
assert "latency_ns" in row, (
|
||||
f"_run_panel row missing latency_ns; keys={sorted(row.keys())}"
|
||||
)
|
||||
assert row["latency_ns"] > 0
|
||||
assert "engine_occupancy_ns" in row, (
|
||||
f"_run_panel row missing engine_occupancy_ns; "
|
||||
f"keys={sorted(row.keys())}"
|
||||
)
|
||||
assert isinstance(row["engine_occupancy_ns"], dict)
|
||||
assert set(row["engine_occupancy_ns"].keys()) >= _EXPECTED_ENGINES
|
||||
|
||||
Reference in New Issue
Block a user