Files
kernbench2/tests/attention/test_gqa_decode_mc.py
T
mukesh d282144339 gqa: ADR-0060/0062/0063/0064 unified GQA kernels + CPU cost model
Land the new GQA fused-attention kernels (ADR-0060) for prefill/decode
across long and short context, the TL discipline primitives they depend
on (ADR-0062 lazy load, ADR-0063 scratch_scope + copy_to), and the
per-op-type CPU issue cost model (ADR-0064). Remove the pre-ADR-0060
mesh-attention baseline now that the unified kernels supersede it.

ADR-0060 (long context)
- _gqa_decode.py: M-fold + 2-level chain reduce-to-root (Level-2
  intra-CUBE row-then-col + Level-1 inter-CUBE) — root-only output.
- _gqa_prefill.py: head-parallel + Ring KV rotation around C CUBEs,
  online-softmax merge per ring step, per-CUBE distributed output.
- Each merge stage wraps in scratch_scope() and persists running
  (m, l, O) via copy_to() to lift the 1 MiB scratch ceiling.

ADR-0060 §B.split.2 (short context, kv_per_cube in {1,2,4,8})
- _gqa_decode_short.py / _gqa_prefill_short.py: no cube-SP; each CUBE
  owns whole KV heads; PE-parallel heads with intra-group chain
  reduce. Prefill has no Ring KV (each head fully resident).

ADR-0062 (lazy tl.load): future-bearing TensorHandle, auto-wait at
first consuming op (dot/MATH/store/send/copy_to/composite).

ADR-0063 (tl.scratch_scope + tl.copy_to): scoped per-tile arena with
copy_to writeback primitive for persistent running state.

ADR-0064 (CPU issue cost model)
- common/cpu_issue_cost.py: per-op-type table (composite=40 ns,
  primitives=5 ns); ratios are load-bearing per D1.
- TLContext: issue_cost_table param; _emit_dispatch_overhead(kind)
  consults table with dispatch_cycles fallback (ADR-0046 §D6
  back-compat).
- Live PE_CPU paths (greenlet + legacy) construct TLContext with
  DEFAULT_CPU_ISSUE_COST so saturation lever (ADR-0060 §1) is
  measurable end-to-end.

P7 headline bench: milestone-gqa-headline writes per-panel
op_log_summary to 1H_milestone_output/gqa_headline/sweep.json. No
figure renderers yet (deferred).

Removals (pre-ADR-0060 baseline now superseded):
- benches: _attention_mesh_kv.py, _attention_mesh_mlo.py,
  _attention_mesh_mlo_2d.py, milestone_gqa_llama70b.py
- tests: test_attention_*, test_mesh_*, test_milestone_gqa_llama70b
- topology: llama70b_4sip.yaml (only consumer was the deleted diag)
- artifacts: 1H_milestone_output/gqa/ (sweep.json + 5 PNGs)
- tests/gqa/ plot helper + test (broken on Windows Tcl/Tkinter)
- ADR-0060/0061 references to deleted file paths cleaned up
  (EN + KO kept in sync).

Tests: 124/124 focused regression green (attention + Phase E + TL
discipline + triton_emu + pe_components). Full regression: 764 pass,
2 pre-existing test_bench_registry failures (stale EXPECTED_NAMES
across multiple benches, not introduced here).

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

