gqa(adr-0065): P5(B) — decode opt2 two-composite kernel + dispatch measurement
New _gqa_attention_decode_opt2.py: per-tile attention as #1 Q.Kt GEMM composite + #2 softmax_merge recipe composite (online merge + P.V + add). Runnable in op_log mode; full data-mode numeric parity is a deferred follow-up (P5-numerics). Measured per-tile PE_CPU dispatch (ADR-0064 Rev2 default): opt3=960ns vs opt2=184ns = 5.22x (gate >2x), FIXED-dominated (command-count reduction) — the ADR-0060/0064 CPU-offload win. opt2<opt3 across R in {0.25,0.0625,0.03125}. K-before-V: softmax_merge prologue carries no DMA; V (ref) streams only in the GEMM tile loop. Fix latent P3 bug surfaced by the first e2e recipe run: prologue MATH stages were FOLDED into the first GEMM tile, but a folded MATH->DMA_READ boundary needs a PE_MATH->PE_DMA token route the pipeline never wires (KeyError pe_dma). Now prologue/post-loop stages are fed as standalone 1-stage tiles (each completes on its component); completion count includes them. Existing benches have no prologue -> tiles + feed order unchanged (byte-equal); full suite 806 pass / 3 pre-existing fail. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -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)
|
||||
@@ -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).
|
||||
|
||||
@@ -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
|
||||
# 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)
|
||||
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))
|
||||
|
||||
return plan
|
||||
|
||||
@@ -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}"
|
||||
Reference in New Issue
Block a user