diff --git a/src/kernbench/benches/_gqa_attention_decode_opt2.py b/src/kernbench/benches/_gqa_attention_decode_opt2.py new file mode 100644 index 0000000..4be0f03 --- /dev/null +++ b/src/kernbench/benches/_gqa_attention_decode_opt2.py @@ -0,0 +1,84 @@ +"""GQA decode opt2 variant — two-composite inner attention (ADR-0065). + +opt2 of decode (ADR-0060 §8 item 4 / ADR-0065): the per-tile attention is +expressed as **two composites** instead of opt3's long primitive sequence — + + #1 Q·Kᵀ → a GEMM composite writing the score tile ``scores``. + #2 softmax → a ``softmax_merge`` recipe composite (online-softmax merge + + P·V of ``(m, l, O)``) whose head GEMM is P·V, with an ``add`` + epilogue folding the P·V contribution into ``O``. + +Fewer PE_CPU-issued commands ⇒ lower dispatch cost under the ADR-0064 Rev2 +structural model (the CPU-offload win). Tile 0 establishes the running +state with primitives (kernbench has no scratch-backed ``-inf`` initializer, +same limitation as the opt3 kernel); the next tile exercises the recipe. + +**Scope (P5 B):** runnable in op_log mode (latency only). Full data-mode +numeric parity (computing the recipe's 8 MATH ops in the DataExecutor) is a +separate follow-up — see DDD-0065 / the P5-numerics note. +""" +from __future__ import annotations + + +def gqa_attention_decode_opt2_kernel( + q_ptr: int, + k_ptr: int, + v_ptr: int, + o_ptr: int, + T_q: int, + S_kv: int, + h_q: int, + h_kv: int, + d_head: int, + C: int, + P: int, + *, + tl, +) -> None: + """Single-rank (C=P=1) GQA decode using the opt2 two-composite form. + + Layout mirrors ``gqa_attention_decode_long_kernel`` (M-fold Q, K loaded + as ``(d_head, S_local)``). The KV slice is split into two sub-tiles so + the second one drives the ``softmax_merge`` recipe composite. + """ + G = h_q // h_kv + S_local = S_kv // (C * P) + KV_ROW_BYTES = d_head * 2 # f16 + + Q = tl.load(q_ptr, shape=(G * T_q, d_head), dtype="f16") + + # Split the local KV slice in two: tile 0 establishes the running state, + # tile 1 merges via the opt2 two-composite path. + half = max(1, S_local // 2) + rest = S_local - half + + # ── Tile 0: establish running (m_local, l_local, O_local) ── + K_T0 = tl.load(k_ptr, shape=(d_head, half), dtype="f16") + V0 = tl.load(v_ptr, shape=(half, d_head), dtype="f16") + scores0 = tl.dot(Q, K_T0) + m_local = tl.max(scores0, axis=-1) + exp0 = tl.exp(scores0 - m_local) + l_local = tl.sum(exp0, axis=-1) + O_local = tl.dot(exp0, V0) + + # ── Tile 1: opt2 two-composite merge (skipped when S_local == 1) ── + if rest > 0: + K_T1 = tl.load(k_ptr + half * KV_ROW_BYTES, + shape=(d_head, rest), dtype="f16") + V1 = tl.ref(v_ptr + half * KV_ROW_BYTES, shape=(rest, d_head), + dtype="f16") + # #1: Q·Kᵀ composite → score tile (TCM-resident, consumed by #2). + scores1 = tl.zeros((G * T_q, rest), dtype="f16") + tl.composite(op="gemm", a=Q, b=K_T1, out=scores1) + # #2: softmax_merge recipe (online merge of (m,l,O)) + P·V + add. + tl.composite( + prologue=[{"op": "softmax_merge", "s": scores1, + "m": m_local, "l": l_local, "O": O_local}], + op="gemm", b=V1, out=O_local, + epilogue=[{"op": "add", "other": O_local}], + ) + + # ── Final normalise + store (root only) ── + if tl.program_id(axis=0) == 0 and tl.program_id(axis=1) == 0: + O_final = O_local / l_local + tl.store(o_ptr, O_final) diff --git a/src/kernbench/components/builtin/pe_scheduler.py b/src/kernbench/components/builtin/pe_scheduler.py index 40a2780..04ba6e9 100644 --- a/src/kernbench/components/builtin/pe_scheduler.py +++ b/src/kernbench/components/builtin/pe_scheduler.py @@ -170,9 +170,13 @@ class PeSchedulerComponent(ComponentBase): plan = self._generate_plan(cmd) self._pipeline_counter += 1 + # Prologue / post-loop single-shot stages each count as a tile for + # completion (ADR-0065 D3). Empty for legacy composites → unchanged. + n_pre = len(getattr(plan, "prologue_stages", ())) + n_post = len(getattr(plan, "epilogue_stages", ())) ctx = PipelineContext( id=f"p{self._pipeline_counter}", - total_tiles=len(plan.tiles), + total_tiles=len(plan.tiles) + n_pre + n_post, done_event=pe_txn.done, ) @@ -198,6 +202,13 @@ class PeSchedulerComponent(ComponentBase): assert self._pending_feeds is not None while True: plan, ctx = yield self._pending_feeds.get() + tid = 0 + # Prologue single-shot stages first (ADR-0065 D3) — each a + # standalone 1-stage tile that runs on its component and + # completes (no inter-component routing). + for st in getattr(plan, "prologue_stages", ()): + yield from self._feed_single_stage(st, ctx, tid) + tid += 1 for tile in plan.tiles: first_stage = tile.stages[0] token = TileToken( @@ -208,6 +219,20 @@ class PeSchedulerComponent(ComponentBase): params=first_stage.params, ) yield self.out_ports[first_stage.component].put(token) + for st in getattr(plan, "epilogue_stages", ()): + yield from self._feed_single_stage(st, ctx, tid) + tid += 1 + + def _feed_single_stage(self, stage: Any, ctx: Any, tile_id: int) -> Generator: + """Feed one standalone stage as a 1-stage tile (prologue/post-loop).""" + from kernbench.components.builtin.pe_types import TilePlan, TileToken + + token = TileToken( + tile_id=tile_id, pipeline_ctx=ctx, + plan=TilePlan(tile_id=tile_id, stages=(stage,)), + stage_idx=0, params=stage.params, + ) + yield self.out_ports[stage.component].put(token) def _generate_plan(self, cmd: Any) -> Any: """Generate a PipelinePlan from a flat-ops CompositeCmd (ADR-0065 D3). diff --git a/src/kernbench/components/builtin/tiling.py b/src/kernbench/components/builtin/tiling.py index 519d1f8..2401a64 100644 --- a/src/kernbench/components/builtin/tiling.py +++ b/src/kernbench/components/builtin/tiling.py @@ -283,22 +283,13 @@ def generate_plan_from_ops( epilogue_specs=tuple(post_ops), ) - pre_stages = tuple(_math_stage(o, pe_prefix) for o in pre_ops) - post_stages = tuple(_math_stage(o, pe_prefix) for o in post_ops - if o.scope == _Scope.KERNEL) - plan.prologue_stages = pre_stages - plan.epilogue_stages = post_stages - - # Fold prologue/post-loop stages into the first/last tile so the feeder - # and completion counting are untouched (existing benches have neither, - # so their tiles are unchanged → byte-equal op_log). - if pre_stages and plan.tiles: - t0 = plan.tiles[0] - plan.tiles[0] = TilePlan(tile_id=t0.tile_id, - stages=(*pre_stages, *t0.stages)) - if post_stages and plan.tiles: - tl = plan.tiles[-1] - plan.tiles[-1] = TilePlan(tile_id=tl.tile_id, - stages=(*tl.stages, *post_stages)) - + # Prologue (pre-GEMM KERNEL) + post-loop KERNEL ops become single-shot + # MATH stages. They are fed as standalone 1-stage tiles by the scheduler + # (PE_SCHEDULER._feed_loop) — NOT folded into the GEMM tiles, because a + # folded MATH→DMA_READ boundary would require a PE_MATH→PE_DMA token + # route that the pipeline does not wire. Existing benches have neither → + # their tiles + feed order are unchanged (byte-equal op_log). + plan.prologue_stages = tuple(_math_stage(o, pe_prefix) for o in pre_ops) + plan.epilogue_stages = tuple(_math_stage(o, pe_prefix) for o in post_ops + if o.scope == _Scope.KERNEL) return plan diff --git a/tests/attention/test_gqa_decode_opt2.py b/tests/attention/test_gqa_decode_opt2.py new file mode 100644 index 0000000..459b228 --- /dev/null +++ b/tests/attention/test_gqa_decode_opt2.py @@ -0,0 +1,173 @@ +"""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.""" + for R in (0.25, 0.0625, 0.03125): + 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}" + + +# ── 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}"