"""Phase 1 spec tests for ADR-0065 P5(B) — decode opt2 dispatch measurement. opt2 replaces the per-tile primitive attention block (opt3: many ``tl.*`` ops) with **two composites**: #1 = Q·Kᵀ GEMM, #2 = ``softmax_merge`` recipe (online-softmax merge) + P·V GEMM + ``add`` (ADR-0060 §8 item 4 / ADR-0065). Fewer PE_CPU-issued commands → lower dispatch cost under the ADR-0064 Rev2 structural cost model. This is the headline CPU-offload win. These tests measure **dispatch cost only** (op_log / command-emission level); numeric parity (full data-mode recipe computation) is a separate follow-up. The dispatch-ratio / R-sweep tests exercise already-shipped features (the P0 cost model + the P2 recipe) and pass now. The e2e test needs the P5 production bench kernel and fails until it lands. """ from __future__ import annotations import math from kernbench.common.pe_commands import PeCpuOverheadCmd, TensorHandle from kernbench.common.pe_cost_model import DEFAULT_PE_COST_MODEL, PeCostModel from kernbench.triton_emu.tl_context import TLContext, run_kernel G, T, D = 8, 64, 128 _HID = [0] def _tcm(addr: int, shape: tuple[int, ...]) -> TensorHandle: _HID[0] += 1 return TensorHandle( id=f"s{_HID[0]}", addr=addr, shape=shape, dtype="f16", nbytes=2 * math.prod(shape), space="tcm", pinned=True, ) # ── per-tile attention emitters ────────────────────────────────────── def _opt3_tile(*, tl) -> None: """opt3: primitive per-tile inner attention + online-softmax merge.""" K_T = tl.load(0x1000, shape=(D, T), dtype="f16") Q = tl.load(0x2000, shape=(G, D), dtype="f16") V = tl.load(0x3000, shape=(T, D), dtype="f16") m_local = _tcm(0x10000, (G, 1)) l_local = _tcm(0x11000, (G, 1)) O_local = _tcm(0x12000, (G, D)) scores = tl.dot(Q, K_T) m_tile = tl.max(scores, axis=-1) centered = scores - m_tile exp_s = tl.exp(centered) l_tile = tl.sum(exp_s, axis=-1) O_tile = tl.dot(exp_s, V) m_new = tl.maximum(m_local, m_tile) scale_old = tl.exp(m_local - m_new) scale_new = tl.exp(m_tile - m_new) l_new = l_local * scale_old + l_tile * scale_new O_new = O_local * scale_old + O_tile * scale_new tl.copy_to(m_local, m_new) tl.copy_to(l_local, l_new) tl.copy_to(O_local, O_new) def _opt2_tile(*, tl) -> None: """opt2: #1 Q·Kᵀ composite + #2 softmax_merge recipe composite.""" K_T = tl.load(0x1000, shape=(D, T), dtype="f16") Q = tl.load(0x2000, shape=(G, D), dtype="f16") V = tl.ref(0x3000, shape=(T, D), dtype="f16") m_local = _tcm(0x10000, (G, 1)) l_local = _tcm(0x11000, (G, 1)) O_local = _tcm(0x12000, (G, D)) scores = _tcm(0x13000, (G, T)) tl.composite(op="gemm", a=Q, b=K_T, out=scores) # #1 tl.composite( # #2 prologue=[{"op": "softmax_merge", "s": scores, "m": m_local, "l": l_local, "O": O_local}], op="gemm", b=V, out=O_local, epilogue=[{"op": "add", "other": O_local}], ) def _dispatch_cycles(emitter, cost_model) -> float: tl = TLContext(pe_id=0, num_programs=1, cost_model=cost_model, scratch_base=0x200000, scratch_size=1 << 20) run_kernel(emitter, tl) return sum(c.cycles for c in tl.commands if isinstance(c, PeCpuOverheadCmd)) # ── dispatch ratio (ADR-0064 Test #9 / ADR-0065 Test #7) ───────────── def test_opt3_dispatch_exceeds_opt2_by_2x(): opt3 = _dispatch_cycles(_opt3_tile, DEFAULT_PE_COST_MODEL) opt2 = _dispatch_cycles(_opt2_tile, DEFAULT_PE_COST_MODEL) assert opt2 > 0 and opt3 > 0 assert opt3 > 2 * opt2, f"opt3={opt3} must exceed 2x opt2={opt2}" def test_dispatch_ratio_R_sensitivity(): """ADR-0064 Test #9 — opt2 < opt3 across the queue-bandwidth range, and the opt3/opt2 ratio increases as R decreases (the cost becomes more FIXED-dominated, i.e. command-count-driven). Absolute ratio values are informative only; the gate is the direction + opt2 < opt3.""" ratios = [] for R in (0.25, 0.0625, 0.03125): # strictly decreasing R cm = PeCostModel(fixed_per_cmd_cycles=40, byte_cycles_recip=R) opt3 = _dispatch_cycles(_opt3_tile, cm) opt2 = _dispatch_cycles(_opt2_tile, cm) assert opt2 < opt3, f"R={R}: opt2={opt2} !< opt3={opt3}" ratios.append(opt3 / opt2) assert ratios[0] < ratios[1] < ratios[2], ( f"opt3/opt2 ratio must increase as R decreases; got {ratios}" ) # ── K-before-V DMA priority (ADR-0065 Test #3) ─────────────────────── def test_k_before_v_in_opt2_plan(): """In opt2's #2 composite, the V (ref) DMA_READ is placed *after* the softmax_merge prologue MATH stages — V is not streamed during the prologue (K-before-V priority).""" from kernbench.common.pe_commands import CompositeCmd from kernbench.components.builtin.pe_types import StageType from kernbench.components.builtin.tiling import generate_plan_from_ops tl = TLContext(pe_id=0, num_programs=1, scratch_base=0x200000, scratch_size=1 << 20) run_kernel(_opt2_tile, tl) composites = [c for c in tl.commands if isinstance(c, CompositeCmd)] cmd2 = composites[-1] # the softmax_merge composite plan = generate_plan_from_ops( ops=cmd2.ops, tile_m=32, tile_k=32, tile_n=32, bytes_per_element=2, pe_prefix="sip0.cube0.pe0", ) # The softmax_merge prologue (8 MATH) carries NO DMA — V is not streamed # during it. The prologue is fed before the GEMM tile loop, where the V # (ref) DMA_READ lives — so K (in #1) loads before V (in #2's GEMM). assert len(plan.prologue_stages) == 8, plan.prologue_stages assert all(s.stage_type == StageType.MATH for s in plan.prologue_stages) assert not any(s.stage_type in (StageType.DMA_READ, StageType.DMA_WRITE) for s in plan.prologue_stages) tile_reads = [s for s in plan.tiles[0].stages if s.stage_type == StageType.DMA_READ] assert len(tile_reads) == 1, "exactly the V tile DMA_READ in the loop" # ── e2e: opt2 bench runs in op_log mode (needs P5 production kernel) ── def test_opt2_bench_completes_oplog_mode(): from pathlib import Path from kernbench.benches._gqa_attention_decode_opt2 import ( # noqa: F401 gqa_attention_decode_opt2_kernel, ) 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 topo_path = Path(__file__).resolve().parents[2] / "topology.yaml" topo = resolve_topology(str(topo_path)) S_KV = 16 def _bench_fn(ctx): dp = DPPolicy(cube="replicate", pe="replicate", num_cubes=1, num_pes=1) q = ctx.zeros((1, 8 * D), dtype="f16", dp=dp, name="q_opt2") k = ctx.zeros((S_KV, D), dtype="f16", dp=dp, name="k_opt2") v = ctx.zeros((S_KV, D), dtype="f16", dp=dp, name="v_opt2") o = ctx.empty((1, 8 * D), dtype="f16", dp=dp, name="o_opt2") ctx.launch("gqa_decode_opt2", gqa_attention_decode_opt2_kernel, q, k, v, o, 1, S_KV, 8, 1, D, 1, 1, _auto_dim_remap=False) result = run_bench( topology=topo, bench_fn=_bench_fn, device=resolve_device(None), engine_factory=lambda t, d: GraphEngine(getattr(t, "topology_obj", t), enable_data=False), ) assert result.completion.ok, f"opt2 decode failed: {result.completion}" # ── opt2 ↔ opt3 numeric parity in data mode (ADR-0065 N4) ───────────── def _run_decode_data(kernel, name, q_d, k_d, v_d, *, S_kv=16): """Run a decode kernel (C=P=1) in data mode with seeded Q/K/V; return the HBM output array.""" from pathlib import Path import numpy as np 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.data_executor import DataExecutor from kernbench.sim_engine.engine import GraphEngine from kernbench.topology.builder import resolve_topology topo = resolve_topology(str(Path(__file__).resolve().parents[2] / "topology.yaml")) cap: dict = {} def bench_fn(ctx): dp = DPPolicy(cube="replicate", pe="replicate", num_cubes=1, num_pes=1) q = ctx.zeros((1, 8 * D), dtype="f16", dp=dp, name=name + "q") k = ctx.zeros((S_kv, D), dtype="f16", dp=dp, name=name + "k") v = ctx.zeros((S_kv, D), dtype="f16", dp=dp, name=name + "v") o = ctx.empty((1, 8 * D), dtype="f16", dp=dp, name=name + "o") q.copy_(ctx.from_numpy(q_d)); k.copy_(ctx.from_numpy(k_d)); v.copy_(ctx.from_numpy(v_d)) cap["o"] = o ctx.launch(name, kernel, q, k, v, o, 1, S_kv, 8, 1, D, 1, 1, _auto_dim_remap=False) r = run_bench(topology=topo, bench_fn=bench_fn, device=resolve_device(None), engine_factory=lambda t, d: GraphEngine(getattr(t, "topology_obj", t), enable_data=True)) assert r.completion.ok DataExecutor(r.engine.op_log, r.engine.memory_store).run() o = cap["o"] return r.engine.memory_store.read("hbm", o._handle.va_base, shape=(1, 8 * D), dtype="f16") def test_opt2_matches_opt3_data_mode(): import numpy as np from kernbench.benches._gqa_attention_decode_long import gqa_attention_decode_long_kernel from kernbench.benches._gqa_attention_decode_opt2 import gqa_attention_decode_opt2_kernel rng = np.random.default_rng(0) q_d = (0.1 * rng.standard_normal((1, 8 * D))).astype(np.float16) k_d = (0.1 * rng.standard_normal((16, D))).astype(np.float16) v_d = (0.1 * rng.standard_normal((16, D))).astype(np.float16) o3 = _run_decode_data(gqa_attention_decode_long_kernel, "p3", q_d, k_d, v_d) o2 = _run_decode_data(gqa_attention_decode_opt2_kernel, "p2", q_d, k_d, v_d) assert np.allclose(np.asarray(o2, np.float32), np.asarray(o3, np.float32), rtol=5e-2, atol=5e-2), f"opt2 {np.asarray(o2).ravel()[:4]} vs opt3 {np.asarray(o3).ravel()[:4]}"