183 lines
6.8 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 P2b GQA decode multi-cube SP (both Level-2 + Level-1).
P2b extends P2a to multiple CUBEs in one CUBE Group. The kernel uses the
canonical full SFR install (``configure_sfr_intercube_multisip``) which
provides disjoint direction namespaces:
- ``intra_E / intra_W / intra_N / intra_S`` — PE↔PE within a CUBE
(logical 2×4 grid, no wrap)
- ``E / W / N / S`` — CUBE↔CUBE inter-CUBE
(mesh, no wrap)
Reduce pattern (chain reduce-to-root, ADR-0060 §A.2 spirit, §4 chain
deviation noted):
Level-2 (intra-CUBE, 2×4 grid):
row-then-col chain — each row reduces leftward along ``intra_W`` to
its col-0 PE, then PE 4 sends to PE 0 along ``intra_N``. 7 chain
steps per CUBE × 3 handles each = 21 ``ipcq_copy`` per CUBE.
Level-1 (inter-CUBE):
only PE 0 of each CUBE participates. Chain leftward along ``W``.
(C-1) chain steps × 3 handles each.
Final store at PE 0 of CUBE 0 only.
Phase 1 (this commit): tests only — production code lands in Phase 2.
Phase 2 also updates ``test_gqa_decode.py`` (add C=1) and
``test_gqa_decode_sp.py`` (switch SFR + add C=1).
"""
from __future__ import annotations
from pathlib import Path
from kernbench.benches._gqa_decode import gqa_decode_kernel # noqa: F401
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"
T_Q = 1
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 _run_decode_mc(*, h_q: int, h_kv: int, C: int, P: int, S_kv: int):
"""Multi-CUBE SP decode: C cubes × P PEs each share the work."""
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)
dp_kv = DPPolicy(cube="row_wise", pe="row_wise",
num_cubes=C, num_pes=P)
# Total KV split across C×P ranks; each rank sees S_kv/(C·P) rows.
q = ctx.zeros((T_Q, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full,
name=f"q_h{h_q}_kv{h_kv}_c{C}_p{P}")
k = ctx.zeros((S_kv, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_kv,
name=f"k_h{h_q}_kv{h_kv}_c{C}_p{P}")
v = ctx.zeros((S_kv, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_kv,
name=f"v_h{h_q}_kv{h_kv}_c{C}_p{P}")
o = ctx.empty((T_Q, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full,
name=f"o_h{h_q}_kv{h_kv}_c{C}_p{P}")
ctx.launch(
f"gqa_decode_mc_h{h_q}_kv{h_kv}_c{C}_p{P}",
gqa_decode_kernel,
q, k, v, o,
T_Q, S_kv, h_q, h_kv, D_HEAD,
C, P,
_auto_dim_remap=False,
)
return run_bench(
topology=topo,
bench_fn=_bench_fn,
device=resolve_device(None),
engine_factory=_engine_factory,
)
def _count(op_log, name: str) -> int:
return sum(1 for r in op_log if r.op_name == name)
# ── Degenerate C=1 P=1 ────────────────────────────────────────────────
def test_mc_c_one_p_one_degenerate():
"""C=1, P=1: single rank, no SP. No IPCQ traffic; one dma_write."""
result = _run_decode_mc(h_q=1, h_kv=1, C=1, P=1, S_kv=16)
assert result.completion.ok, (
f"C=1 P=1 degenerate failed: {result.completion}"
)
assert _count(result.engine.op_log, "ipcq_copy") == 0
assert _count(result.engine.op_log, "dma_write") == 1
# ── Intra-CUBE only (single CUBE, P=8 on 2×4 grid) ────────────────────
def test_mc_c_one_p_eight_intracube_grid():
"""C=1, P=8: intra-cube row-then-col chain on the 2×4 grid.
7 chain steps × 3 handles (m, , O) = 21 ipcq_copy."""
result = _run_decode_mc(h_q=1, h_kv=1, C=1, P=8, S_kv=64)
assert result.completion.ok, (
f"C=1 P=8 intra-cube failed: {result.completion}"
)
assert _count(result.engine.op_log, "dma_write") == 1
n_copy = _count(result.engine.op_log, "ipcq_copy")
assert n_copy == 21, (
f"C=1 P=8: expected 21 ipcq_copy (7 chain × 3 handles); got {n_copy}"
)
# ── Multi-CUBE root-only write ────────────────────────────────────────
def test_mc_c_two_p_eight_root_only_writes_o():
"""C=2, P=8: 16 ranks total. Only PE 0 of CUBE 0 writes O."""
result = _run_decode_mc(h_q=1, h_kv=1, C=2, P=8, S_kv=128)
assert result.completion.ok, (
f"C=2 P=8 multi-cube failed: {result.completion}"
)
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 1, (
f"root-only write must hold for C=2 P=8 (16 ranks); got {n_writes}"
)
# ── Multi-CUBE total IPCQ chain count ─────────────────────────────────
def test_mc_c_two_p_eight_total_ipcq_count():
"""C=2, P=8: 21 intra-cube ipcq_copy per CUBE × 2 CUBEs + 3 inter-cube
chain ipcq_copy (C-1=1 step × 3 handles) = 45 total."""
result = _run_decode_mc(h_q=1, h_kv=1, C=2, P=8, S_kv=128)
assert result.completion.ok, (
f"C=2 P=8 multi-cube failed: {result.completion}"
)
n_copy = _count(result.engine.op_log, "ipcq_copy")
expected = 21 * 2 + (2 - 1) * 3
assert n_copy == expected, (
f"C=2 P=8: expected {expected} ipcq_copy (21 intra × 2 CUBEs + 3 "
f"inter); got {n_copy}"
)
# ── Real GQA × multi-CUBE SP combined ─────────────────────────────────
def test_mc_real_gqa_c_two_p_eight():
"""Headline: real GQA (h_q = G·h_kv with G=8) AND multi-CUBE SP
(C=2, P=8) together — the case the original baseline can express
neither part of."""
result = _run_decode_mc(h_q=8, h_kv=1, C=2, P=8, S_kv=128)
assert result.completion.ok, (
f"real GQA + multi-CUBE SP combined failed: {result.completion}"
)
n_writes = _count(result.engine.op_log, "dma_write")
assert n_writes == 1, (
f"root-only write must hold under M-fold + multi-CUBE; "
f"got {n_writes}"
)