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:
@@ -151,7 +151,7 @@ def test_k_before_v_in_opt2_plan():
|
||||
def test_opt2_bench_completes_oplog_mode():
|
||||
from pathlib import Path
|
||||
|
||||
from kernbench.benches._gqa_attention_decode_opt2 import ( # noqa: F401
|
||||
from kernbench.benches.gqa_helpers.shared._gqa_attention_decode_opt2 import ( # noqa: F401
|
||||
gqa_attention_decode_opt2_kernel,
|
||||
)
|
||||
from kernbench.policy.placement.dp import DPPolicy
|
||||
|
||||
@@ -1,118 +0,0 @@
|
||||
"""Phase 1 spec test for P6a GQA prefill kernel (head-parallel, C=1 baseline).
|
||||
|
||||
P6a introduces ``_gqa_attention_prefill_long.py`` with the head-parallel structure (one
|
||||
Q head per CUBE, per-CUBE distributed output, no reduce). C=1 is the
|
||||
degenerate case — no Ring KV, no IPCQ traffic. Validates kernel
|
||||
structure and T_q > 1 attention.
|
||||
|
||||
P6b (deferred) adds the Ring KV rotation for C > 1, which needs either
|
||||
a new SFR install function (intra_* + wrapped E/W at CUBE level) or a
|
||||
topology-specific config — separate design call.
|
||||
|
||||
The prefill kernel differs from decode (P1a/P2a/P2b) in three ways
|
||||
(ADR-0060 §5.5 / TL;DR):
|
||||
1. Q has T_q > 1 rows (not just decode's single timestep).
|
||||
2. Head-parallel placement: each CUBE owns ONE Q head — no M-fold.
|
||||
3. Each CUBE writes its own head's output — NO reduce.
|
||||
|
||||
Phase 1 (this commit): tests only — production code lands in Phase 2.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401
|
||||
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 _engine_factory(t, d):
|
||||
return GraphEngine(getattr(t, "topology_obj", t), enable_data=True)
|
||||
|
||||
|
||||
def _run_prefill(*, T_q: int, S_kv: int, C: int = 1):
|
||||
"""C=1 head-parallel prefill: single CUBE owns the one head + full KV."""
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
def _bench_fn(ctx):
|
||||
dp = DPPolicy(cube="replicate", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
# Q: (T_q, d_head) — one head per CUBE (head-parallel; for C=1
|
||||
# only one head total). 2D layout matches what the kernel loads.
|
||||
# P6b will use a 3D (h_q, T_q, d_head) Q with cube_row_wise
|
||||
# sharding so each CUBE owns its head.
|
||||
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp,
|
||||
name=f"q_t{T_q}_c{C}")
|
||||
# K, V: full local for C=1 (no ring). Kernel loads K as
|
||||
# (d_head, S_kv) via byte-conserving reshape.
|
||||
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp,
|
||||
name=f"k_t{T_q}_c{C}")
|
||||
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp,
|
||||
name=f"v_t{T_q}_c{C}")
|
||||
# O: (T_q, d_head) — per-CUBE distributed output.
|
||||
o = ctx.empty((T_q, D_HEAD), dtype=DTYPE, dp=dp,
|
||||
name=f"o_t{T_q}_c{C}")
|
||||
ctx.launch(
|
||||
f"gqa_prefill_p6a_t{T_q}_s{S_kv}_c{C}",
|
||||
gqa_attention_prefill_long_kernel,
|
||||
q, k, v, o,
|
||||
T_q, S_kv, D_HEAD, C,
|
||||
_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)
|
||||
|
||||
|
||||
def test_prefill_c_one_t_q_one_completes():
|
||||
"""C=1, T_q=1: smallest workload (decode-like)."""
|
||||
result = _run_prefill(T_q=1, S_kv=16, C=1)
|
||||
assert result.completion.ok, (
|
||||
f"prefill C=1 T_q=1 failed: {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
def test_prefill_c_one_t_q_four_completes():
|
||||
"""C=1, T_q=4: real prefill (Q has multiple rows) — distinguishes
|
||||
prefill from decode (T_q=1)."""
|
||||
result = _run_prefill(T_q=4, S_kv=16, C=1)
|
||||
assert result.completion.ok, (
|
||||
f"prefill C=1 T_q=4 failed: {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
def test_prefill_c_one_no_ipcq_traffic():
|
||||
"""C=1: no ring step, no IPCQ traffic."""
|
||||
result = _run_prefill(T_q=4, S_kv=16, C=1)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
assert n_copy == 0, (
|
||||
f"C=1 must have no IPCQ traffic (no ring); got {n_copy}"
|
||||
)
|
||||
|
||||
|
||||
def test_prefill_c_one_one_dma_write():
|
||||
"""C=1, one head: exactly one dma_write (per-CUBE distributed output;
|
||||
no reduce). For C > 1 in P6b this becomes dma_write_count == C."""
|
||||
result = _run_prefill(T_q=4, S_kv=16, C=1)
|
||||
assert result.completion.ok
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == 1, (
|
||||
f"C=1 prefill: expected 1 dma_write (one head per cube); "
|
||||
f"got {n_writes}"
|
||||
)
|
||||
@@ -1,154 +0,0 @@
|
||||
"""Tests for prefill_long Ring KV at C=8 over a 2×4 snake sub-mesh.
|
||||
|
||||
ADR-0060 §5.5 design target for the single-KV-group LLaMA-3.1-70B
|
||||
configuration: ``C = G = 8`` (one Q head per CUBE), Ring KV rotates
|
||||
the 8 KV slices around all 8 CUBEs.
|
||||
|
||||
Increment 1 wired ``configure_sfr_intercube_ring(submesh_shape=(2, 4))``
|
||||
to map the kernel's logical E/W to a snake/serpentine 1-hop path
|
||||
through the top 2×4 sub-mesh of the 4×4 SIP CUBE mesh. The kernel
|
||||
itself uses only logical ``dir="W"`` / ``dir="E"`` — the snake is
|
||||
transparent at kernel level.
|
||||
|
||||
This file verifies the assembly works end-to-end at C=8:
|
||||
T1 kernel completes (no scratch overflow, no deadlock)
|
||||
T2 per-CUBE head-parallel output (8 dma_writes, one per CUBE)
|
||||
T3 tile-granular ring traffic at C=8 follows the same
|
||||
``(C-1)·n_tiles·2·C`` formula as the existing tile-ring tests
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401
|
||||
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
|
||||
from kernbench.ccl.sfr_config import configure_sfr_intercube_ring
|
||||
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)
|
||||
|
||||
|
||||
def _run_prefill_c8_snake(*, T_q: int, S_kv: int):
|
||||
"""Head-parallel prefill at C=8 with snake-mapped Ring KV over the
|
||||
top 2×4 sub-mesh of the 4×4 SIP CUBE mesh."""
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
C = 8
|
||||
|
||||
def _bench_fn(ctx):
|
||||
configure_sfr_intercube_ring(
|
||||
ctx.engine, ctx.spec, _ccl_cfg(),
|
||||
submesh_shape=(2, 4),
|
||||
)
|
||||
dp_q = DPPolicy(cube="replicate", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
dp_kv = DPPolicy(cube="row_wise", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
dp_o = DPPolicy(cube="row_wise", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q,
|
||||
name=f"q_c8_snake_t{T_q}_s{S_kv}")
|
||||
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"k_c8_snake_t{T_q}_s{S_kv}")
|
||||
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"v_c8_snake_t{T_q}_s{S_kv}")
|
||||
o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o,
|
||||
name=f"o_c8_snake_t{T_q}_s{S_kv}")
|
||||
ctx.launch(
|
||||
f"gqa_prefill_long_c8_snake_t{T_q}_s{S_kv}",
|
||||
gqa_attention_prefill_long_kernel,
|
||||
q, k, v, o,
|
||||
T_q, S_kv, D_HEAD, C,
|
||||
_auto_dim_remap=False,
|
||||
)
|
||||
|
||||
return run_bench(
|
||||
topology=topo, bench_fn=_bench_fn,
|
||||
device=resolve_device(None),
|
||||
engine_factory=_engine_factory,
|
||||
)
|
||||
|
||||
|
||||
# ── T1: kernel completes end-to-end at C=8 ───────────────────────────
|
||||
|
||||
|
||||
def test_prefill_long_c8_snake_completes():
|
||||
"""ADR-0060 §5.5 design target: prefill Ring KV at C=G=8 over the
|
||||
2×4 snake sub-mesh. With zero inputs, output is trivially zero;
|
||||
this test guards that the kernel reaches completion without
|
||||
deadlock, scratch overflow, or routing failure.
|
||||
|
||||
Configuration: T_q=4, S_kv=8192 → S_local=1024, n_tiles=1
|
||||
(TILE_S_KV=1024). Bounded scratch.
|
||||
"""
|
||||
result = _run_prefill_c8_snake(T_q=4, S_kv=8192)
|
||||
assert result.completion.ok, (
|
||||
f"prefill at C=8 with snake ring must complete; "
|
||||
f"got {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
# ── T2: 8 dma_writes (head-parallel, one head per CUBE) ──────────────
|
||||
|
||||
|
||||
def test_prefill_long_c8_snake_distributed_output_count():
|
||||
"""Head-parallel: each CUBE owns one Q head and writes its own
|
||||
head's output (T_q, d_head) rows. No inter-CUBE reduce on the
|
||||
output side — so ``dma_write_count == C == 8``.
|
||||
|
||||
This is the C=8 analogue of
|
||||
``test_prefill_long_tile_ring_dma_write_count`` at C=4.
|
||||
"""
|
||||
result = _run_prefill_c8_snake(T_q=4, S_kv=8192)
|
||||
assert result.completion.ok
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == 8, (
|
||||
f"head-parallel C=8: expected 8 dma_writes (one per CUBE); "
|
||||
f"got {n_writes}"
|
||||
)
|
||||
|
||||
|
||||
# ── T3: tile-granular ring ipcq_copy count scales as (C-1)·n_tiles·2·C
|
||||
|
||||
|
||||
def test_prefill_long_c8_snake_tile_ipcq_count():
|
||||
"""ADR-0060 §5.5.1 tile-granular ring: per-CUBE send count is
|
||||
``2·n_tiles·(C-1)`` (K + V per tile per ring step). Aggregated
|
||||
across C CUBEs, total ipcq_copy = ``(C-1)·n_tiles·2·C``.
|
||||
|
||||
Configuration: T_q=4, S_kv=8192, C=8 → S_local=1024, n_tiles=1
|
||||
(TILE_S_KV=1024). Expected total = ``7·1·2·8 = 112``.
|
||||
|
||||
Verifies that the snake-mapped 1-hop physical links carry the
|
||||
same logical ring traffic as the existing C=4 1D-row ring tests.
|
||||
"""
|
||||
C = 8
|
||||
n_tiles = 1 # S_local=1024 / TILE_S_KV=1024
|
||||
result = _run_prefill_c8_snake(T_q=4, S_kv=8192)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
expected = (C - 1) * n_tiles * 2 * C
|
||||
assert n_copy == expected, (
|
||||
f"tile-granular ring at C=8: expected {expected} ipcq_copy "
|
||||
f"((C-1)·n_tiles·2·C = {C - 1}·{n_tiles}·2·{C}); got {n_copy}"
|
||||
)
|
||||
@@ -1,233 +0,0 @@
|
||||
"""Tests for intra-CUBE PE-SP in prefill_long (query-axis split).
|
||||
|
||||
ADR-0060 §5.5 last bullet + §B-item-3: split the head's query rows
|
||||
``[T_q, d_head]`` across the ``P`` PEs of each CUBE so all P PEs work
|
||||
in parallel. Output rows are disjoint across PEs ⇒ no intra-CUBE
|
||||
reduce needed.
|
||||
|
||||
Same-lane SFR wiring: PE ``i`` in CUBE A has its own E/W ring link to
|
||||
PE ``i`` in CUBE B (the snake's prev/next). All P PEs of a CUBE see
|
||||
the same K, V (HBM-resident, ``pe="replicate"``) ⇒ P parallel rings
|
||||
run in lockstep, each PE rotating its own K/V copies via its own IPCQ
|
||||
channels.
|
||||
|
||||
Activation contract:
|
||||
- ``P == 1`` (default; omitted from launch args) → existing
|
||||
PE-0-only behavior. Byte-for-byte unchanged.
|
||||
- ``P > 1`` → all P PEs active; each handles ``T_q // P`` query
|
||||
rows. Requires ``T_q % P == 0`` (degenerate T_q < P is rejected
|
||||
— caller must use ``P=1`` for that workload, ADR-0060 §B-item-3).
|
||||
|
||||
Phase 1: tests only — production code lands in Phase 2.
|
||||
T1, T2, T3, T5 fail today (TypeError: kernel signature has no P).
|
||||
T4 passes today as the backward-compat anchor.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401
|
||||
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
|
||||
from kernbench.ccl.sfr_config import configure_sfr_intercube_ring
|
||||
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)
|
||||
|
||||
|
||||
def _run_prefill_pe_sp(
|
||||
*, T_q: int, S_kv: int, C: int, P: int | None,
|
||||
snake: bool,
|
||||
):
|
||||
"""Drive a prefill_long launch with optional intra-CUBE PE-SP.
|
||||
|
||||
``P=None`` → omit P from the launch args (exercises the kernel's
|
||||
default behaviour).
|
||||
``snake=True`` → install the 2×4 snake ring SFR (Increment 1);
|
||||
else install the 1D-row ring at ``ring_size=C``.
|
||||
"""
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
def _bench_fn(ctx):
|
||||
if snake:
|
||||
configure_sfr_intercube_ring(
|
||||
ctx.engine, ctx.spec, _ccl_cfg(),
|
||||
submesh_shape=(2, 4),
|
||||
)
|
||||
else:
|
||||
configure_sfr_intercube_ring(
|
||||
ctx.engine, ctx.spec, _ccl_cfg(),
|
||||
ring_size=C,
|
||||
)
|
||||
num_pes = P if (P is not None and P > 1) else 1
|
||||
# When PE-SP is active, Q and O are split row-wise across PEs;
|
||||
# K, V remain replicated within a CUBE.
|
||||
q_pe = "row_wise" if num_pes > 1 else "replicate"
|
||||
o_pe = "row_wise" if num_pes > 1 else "replicate"
|
||||
dp_q = DPPolicy(cube="replicate", pe=q_pe,
|
||||
num_cubes=C, num_pes=num_pes)
|
||||
dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate",
|
||||
pe="replicate", num_cubes=C, num_pes=num_pes)
|
||||
dp_o = DPPolicy(cube="row_wise" if C > 1 else "replicate",
|
||||
pe=o_pe, num_cubes=C, num_pes=num_pes)
|
||||
suffix = f"c{C}_p{P}_t{T_q}_s{S_kv}"
|
||||
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q,
|
||||
name=f"q_pe_sp_{suffix}")
|
||||
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"k_pe_sp_{suffix}")
|
||||
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"v_pe_sp_{suffix}")
|
||||
o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o,
|
||||
name=f"o_pe_sp_{suffix}")
|
||||
launch_args = [q, k, v, o, T_q, S_kv, D_HEAD, C]
|
||||
if P is not None:
|
||||
launch_args.append(P)
|
||||
ctx.launch(
|
||||
f"gqa_prefill_long_pe_sp_{suffix}",
|
||||
gqa_attention_prefill_long_kernel,
|
||||
*launch_args,
|
||||
_auto_dim_remap=False,
|
||||
)
|
||||
|
||||
return run_bench(
|
||||
topology=topo, bench_fn=_bench_fn,
|
||||
device=resolve_device(None),
|
||||
engine_factory=_engine_factory,
|
||||
)
|
||||
|
||||
|
||||
# ── T1: C=8 P=8 completes end-to-end ─────────────────────────────────
|
||||
|
||||
|
||||
def test_prefill_c8_p8_pe_sp_completes():
|
||||
"""ADR-0060 §5.5 + §B-item-3 design target: 64 ranks active
|
||||
(8 CUBEs × 8 PEs), query-axis split across PEs, snake-mapped
|
||||
Ring KV. Guards no scratch overflow, no deadlock across the
|
||||
P parallel same-lane rings.
|
||||
"""
|
||||
result = _run_prefill_pe_sp(
|
||||
T_q=8, S_kv=8192, C=8, P=8, snake=True,
|
||||
)
|
||||
assert result.completion.ok, (
|
||||
f"prefill at C=8, P=8 (PE-SP) must complete; "
|
||||
f"got {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
# ── T2: 64 dma_writes (one per PE, disjoint query rows) ──────────────
|
||||
|
||||
|
||||
def test_prefill_c8_p8_pe_sp_64_dma_writes():
|
||||
"""With query-axis split, each PE owns ``T_q/P = 1`` query row
|
||||
and stores its own ``(1, d_head)`` slice. Across all 64 ranks
|
||||
(8 CUBEs × 8 PEs), ``dma_write_count == 64``.
|
||||
|
||||
This is the structural signal that PE-SP is actually wired:
|
||||
PE-0-only would give 8 dma_writes (one per CUBE).
|
||||
"""
|
||||
result = _run_prefill_pe_sp(
|
||||
T_q=8, S_kv=8192, C=8, P=8, snake=True,
|
||||
)
|
||||
assert result.completion.ok
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == 64, (
|
||||
f"PE-SP at C=8, P=8: expected 64 dma_writes "
|
||||
f"(8 cubes × 8 PEs, each writes its T_q/P=1 query row); "
|
||||
f"got {n_writes}"
|
||||
)
|
||||
|
||||
|
||||
# ── T3: P parallel rings — ipcq_copy scales by P ─────────────────────
|
||||
|
||||
|
||||
def test_prefill_c8_p8_pe_sp_ring_ipcq_count():
|
||||
"""Per-PE same-lane rings: each PE runs its own ring traffic via
|
||||
its own IPCQ channels. Total inter-CUBE ipcq_copy:
|
||||
``(C-1) · n_tiles · 2 · C · P`` (= existing C=8 formula × P).
|
||||
|
||||
Configuration: T_q=8, S_kv=8192, C=8 → S_local=1024, n_tiles=1
|
||||
(TILE_S_KV=1024). Expected = ``7·1·2·8·8`` = **896**.
|
||||
"""
|
||||
C = 8
|
||||
P = 8
|
||||
n_tiles = 1
|
||||
result = _run_prefill_pe_sp(
|
||||
T_q=8, S_kv=8192, C=C, P=P, snake=True,
|
||||
)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
expected = (C - 1) * n_tiles * 2 * C * P
|
||||
assert n_copy == expected, (
|
||||
f"PE-SP parallel rings at C=8, P=8: expected {expected} "
|
||||
f"ipcq_copy ((C-1)·n_tiles·2·C·P = "
|
||||
f"{C - 1}·{n_tiles}·2·{C}·{P}); got {n_copy}"
|
||||
)
|
||||
|
||||
|
||||
# ── T4: backward-compat — P omitted → existing PE-0-only behaviour ───
|
||||
|
||||
|
||||
def test_prefill_p_default_1_backward_compat():
|
||||
"""When ``P`` is omitted from the launch args, the kernel must
|
||||
behave identically to today: only PE 0 of each CUBE participates;
|
||||
one head per CUBE; one dma_write per CUBE.
|
||||
|
||||
At C=4 this gives 4 dma_writes (matches the existing
|
||||
``test_prefill_long_tile_ring_dma_write_count``). Phase 2 must
|
||||
not regress this path.
|
||||
|
||||
Passes today AND after Phase 2.
|
||||
"""
|
||||
C = 4
|
||||
result = _run_prefill_pe_sp(
|
||||
T_q=4, S_kv=8192, C=C, P=None, snake=False,
|
||||
)
|
||||
assert result.completion.ok, (
|
||||
f"prefill at C=4 with default P (PE-0-only) must complete; "
|
||||
f"got {result.completion}"
|
||||
)
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == C, (
|
||||
f"default P=1 (PE-0-only): expected {C} dma_writes "
|
||||
f"(one per CUBE); got {n_writes}"
|
||||
)
|
||||
|
||||
|
||||
# ── T5: validation — T_q must be divisible by P when P > 1 ───────────
|
||||
|
||||
|
||||
def test_prefill_pe_sp_rejects_non_divisible_t_q():
|
||||
"""ADR-0060 §B-item-3 fallback (KV-block split + intra-CUBE
|
||||
reduce for T_q < P) is deferred. Callers must request ``P=1`` for
|
||||
workloads where T_q < P or T_q % P != 0.
|
||||
|
||||
C=4, P=8, T_q=4 violates ``T_q % P == 0`` (and also T_q < P).
|
||||
The kernel must raise ValueError with a clear error message
|
||||
rather than silently producing wrong results.
|
||||
"""
|
||||
with pytest.raises((ValueError, AssertionError), match=r"T_q"):
|
||||
_run_prefill_pe_sp(
|
||||
T_q=4, S_kv=8192, C=4, P=8, snake=False,
|
||||
)
|
||||
@@ -1,162 +0,0 @@
|
||||
"""Phase 1 spec test for P3c: tile-granular Ring KV in prefill_long
|
||||
(ADR-0060 §5.5 amendment + ADR-0063 §A.2).
|
||||
|
||||
Today's prefill_long ring sends and receives full ``(d_head, S_local)``
|
||||
KV slices per step. Step 0's local-attention intermediates also live
|
||||
outside any ``scratch_scope`` (they're loaded as full-slice ``Kc``,
|
||||
``Vc`` and feed ``scores``, ``exp_scores`` as persistent allocations).
|
||||
At larger ``S_local`` both the step-0 leak and the ring step's
|
||||
in-scope intermediates grow linearly with ``S_local``, and at
|
||||
``S_local = 32K`` (S_kv=128K, C=4, T_q=4) the peak overflows the
|
||||
1 MiB pool.
|
||||
|
||||
P3c converts the ring to **tile-granular**: a nested loop
|
||||
``for k in range(C): for t in range(n_tiles): ...`` where each
|
||||
iteration sends/recvs one ``(d_head, TILE_S_KV)`` tile (and its V
|
||||
counterpart). The persistent state shrinks to ``(m, ℓ, O)`` only
|
||||
(~1 KB); per-tile in-scope scratch is bounded by
|
||||
``TILE_S_KV`` regardless of ``S_local``. Ceiling lifted.
|
||||
|
||||
Trade-off: the per-CUBE send count grows from ``2·(C-1)`` to
|
||||
``2·n_tiles·(C-1)``. At ``n_tiles=1`` (small ``S_local``) the count is
|
||||
unchanged, so the existing ``test_prefill_ring_c_*`` tests at
|
||||
``S_kv ∈ {16, 32}`` still pass.
|
||||
|
||||
Phase 1 (this commit): tests only — production code lands in Phase 2.
|
||||
T1 fails today with a ``TLContext scratch overflow``; T2 fails today
|
||||
with the slice-granular ipcq_copy count; T5 passes today and is a
|
||||
regression guard for the per-CUBE output write count.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401
|
||||
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
|
||||
from kernbench.ccl.sfr_config import configure_sfr_intercube_ring
|
||||
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)
|
||||
|
||||
|
||||
def _run_prefill_ring(*, T_q: int, S_kv: int, C: int):
|
||||
"""Head-parallel prefill with Ring KV across C CUBEs."""
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
def _bench_fn(ctx):
|
||||
configure_sfr_intercube_ring(
|
||||
ctx.engine, ctx.spec, _ccl_cfg(), ring_size=C,
|
||||
)
|
||||
dp_q = DPPolicy(cube="replicate", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate",
|
||||
pe="replicate", num_cubes=C, num_pes=1)
|
||||
dp_o = DPPolicy(cube="row_wise" if C > 1 else "replicate",
|
||||
pe="replicate", num_cubes=C, num_pes=1)
|
||||
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q,
|
||||
name=f"q_tlr_t{T_q}_c{C}_s{S_kv}")
|
||||
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"k_tlr_t{T_q}_c{C}_s{S_kv}")
|
||||
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"v_tlr_t{T_q}_c{C}_s{S_kv}")
|
||||
o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o,
|
||||
name=f"o_tlr_t{T_q}_c{C}_s{S_kv}")
|
||||
ctx.launch(
|
||||
f"gqa_prefill_long_tile_ring_t{T_q}_c{C}_s{S_kv}",
|
||||
gqa_attention_prefill_long_kernel,
|
||||
q, k, v, o,
|
||||
T_q, S_kv, D_HEAD, C,
|
||||
_auto_dim_remap=False,
|
||||
)
|
||||
|
||||
return run_bench(
|
||||
topology=topo, bench_fn=_bench_fn,
|
||||
device=resolve_device(None),
|
||||
engine_factory=_engine_factory,
|
||||
)
|
||||
|
||||
|
||||
# ── T1: 128K ceiling lift ────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_prefill_long_context_128k_completes():
|
||||
"""ADR-0063 §A.2 headline ceiling lift. At S_kv=128K with C=4,
|
||||
S_local=32K. Today: step-0 score-stack (~768 KB persistent) + ring
|
||||
scope (~768 KB) → 1.5 MB peak → TLContext scratch overflow. After
|
||||
P3c the persistent state is just ``(m, ℓ, O)`` (≈ 1 KB) and the
|
||||
per-tile in-scope scratch is bounded by TILE_S_KV.
|
||||
"""
|
||||
result = _run_prefill_ring(T_q=4, S_kv=131_072, C=4)
|
||||
assert result.completion.ok, (
|
||||
f"prefill_long at S_kv=128K must complete after tile-granular "
|
||||
f"ring lands; got {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
# ── T2: tile-granular ipcq_copy count ────────────────────────────────
|
||||
|
||||
|
||||
def test_prefill_long_tile_granular_ipcq_count():
|
||||
"""ADR-0060 §5.5 amendment: with tile-granular sends, the per-CUBE
|
||||
send count grows from ``2·(C-1)`` to ``2·n_tiles·(C-1)``. Aggregated
|
||||
across all C CUBEs the total ipcq_copy becomes
|
||||
``2·n_tiles·(C-1)·C``.
|
||||
|
||||
Config: T_q=4, S_kv=4096, C=2 → S_local=2048, n_tiles=2 (with
|
||||
TILE_S_KV=1024). Today: slice-granular total =
|
||||
``(C-1)·2·C = 4``. After P3c: tile-granular total =
|
||||
``(C-1)·n_tiles·2·C = 8``.
|
||||
"""
|
||||
C = 2
|
||||
n_tiles = 2 # S_local=2048 / TILE_S_KV=1024
|
||||
result = _run_prefill_ring(T_q=4, S_kv=4096, C=C)
|
||||
assert result.completion.ok, (
|
||||
f"prefill_long multi-tile ring must complete; got {result.completion}"
|
||||
)
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
expected = (C - 1) * n_tiles * 2 * C
|
||||
assert n_copy == expected, (
|
||||
f"tile-granular ring: expected {expected} ipcq_copy "
|
||||
f"((C-1)·n_tiles·2·C = {C - 1}·{n_tiles}·2·{C}); got {n_copy}"
|
||||
)
|
||||
|
||||
|
||||
# ── T5: per-CUBE distributed output unchanged (regression guard) ─────
|
||||
|
||||
|
||||
def test_prefill_long_tile_ring_dma_write_count():
|
||||
"""ADR-0060 §5.5: per-CUBE distributed output must hold under the
|
||||
tile-granular ring rewrite. Each CUBE still writes its own head's
|
||||
output (no inter-CUBE reduce); dma_write_count == C.
|
||||
|
||||
Today passes; must continue to pass after P3c.
|
||||
"""
|
||||
C = 4
|
||||
result = _run_prefill_ring(T_q=4, S_kv=4096, C=C)
|
||||
assert result.completion.ok
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == C, (
|
||||
f"per-CUBE distributed output: expected {C} dma_writes (one per "
|
||||
f"CUBE); got {n_writes}"
|
||||
)
|
||||
@@ -1,161 +0,0 @@
|
||||
"""Phase 1 spec test for P6b GQA prefill Ring KV (head-parallel, C>1).
|
||||
|
||||
P6b adds the Ring KV rotation (ADR-0060 §5.5) to the prefill kernel.
|
||||
Each CUBE owns one Q head + its KV slice; over C ring steps the KV
|
||||
blocks rotate around the C CUBEs so every CUBE sees every block. The
|
||||
running ``(m, ℓ, O)`` is folded inside the loop. No reduce — each CUBE
|
||||
writes its own head's output.
|
||||
|
||||
Requires a new SFR install ``configure_sfr_intercube_ring(ring_size=C)``
|
||||
that wires:
|
||||
- ``intra_*`` : 2×4 PE grid within a CUBE (same as multisip)
|
||||
- ``E/W`` : 1D ring of CUBEs 0..ring_size-1 WITH WRAP
|
||||
(symmetric to ``configure_sfr_intracube_pe_ring``,
|
||||
applied at CUBE level)
|
||||
- ``global_*`` : SIP-level (same as multisip)
|
||||
|
||||
The kernel ring body:
|
||||
for step in range(1, C):
|
||||
tl.send(dir="W", src=Kc)
|
||||
tl.send(dir="W", src=Vc)
|
||||
Kc = tl.recv(dir="E", ...)
|
||||
Vc = tl.recv(dir="E", ...)
|
||||
... local partial + online-softmax merge into (m, ℓ, O) ...
|
||||
|
||||
Per CUBE per step: 2 sends (K, V) → 2 ``ipcq_copy``. Across all CUBEs:
|
||||
``(C-1) * 2 * C`` total ipcq_copy.
|
||||
|
||||
Restriction in P6b first cut: ``C ∈ {1, 2, 4}`` (single row of the 4×4
|
||||
cube mesh). C=8 ring would span rows — follow-on.
|
||||
|
||||
Phase 1 (this commit): tests only — production code lands in Phase 2.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401
|
||||
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
|
||||
from kernbench.ccl.sfr_config import configure_sfr_intercube_ring # noqa: F401 (Phase 2)
|
||||
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 _run_prefill_ring(*, T_q: int, S_kv: int, C: int):
|
||||
"""Head-parallel prefill with Ring KV across C CUBEs."""
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
def _bench_fn(ctx):
|
||||
# P6b: new SFR install with cube-level ring wrap.
|
||||
configure_sfr_intercube_ring(
|
||||
ctx.engine, ctx.spec, _ccl_cfg(), ring_size=C,
|
||||
)
|
||||
# Q replicated on every CUBE (zeros for testing; per-CUBE head
|
||||
# indexing is implicit). KV sequence-sharded by CUBE. O
|
||||
# distributed — each CUBE writes its slice of (T_q*C, d_head).
|
||||
dp_q = DPPolicy(cube="replicate", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
dp_kv = DPPolicy(cube="row_wise", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
dp_o = DPPolicy(cube="row_wise", pe="replicate",
|
||||
num_cubes=C, num_pes=1)
|
||||
q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q,
|
||||
name=f"q_t{T_q}_c{C}_ring")
|
||||
k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"k_t{T_q}_c{C}_ring")
|
||||
v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv,
|
||||
name=f"v_t{T_q}_c{C}_ring")
|
||||
# O: (T_q * C, d_head), each CUBE writes (T_q, d_head) at its
|
||||
# slice. dma_write_count = C (per-CUBE distributed output).
|
||||
o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o,
|
||||
name=f"o_t{T_q}_c{C}_ring")
|
||||
ctx.launch(
|
||||
f"gqa_prefill_ring_t{T_q}_s{S_kv}_c{C}",
|
||||
gqa_attention_prefill_long_kernel,
|
||||
q, k, v, o,
|
||||
T_q, S_kv, D_HEAD, C,
|
||||
_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)
|
||||
|
||||
|
||||
# ── C=2 Ring KV ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_prefill_ring_c_two_completes():
|
||||
"""C=2: 1 ring step rotates KV between the 2 CUBEs."""
|
||||
result = _run_prefill_ring(T_q=4, S_kv=16, C=2)
|
||||
assert result.completion.ok, (
|
||||
f"prefill ring C=2 failed: {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
def test_prefill_ring_c_two_distributed_output():
|
||||
"""C=2: per-CUBE distributed output, no reduce → dma_write_count == 2."""
|
||||
result = _run_prefill_ring(T_q=4, S_kv=16, C=2)
|
||||
assert result.completion.ok
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == 2, (
|
||||
f"C=2 prefill: expected 2 dma_write (one per CUBE); got {n_writes}"
|
||||
)
|
||||
|
||||
|
||||
def test_prefill_ring_c_two_ipcq_count():
|
||||
"""C=2: 1 ring step × 2 handles (K, V) × 2 CUBEs = 4 ipcq_copy."""
|
||||
result = _run_prefill_ring(T_q=4, S_kv=16, C=2)
|
||||
assert result.completion.ok
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
expected = (2 - 1) * 2 * 2
|
||||
assert n_copy == expected, (
|
||||
f"C=2 ring: expected {expected} ipcq_copy "
|
||||
f"((C-1)·2·C = 1·2·2); got {n_copy}"
|
||||
)
|
||||
|
||||
|
||||
# ── C=4 Ring KV (combined assertions) ─────────────────────────────────
|
||||
|
||||
|
||||
def test_prefill_ring_c_four_combined():
|
||||
"""C=4: 3 ring steps rotate KV around 4 CUBEs.
|
||||
Expected: completes; 4 dma_writes; (4-1)·2·4 = 24 ipcq_copy."""
|
||||
result = _run_prefill_ring(T_q=4, S_kv=32, C=4)
|
||||
assert result.completion.ok, (
|
||||
f"prefill ring C=4 failed: {result.completion}"
|
||||
)
|
||||
n_writes = _count(result.engine.op_log, "dma_write")
|
||||
assert n_writes == 4, (
|
||||
f"C=4 prefill: expected 4 dma_write (one per CUBE); got {n_writes}"
|
||||
)
|
||||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||||
expected = (4 - 1) * 2 * 4
|
||||
assert n_copy == expected, (
|
||||
f"C=4 ring: expected {expected} ipcq_copy "
|
||||
f"((C-1)·2·C = 3·2·4); got {n_copy}"
|
||||
)
|
||||
@@ -76,7 +76,7 @@ def _run_decode_short(*, h_q: int, h_kv: int, kv_per_cube: int,
|
||||
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
|
||||
from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_decode_short import gqa_attention_decode_short_kernel # Phase 2
|
||||
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
@@ -196,7 +196,7 @@ def _run_prefill_short(*, h_kv: int, kv_per_cube: int,
|
||||
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
|
||||
from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_prefill_short import gqa_attention_prefill_short_kernel # Phase 2
|
||||
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -1,89 +0,0 @@
|
||||
"""Tests for the milestone-gqa-headline bench.
|
||||
|
||||
The bench wires the new ``_gqa_attention_prefill_long`` kernel into a
|
||||
single-KV-group milestone panel (LLaMA-3.1-70B target, C=8 P=8). The
|
||||
4 legacy C=1 / C=4 panels (single_user_*, multi_user_*) were removed
|
||||
once pytest regression covered those configurations comprehensively
|
||||
and the comparative decode work moved into the
|
||||
``milestone-gqa-decode-4cases`` bench.
|
||||
|
||||
Single SIP scope (default ``topology.yaml`` 4×4 cube mesh; 4-SIP
|
||||
headline deferred per ADR-0060 §B-item-1).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import kernbench.benches.milestone_gqa_headline as bench_mod # noqa: F401 (Phase 2)
|
||||
from kernbench.benches.registry import resolve
|
||||
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
|
||||
|
||||
BENCH_NAME = "milestone-gqa-headline"
|
||||
|
||||
PANELS = (
|
||||
"single_kv_group_prefill_gqa_c8_p8",
|
||||
)
|
||||
|
||||
|
||||
def _run_validation():
|
||||
topo = resolve_topology("topology.yaml")
|
||||
return run_bench(
|
||||
topology=topo,
|
||||
bench_fn=resolve(BENCH_NAME).run,
|
||||
device=resolve_device(None),
|
||||
engine_factory=lambda t, d: GraphEngine(
|
||||
getattr(t, "topology_obj", t), enable_data=True,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _sweep_json(monkeypatch) -> dict:
|
||||
monkeypatch.setenv("GQA_HEADLINE_RUN", "1")
|
||||
out = bench_mod._OUTPUT_DIR / "sweep.json"
|
||||
if not out.exists():
|
||||
result = _run_validation()
|
||||
assert result.completion.ok, result.completion
|
||||
assert out.exists(), f"missing {out}"
|
||||
return json.loads(out.read_text())
|
||||
|
||||
|
||||
# ── Registration ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_bench_registered():
|
||||
spec = resolve(BENCH_NAME)
|
||||
assert spec.name == BENCH_NAME
|
||||
assert callable(spec.run)
|
||||
assert spec.description.strip(), "description must be non-empty"
|
||||
|
||||
|
||||
# ── Validation run completes ──────────────────────────────────────────
|
||||
|
||||
|
||||
def test_validation_run_completes_ok(monkeypatch):
|
||||
monkeypatch.setenv("GQA_HEADLINE_RUN", "1")
|
||||
result = _run_validation()
|
||||
assert result.completion.ok, (
|
||||
f"headline validation run failed: {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
# ── sweep.json shape ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_sweep_json_has_expected_panels(monkeypatch):
|
||||
data = _sweep_json(monkeypatch)
|
||||
assert set(data["panels"]) == set(PANELS), (
|
||||
f"panels mismatch: expected {set(PANELS)}, got {set(data['panels'])}"
|
||||
)
|
||||
assert len(data["rows"]) == len(PANELS)
|
||||
assert {r["panel"] for r in data["rows"]} == set(PANELS)
|
||||
|
||||
|
||||
# Per-panel architectural assertions for the surviving single-KV-group
|
||||
# panel live in tests/attention/test_milestone_gqa_single_kv_group_prefill_panel.py.
|
||||
# Decode architectural assertions live in
|
||||
# tests/attention/test_milestone_gqa_decode_4cases.py.
|
||||
@@ -1,121 +0,0 @@
|
||||
"""Tests for the single-KV-group prefill panel in milestone_gqa_headline.
|
||||
|
||||
The single-KV-group (LLaMA-3.1-70B target) prefill panel wires together
|
||||
all Increment 1–4 work in a single bench panel:
|
||||
- C=8 head-parallel + Ring KV via snake-mapped 2×4 sub-mesh (Inc 1, 3)
|
||||
- P=8 intra-CUBE PE-SP (Inc 4)
|
||||
- d_head=128 (LLaMA-3.1-70B), one-shot prefill T_q = S_kv = 1K
|
||||
(scratch-limited; LLaMA 32K headline awaits Q-axis kernel tiling)
|
||||
|
||||
This file verifies the bench-config wiring:
|
||||
T1 the new ``single_kv_group_prefill_gqa_c8_p8`` panel is registered
|
||||
with the expected dims in ``_PANELS`` + ``_PANEL_DISPATCH``.
|
||||
T2 ``_run_prefill_panel`` accepts the new ``P``, ``T_q``, ``d_head``
|
||||
kwargs and drives the prefill kernel to completion at the
|
||||
milestone-target ``(C, P) = (8, 8)``. Uses smaller T_q/S_kv to
|
||||
keep test time bounded — full 32K runs come from
|
||||
``kernbench run milestone-gqa-headline``.
|
||||
T3 Existing prefill panels still work when ``_run_prefill_panel``
|
||||
is called without the new kwargs (backward compat anchor).
|
||||
|
||||
Phase 1: tests only — production code lands in Phase 2.
|
||||
T1 fails today (panel not registered).
|
||||
T2 fails today (helper doesn't accept the new kwargs).
|
||||
T3 passes today as the backward-compat anchor.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from kernbench.benches.milestone_gqa_headline import ( # noqa: F401
|
||||
_PANEL_DISPATCH,
|
||||
_PANELS,
|
||||
_run_prefill_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"
|
||||
|
||||
|
||||
def _engine_factory(t, d):
|
||||
return GraphEngine(getattr(t, "topology_obj", t), enable_data=True)
|
||||
|
||||
|
||||
# ── T1: new panel registered ─────────────────────────────────────────
|
||||
|
||||
|
||||
def test_single_kv_group_prefill_panel_registered():
|
||||
"""The ``single_kv_group_prefill_gqa_c8_p8`` panel must be in ``_PANELS`` and
|
||||
its dispatch entry must have the expected LLaMA-3.1-70B params.
|
||||
|
||||
Headline config (scratch-limited; LLaMA 32K headline awaits
|
||||
Q-axis kernel tiling):
|
||||
C = 8 (head-parallel, snake sub-mesh)
|
||||
P = 8 (intra-CUBE PE-SP)
|
||||
T_q = 1_024 (one-shot long-context prefill)
|
||||
S_kv = 1_024
|
||||
(scratch-limited; LLaMA 32K headline awaits Q-axis kernel tiling)
|
||||
d_head = 128 (LLaMA-3.1-70B)
|
||||
"""
|
||||
panel_name = "single_kv_group_prefill_gqa_c8_p8"
|
||||
assert panel_name in _PANELS, (
|
||||
f"{panel_name!r} not in _PANELS; got {_PANELS}"
|
||||
)
|
||||
assert panel_name in _PANEL_DISPATCH, (
|
||||
f"{panel_name!r} not in _PANEL_DISPATCH"
|
||||
)
|
||||
kind, params = _PANEL_DISPATCH[panel_name]
|
||||
assert kind == "prefill", f"kind={kind!r}, expected 'prefill'"
|
||||
assert params.get("C") == 8, f"C={params.get('C')}, expected 8"
|
||||
assert params.get("P") == 8, f"P={params.get('P')}, expected 8"
|
||||
assert params.get("T_q") == 1_024, (
|
||||
f"T_q={params.get('T_q')}, expected 1_024"
|
||||
)
|
||||
assert params.get("S_kv") == 1_024, (
|
||||
f"S_kv={params.get('S_kv')}, expected 1_024"
|
||||
)
|
||||
assert params.get("d_head") == 128, (
|
||||
f"d_head={params.get('d_head')}, expected 128"
|
||||
)
|
||||
|
||||
|
||||
# ── T2: helper drives the C=8, P=8 prefill kernel to completion ──────
|
||||
|
||||
|
||||
def test_single_kv_group_prefill_panel_runner_smoke():
|
||||
"""``_run_prefill_panel`` must accept the new ``P``, ``T_q``,
|
||||
``d_head`` kwargs and successfully launch the prefill kernel at
|
||||
``(C, P) = (8, 8)`` with the snake-mapped 2×4 SFR.
|
||||
|
||||
Uses **smaller** T_q/S_kv than the headline panel so the test
|
||||
completes quickly. The headline 32K dims are exercised via
|
||||
``kernbench run milestone-gqa-headline``, not pytest.
|
||||
"""
|
||||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||||
|
||||
def _bench_fn(ctx):
|
||||
_run_prefill_panel(
|
||||
ctx,
|
||||
panel="single_kv_group_prefill_gqa_c8_p8",
|
||||
C=8, P=8,
|
||||
T_q=8, S_kv=8192, # smoke dims, not the 32K headline
|
||||
d_head=128,
|
||||
)
|
||||
|
||||
result = run_bench(
|
||||
topology=topo, bench_fn=_bench_fn,
|
||||
device=resolve_device(None),
|
||||
engine_factory=_engine_factory,
|
||||
)
|
||||
assert result.completion.ok, (
|
||||
f"single_kv_group prefill panel smoke at C=8 P=8 must complete; "
|
||||
f"got {result.completion}"
|
||||
)
|
||||
|
||||
|
||||
# Backward-compat test removed when the legacy multi_user_prefill_gqa
|
||||
# panel was dropped; signature-extension contract is now exercised
|
||||
# implicitly by the live single_kv_group_prefill_gqa_c8_p8 panel.
|
||||
Reference in New Issue
Block a user