d282144339
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>
199 lines
7.6 KiB
Python
199 lines
7.6 KiB
Python
"""Phase 1 spec test for Phase C: scratch_scope + tl.copy_to discipline
|
||
in the GQA kernels (ADR-0060 §5.2 / §5.5 + ADR-0063 §D3 / §D3.1).
|
||
|
||
ADR-0060 §5.2 (decode pseudocode line 75 / §5.5 (prefill pseudocode line
|
||
96) both wrap per-tile / per-ring-step intermediates in
|
||
``with tl.scratch_scope():`` and persist the merged running ``(m, ℓ, O)``
|
||
to a persistent arena allocated outside the scope. The original ADR-0063
|
||
§D3 specifies the two-arena pattern; §D3.1 specifies the
|
||
``tl.copy_to(dst, src)`` writeback primitive used to persist scoped
|
||
results.
|
||
|
||
Currently neither kernel uses ``scratch_scope`` or ``copy_to``; their
|
||
chain-merge / ring-merge bodies allocate every intermediate from the
|
||
bump cursor and never recycle. Result: op_log has 0 ``copy`` entries.
|
||
|
||
After Phase 2: each per-tile / per-step merge writes the new running
|
||
``(m, ℓ, O)`` via ``copy_to`` to the persistent arena. Per merge step
|
||
→ 3 copy ops (m, ℓ, O).
|
||
|
||
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 import gqa_decode_kernel # noqa: F401
|
||
from kernbench.benches._gqa_prefill import gqa_prefill_kernel # noqa: F401
|
||
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
|
||
from kernbench.ccl.sfr_config import (
|
||
configure_sfr_intercube_multisip,
|
||
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)
|
||
|
||
|
||
# ── Decode chain merges must use scratch_scope + copy_to ─────────────
|
||
|
||
|
||
def _run_decode_sp(*, h_q: int, h_kv: int, P: int, S_kv: int):
|
||
"""Single-CUBE SP decode (C=1, P PEs along intra-cube chain)."""
|
||
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((1, h_q * D_HEAD),
|
||
dtype=DTYPE, dp=dp_full, name=f"q_sc_{P}")
|
||
k = ctx.zeros((S_kv, h_kv * D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"k_sc_{P}")
|
||
v = ctx.zeros((S_kv, h_kv * D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"v_sc_{P}")
|
||
o = ctx.empty((1, h_q * D_HEAD),
|
||
dtype=DTYPE, dp=dp_full, name=f"o_sc_{P}")
|
||
ctx.launch(
|
||
f"gqa_decode_scoped_{P}",
|
||
gqa_decode_kernel,
|
||
q, k, v, o,
|
||
1, S_kv, h_q, h_kv, D_HEAD,
|
||
1, P,
|
||
_auto_dim_remap=False,
|
||
)
|
||
|
||
return run_bench(
|
||
topology=topo, bench_fn=_bench_fn,
|
||
device=resolve_device(None), engine_factory=_engine_factory,
|
||
)
|
||
|
||
|
||
def test_decode_chain_merges_emit_copy_to_writeback():
|
||
"""ADR-0060 §5.2 + ADR-0063 §D3.1: each chain-merge step must wrap
|
||
its intermediates in ``scratch_scope`` and persist the new running
|
||
``(m, ℓ, O)`` via ``tl.copy_to``.
|
||
|
||
For (C=1, P=8): 7 intra-cube chain merges × 3 handles (m, ℓ, O) per
|
||
merge ⇒ 21 ``copy`` entries.
|
||
|
||
Currently 0 because the kernel never calls ``tl.copy_to``.
|
||
"""
|
||
result = _run_decode_sp(h_q=8, h_kv=1, P=8, S_kv=64)
|
||
assert result.completion.ok, f"decode SP failed: {result.completion}"
|
||
n_copy = _count(result.engine.op_log, "copy")
|
||
assert n_copy > 0, (
|
||
f"decode kernel must emit copy_to writeback per merge step "
|
||
f"(ADR-0060 §5.2 + ADR-0063 §D3.1); got 0 ``copy`` entries"
|
||
)
|
||
|
||
|
||
# ── Prefill Ring KV merges must use scratch_scope + copy_to ──────────
|
||
|
||
|
||
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_ring_{C}")
|
||
k = ctx.zeros((S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"k_ring_{C}")
|
||
v = ctx.zeros((S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"v_ring_{C}")
|
||
o = ctx.empty((T_q * C, D_HEAD),
|
||
dtype=DTYPE, dp=dp_o, name=f"o_ring_{C}")
|
||
ctx.launch(
|
||
f"gqa_prefill_scoped_{C}",
|
||
gqa_prefill_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 test_prefill_ring_step_merges_emit_copy_to_writeback():
|
||
"""ADR-0060 §5.5 + ADR-0063 §D3.1: each Ring KV step's online-softmax
|
||
merge must wrap its intermediates in ``scratch_scope`` and persist
|
||
the new running ``(m, ℓ, O)`` via ``tl.copy_to``.
|
||
|
||
For C=4: 3 ring-step merges (steps 1..C-1) × 3 handles (m, ℓ, O) per
|
||
merge ⇒ 9 ``copy`` entries per participating CUBE. Aggregated across
|
||
C CUBEs: ⇒ 36 ``copy`` entries.
|
||
|
||
Currently 0 because the kernel never calls ``tl.copy_to``.
|
||
"""
|
||
result = _run_prefill_ring(T_q=4, S_kv=16, C=4)
|
||
assert result.completion.ok, f"prefill ring failed: {result.completion}"
|
||
n_copy = _count(result.engine.op_log, "copy")
|
||
assert n_copy > 0, (
|
||
f"prefill ring kernel must emit copy_to writeback per merge step "
|
||
f"(ADR-0060 §5.5 + ADR-0063 §D3.1); got 0 ``copy`` entries"
|
||
)
|
||
|
||
|
||
# ── Scoped kernels still produce the same op_log shape (regression) ──
|
||
|
||
|
||
def test_decode_scoped_still_has_root_only_write():
|
||
"""ADR-0060 §A.2 root-only output must hold under the rewrite:
|
||
adding scratch_scope + copy_to should not change the reduce
|
||
topology; only the per-PE scratch usage. Single PE 0 writes O."""
|
||
result = _run_decode_sp(h_q=8, h_kv=1, P=8, S_kv=64)
|
||
assert result.completion.ok
|
||
n_writes = _count(result.engine.op_log, "dma_write")
|
||
assert n_writes == 1, (
|
||
f"scoped decode must still produce 1 dma_write (root-only); "
|
||
f"got {n_writes}"
|
||
)
|
||
|
||
|
||
def test_prefill_scoped_still_has_per_cube_distributed_output():
|
||
"""ADR-0060 §5.5 per-CUBE distributed output must hold under the
|
||
rewrite: scoped prefill still writes one O slice per CUBE."""
|
||
result = _run_prefill_ring(T_q=4, S_kv=16, C=4)
|
||
assert result.completion.ok
|
||
n_writes = _count(result.engine.op_log, "dma_write")
|
||
assert n_writes == 4, (
|
||
f"scoped prefill must write one O per CUBE (C=4 → 4 dma_write); "
|
||
f"got {n_writes}"
|
||
)
|