Files
kernbench2/tests/attention/test_gqa_scoped_writeback.py
T
mukesh 5a76ed4f6a gqa: rename long-context kernels to *_long for symmetry with *_short
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>
2026-06-10 10:05:54 -07:00

199 lines
7.6 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 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_long import gqa_decode_long_kernel # noqa: F401
from kernbench.benches._gqa_prefill_long import gqa_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_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_long_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_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 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}"
)