5a76ed4f6a
Make the file + function naming symmetric:
_gqa_decode.py -> _gqa_decode_long.py
_gqa_prefill.py -> _gqa_prefill_long.py
gqa_decode_kernel -> gqa_decode_long_kernel
gqa_prefill_kernel -> gqa_prefill_long_kernel
Mirrors the existing _gqa_{decode,prefill}_short.py naming. Updates the
two imports + two call sites in milestone_gqa_headline.py and the 9
attention tests that import the kernels.
Tests: 72/72 focused regression green (tests/attention/ + Phase E + TL
discipline). milestone-gqa-headline bench passes its 7 panel/schema
assertions.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
154 lines
6.5 KiB
Python
154 lines
6.5 KiB
Python
"""Phase 1 spec test for P2a GQA decode SP (chain reduce-to-root, Level-2 only).
|
||
|
||
P2a is the first half of DDD-0060 P2: the kernel becomes multi-PE within
|
||
one CUBE and reduces to root (PE 0) using a chain over the 1D intra-cube
|
||
ring (W direction). This **replaces the baseline's bidirectional O(N)
|
||
fan-out** where every rank ends with O — ADR-0060 §A.2's headline.
|
||
|
||
Deviation from DDD-0060 §7 P2 gate: the gate text asks for
|
||
``⌈log₂ P⌉`` reduce rounds. The intra-cube SFR install
|
||
(``configure_sfr_intracube_pe_ring``) wires only a 1D E/W ring, so a
|
||
true tree would require either multi-hop forwarding or a new SFR install
|
||
(future ADR). P2a uses a **chain reduce-to-root**: ``P-1`` rounds along
|
||
W. The architectural property the ADR cares about
|
||
(root-only output vs every-rank-has-O) is preserved; the logarithmic
|
||
collective is deferred.
|
||
|
||
P2b (deferred) covers Level-1 inter-CUBE center-mesh reduce (C>1).
|
||
|
||
Phase 1 (this commit): tests only — production code lands in Phase 2.
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from pathlib import Path
|
||
|
||
from kernbench.benches._gqa_decode_long import gqa_decode_long_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_sp(*, h_q: int, h_kv: int, P: int, S_kv: int):
|
||
"""Single-CUBE SP decode: P PEs share the work along the intra-cube ring."""
|
||
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=1, num_pes=P)
|
||
dp_kv = DPPolicy(cube="replicate", pe="row_wise",
|
||
num_cubes=1, num_pes=P)
|
||
q = ctx.zeros((T_Q, h_q * D_HEAD),
|
||
dtype=DTYPE, dp=dp_full, name=f"q_h{h_q}_kv{h_kv}_p{P}")
|
||
# KV: total S_kv split across P PEs along axis 0 (row_wise sharding).
|
||
# Each PE sees (S_kv/P, h_kv·D_HEAD).
|
||
k = ctx.zeros((S_kv, h_kv * D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"k_h{h_q}_kv{h_kv}_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}_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}_p{P}")
|
||
ctx.launch(
|
||
f"gqa_decode_sp_h{h_q}_kv{h_kv}_p{P}",
|
||
gqa_decode_long_kernel,
|
||
q, k, v, o,
|
||
T_Q, S_kv, h_q, h_kv, D_HEAD,
|
||
1, P, # C=1, P=P (single-CUBE SP)
|
||
_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)
|
||
|
||
|
||
# ── Root-only write ────────────────────────────────────────────────────
|
||
|
||
|
||
def test_sp_chain_reduce_root_only_writes_o():
|
||
"""ADR-0060 §A.2: only the root rank (PE 0) writes O. Baseline today
|
||
has every rank write the full final O (bidirectional fan-out)."""
|
||
result = _run_decode_sp(h_q=1, h_kv=1, P=8, S_kv=64)
|
||
assert result.completion.ok, f"P=8 chain reduce failed: {result.completion}"
|
||
n_writes = _count(result.engine.op_log, "dma_write")
|
||
assert n_writes == 1, (
|
||
f"reduce-to-root must produce exactly 1 dma_write (PE 0); "
|
||
f"got {n_writes}"
|
||
)
|
||
|
||
|
||
# ── Chain step count ───────────────────────────────────────────────────
|
||
|
||
|
||
def test_sp_chain_reduce_p_minus_one_ipcq_pairs():
|
||
"""Chain reduce-to-root has P-1 send→recv pairs along the W chain;
|
||
each pair logs one ``ipcq_copy`` (inbound DMA, per
|
||
``milestone_gqa_llama70b._summarize_op_log``). Each chain step ships
|
||
the triplet (m, ℓ, O) → 3 handles per step → 7 steps × 3 = 21."""
|
||
result = _run_decode_sp(h_q=1, h_kv=1, P=8, S_kv=64)
|
||
assert result.completion.ok, f"P=8 chain reduce failed: {result.completion}"
|
||
n_copy = _count(result.engine.op_log, "ipcq_copy")
|
||
expected = (8 - 1) * 3
|
||
assert n_copy == expected, (
|
||
f"chain reduce: expected {expected} ipcq_copy (P-1=7 steps × "
|
||
f"3 handles m/ℓ/O); got {n_copy}"
|
||
)
|
||
|
||
|
||
# ── Real GQA × SP combined ─────────────────────────────────────────────
|
||
|
||
|
||
def test_sp_real_gqa_h_q_eight_h_kv_one_p_eight():
|
||
"""The combined unlock: real GQA (h_q=G·h_kv with G=8) AND SP
|
||
(P=8) together — neither expressible by the baseline."""
|
||
result = _run_decode_sp(h_q=8, h_kv=1, P=8, S_kv=64)
|
||
assert result.completion.ok, (
|
||
f"real GQA + SP combined run 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 too; got {n_writes}"
|
||
)
|
||
|
||
|
||
# ── Degenerate P=1 ────────────────────────────────────────────────────
|
||
|
||
|
||
def test_sp_p_one_degenerate_no_ipcq_traffic():
|
||
"""P=1: SP degenerates to a single rank. No IPCQ traffic; one dma_write."""
|
||
result = _run_decode_sp(h_q=8, h_kv=1, P=1, S_kv=16)
|
||
assert result.completion.ok, f"P=1 degenerate failed: {result.completion}"
|
||
n_send = _count(result.engine.op_log, "ipcq_send")
|
||
n_recv = _count(result.engine.op_log, "ipcq_recv")
|
||
assert n_send == 0, f"P=1 must have no ipcq_send; got {n_send}"
|
||
assert n_recv == 0, f"P=1 must have no ipcq_recv; got {n_recv}"
|
||
n_writes = _count(result.engine.op_log, "dma_write")
|
||
assert n_writes == 1, f"P=1: one dma_write; got {n_writes}"
|