5 Commits

Author SHA1 Message Date
ywkang a4a2683aad paper(gqa): single-rank decode-streaming study — composite's masked bandwidth win
Adds the memory-bound mirror of the compute-bound prefill figure. New single-rank decode sweep (T_q=1, M=8) reuses the prefill_compute_bound kernels, sweeping per-rank S_kv with NO cross-CUBE reduce — isolating the local attention that the 64-way Case-6 decode masks under its reduce tail.

Finding: even memory-bound decode benefits from the composite command. Its scheduler-streamed concurrent per-tile DMAs reach ~233 GB/s (91% of the 256 GB/s per-rank roofline) while the primitive's blocking tl.dot serializes one tile DMA and plateaus at ~166 GB/s — a 25-28% latency win that widens with context. This refines the paper's 'decode is latency-neutral' claim: neutral at 64-way production scale (reduce-dominated), but the composite extracts bandwidth at the local level. The bandwidth roofline here mirrors the MAC roofline in prefill.

Adds: gqa_decode_streaming.py sweep (wired into the milestone umbrella as GQA_1H_SWEEPS=decode_streaming), paper_plot_gqa_decode_streaming.py, the figure, and a new fig:gqa-decode-stream + paragraph in 05-gqa.tex; rebuilt main.pdf. Prefill (Experiment B) re-verified byte-identical (794.1/668.1/646.9us).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-23 19:39:45 -07:00
ywkang faf011dae5 paper(ipcq): regen alternatives figures + comparison plot + allreduce section
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-23 19:19:55 -07:00
ywkang 1dade267ff experiment(gqa-decode): 16x16x16 MAC-block dispatch model
Research artifact (NOT wired into the production sweep) for the decode composite-vs-primitive investigation. Models a primitive decode kernel that hand-blocks each tl.dot into 16x16x16 GemmCmds, charging per-MAC-block PE_CPU dispatch -- testing whether faithfully charging the dispatch that the composite form offloads to PE_SCHEDULER flips the 'composite gives no decode latency benefit' result.

Finding: it does NOT. Even with up to 512 serialized blocking GemmCmds per matmul (vs 2 coarse tl.dot), end-to-end decode latency is unchanged (8192: 28.9us, 32768: 114.9us) -- the inter-cube (m,l,O) reduce DMA tail dominates and the local-attention GEMM/dispatch sits in critical-path slack. Confirms the two-regime conclusion (24c7054): composite helps compute-bound prefill, not memory/reduce-bound decode.

Caveats: measured at S_kv 8K/32K (reduce-dominated); large-S_kv streaming regime extrapolated only. The -5.5% tiled<untiled is reduce-tail ordering noise, not fully nailed down.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-23 18:24:26 -07:00
ywkang 24c705419e paper(gqa): two-regime "use of composite commands" — decode + compute-bound prefill
Complete the composite-command study across both attention regimes and
write up the result, resolving when the composite command helps latency.

Decode (memory-bound, T_q=1, M=8): the command form is latency-neutral —
the kernel is bound by KV streaming and the MAC array is near-idle, so the
composite's only benefit here is host-issue offload (O(n_tiles) -> O(1)
PE_CPU commands). Figure: gqa_decode_long_ctx_composite (latency flat,
command count saturates).

Prefill (compute-bound, large M=G*T_q tile-filling): add three command-form
variants of a single-rank FlashAttention prefill kernel
(_gqa_prefill_compute_bound: primitive / composite / composite_extended).
Here the composite's scheduler-internal per-tile DMA<->compute pipeline
keeps the MAC array fed while the primitive's blocking tl.dot starves it,
so composite wins on MAC utilization (68% flat -> 83%) and wall-clock
(19% faster at ctx=1024), with the margin growing with context (deeper
P.V reduction = more tiles to pipeline) — the compute-bound mirror of the
GEMM result in section 3. Figure: gqa_prefill_compute_bound (latency +
MAC util). Sweep wired as GQA_1H_SWEEPS=prefill_cb; tests cover structure,
e2e, and composite<primitive in the compute-bound regime.

The synthesis: composite has two benefits set by roofline position —
host-issue offload (always) and MAC-array feeding (compute-bound only).
Decode exercises the first, prefill both.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-23 17:04:26 -07:00
ywkang cd6f0ed91d feat(gqa): Case-6 composite-command decode variants + on-chip-operand DMA fix
Add two command-form variants of the Case-6 (Cube-SP x PE-SP) long-context
decode kernel alongside the primitive baseline:
  - composite: per-tile GEMMs issued as one coarse tl.composite over the
    full S_local (PE_SCHEDULER tiles/streams K,V) -- O(1) PE_CPU commands
    vs the primitive kernel's O(n_tiles).
  - composite_extended: Q.K^T composite + softmax_merge recipe composite
    (ADR-0065), folding the per-tile online merge + P.V.

Extract the shared two-level (m,l,O) reduce into _gqa_mlo_reduce and
refactor the baseline to use it (byte-equal, guarded by a 64-rank
command-stream digest test).

Engine fix: _make_compute_out omitted pinned=True, so an on-chip TCM
compute result consumed as a composite GEMM operand was scheduled as an
HBM DMA_READ of its bit-61 scratch address -> PhysAddrError in data mode.
Mark compute outputs pinned (resident), matching the composite
auto-output and recipe scratch handles. .pinned is read only by the
tiling DMA gate, so this is byte-equal for existing benches.

Add the composite sweep (emit-level PE_CPU dispatch to 1M + data-mode
latency to 128K), its plot, umbrella wiring (GQA_1H_SWEEPS=composite),
and tests (refactor guard, dispatch saturation, engine-fix, e2e).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-23 16:07:41 -07:00
41 changed files with 2844 additions and 230 deletions
Binary file not shown.
Binary file not shown.

After

Width:  |  Height:  |  Size: 136 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 136 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 145 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 230 KiB

After

Width:  |  Height:  |  Size: 305 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 260 KiB

After

Width:  |  Height:  |  Size: 381 KiB

@@ -145,9 +145,9 @@ with credit return, splitting the control plane into PE\_IPCQ and the
data plane into PE\_DMA, with head updates riding the payload and tail
updates riding a 16\,B side-channel credit (\S\ref{sec:allreduce}).
\begin{figure}[t]
\begin{figure*}[t]
\centering
\includegraphics[width=\linewidth]{ipcq_alternatives_architecture_stacked.png}
\includegraphics[width=0.78\linewidth]{ipcq_alternatives_architecture_flow.png}
\caption{Per-send data and control flow for the four PE-to-PE
signalling mechanisms (sender\,$\rightarrow$\,NoC\,$\rightarrow$\,receiver).
Doorbell and RDMA-CQ each issue two fabric transactions (payload then
@@ -158,7 +158,7 @@ the payload flit train and returns the tail credit on a side channel,
so a send is one MMIO write and a receive is a flip-flop read. This is
a \emph{design schematic}, not a measured comparison.}
\label{fig:ipcq-arch}
\end{figure}
\end{figure*}
\begin{figure}[t]
\centering
@@ -122,9 +122,10 @@ overlaps score computation; and per-tile \emph{scratch recycling} that
keeps the running softmax accumulators ($m,\ell,O$) in a persistent arena
while freeing per-tile temporaries, so the kernel fits the
\SI{1}{\mebi\byte} scratch budget across many tiles. A further refinement
that restructures the decode step into two stateful composites (a named
\textsf{softmax\_merge} recipe) is designed but not yet wired into the
measured path; results below reflect the implemented kernel only.
restructures the decode step so the per-tile matrix products and the
online-softmax merge are issued as \emph{composite} commands rather than
hand-tiled primitives; \S\ref{sec:gqa-composite} measures it against the
primitive baseline at long context.
% TODO: CUBE <-> KV-head mapping diagram for the short-context regime
% (h_kv=8 KV heads -> 8 CUBEs, 1:1; intra-CUBE PE usage).
@@ -221,6 +222,171 @@ memory-feasible family, and the cross-PE softmax reduction it does pay is
precisely the traffic the communication-side codesign of this report is
built to move quickly.
\subsection{Use of Composite Commands}
\label{sec:gqa-composite}
The decode kernel of \S\ref{sec:gqa-long} issues its local attention as
primitive operations: it walks each PE's $S_{\text{local}}$ token slice in
\SI{1024}{}-token tiles, and for every tile issues a $Q\!\cdot\!K^{\top}$
\textsf{dot}, the online-softmax primitives, and a $P\!\cdot\!V$
\textsf{dot}, merging the running $(m,\ell,O)$ state by hand. The number
of PE\_CPU commands this costs grows with the context. At a production
context of $S_{kv}{=}1\,\text{M}$ tokens the Case-6 64-way split gives
each PE $S_{\text{local}}{=}16384$ tokens, so its local attention is a
$Q\!\cdot\!K^{\top}$ of $(8,128)\!\cdot\!(128,16384)$ and a
$P\!\cdot\!V$ of $(8,16384)\!\cdot\!(16384,128)$---sixteen hand-issued
tiles, each a fresh batch of CPU commands.
The composite command lets the kernel hand that tiling to PE\_SCHEDULER.
We compare three command forms of the \emph{same} Case-6 kernel---identical
placement and identical $(m,\ell,O)$ reduce, differing only in how the
local attention is issued:
\begin{itemize}
\item \textbf{primitive}---the hand-tiled \textsf{dot}/softmax kernel
above (the baseline of \S\ref{sec:gqa-long}).
\item \textbf{composite}---each matrix product is one coarse
\textsf{composite} GEMM over the \emph{whole} $S_{\text{local}}$, with
$K$ and $V$ passed as HBM references so PE\_SCHEDULER streams and
tiles them on the fixed $32\!\times\!64\!\times\!32$ MAC tile; the
softmax stays primitive.
\item \textbf{composite\,+\,softmax\_merge}---additionally folds the
online-softmax merge and $P\!\cdot\!V$ into a single stateful
\textsf{softmax\_merge} recipe composite.
\end{itemize}
\begin{figure}[t]
\centering
\includegraphics[width=\linewidth]{gqa_decode_long_ctx_composite.png}
\caption{Three command forms of the Case-6 decode kernel, swept over
context length ($S_{\text{local}}{=}S_{kv}/64$ per PE). \emph{Right:}
PE\_CPU commands issued. The hand-tiled primitive kernel rises
$O(n_{\text{tiles}})$---from 96 commands at one tile to 426 at the
1\,M-token, sixteen-tile production point---while both composite forms
issue a context-\emph{independent} $O(1)$ count (94 and 98) that
\emph{saturates}: one coarse descriptor offloads the entire per-tile
fan-out. \emph{Left:} the consequence for wall-clock latency is none---all
three land on the same curve (\SI{30.6}{}, \SI{231}{},
\SI{461}{\micro\second} at 8\,K\,/\,64\,K\,/\,128\,K), because decode is
bound by streaming the KV cache, not by issue. Command-count is measured
at emit time (exact, to 1\,M); latency on the data-mode engine over the
tractable range.}
\label{fig:gqa-composite}
\end{figure}
Figure~\ref{fig:gqa-composite} reads off the two quantities that matter,
and they point in opposite directions. The PE\_CPU command count (right)
collapses from a context-growing $O(n_{\text{tiles}})$ to a flat $O(1)$:
at 1\,M tokens the composite form issues \num{94} commands against the
primitive kernel's \num{426}, a $4.5\times$ reduction
($4.2\times$ in modeled dispatch cost), and---crucially---that number no
longer grows with context. The wall-clock latency (left), by contrast,
is unchanged across all three forms: decode is bound by streaming the KV
cache out of HBM, so the command form does not move the critical path.
That juxtaposition is the point. The composite command is not a latency
optimization for this memory-bound decode \emph{at the 64-way production scale} (a
single-rank caveat follows, Figure~\ref{fig:gqa-decode-stream}); it is a
\emph{CPU-issue}
optimization. Its value is removing the per-tile dispatch work that would
otherwise grow without bound as context grows, freeing PE\_CPU to run
ahead and keep the engines fed---which is exactly what lets the
data-movement cost analyzed next show through as the true bottleneck
rather than being masked by issue overhead. The \textsf{softmax\_merge}
recipe folds the online merge into the same descriptor; on this
memory-bound path its marginal cost over the plain GEMM composite is small
(98 vs.\ 94 commands), and like the plain composite it keeps the issued
count flat as context scales.
\paragraph{Isolating the rank: the masked streaming win.} That neutrality
is a property of the \emph{full 64-way} critical path, not of the local
attention: at production scale the inter-CUBE $(m,\ell,O)$ reduce tail and
shared-HBM contention set the wall clock, so a faster local attention does
not surface. Stripping those away---a single rank, no cross-CUBE reduce,
swept over the per-rank context $S_{kv}$ (so $S_{kv}{=}16$\,K here is the
per-PE load of a 1M-token, 64-way-sharded decode)---exposes the local
attention directly (Figure~\ref{fig:gqa-decode-stream}), and the composite
\emph{does} win, by \SI{25}{}--\SI{28}{\percent}. The reason is the
memory-bound mirror of prefill: its scheduler-streamed concurrent per-tile
DMAs keep the HBM pipeline full and reach \SI{233}{\giga\byte\per\second}
---\SI{91}{\percent} of the per-rank \SI{256}{\giga\byte\per\second}
roofline---whereas the primitive kernel's blocking \textsf{tl.dot}
serializes one tile DMA at a time and plateaus at
\SI{166}{\giga\byte\per\second}. So even for memory-bound decode the
composite is not \emph{only} a CPU-issue optimization---it also extracts
bandwidth---but that latency benefit materializes only when the local
attention is on the critical path, which at 64-way production scale it is
not.
\begin{figure}[t]
\centering
\includegraphics[width=\linewidth]{gqa_decode_streaming.png}
\caption{Single-rank memory-bound decode ($T_q{=}1$, $M{=}8$), three
command forms, swept over per-rank context. \emph{Left:} end-to-end
latency---the composite forms run \SI{25}{}--\SI{28}{\percent} below the
primitive, a gap that widens with context. \emph{Right:} achieved HBM
bandwidth against the per-rank \SI{256}{\giga\byte\per\second} roofline.
The primitive's blocking load$\rightarrow$dot serializes the KV stream and
plateaus at \SI{166}{\giga\byte\per\second}; the composite forms pipeline
concurrent per-tile DMAs through the scheduler and reach
\SI{233}{\giga\byte\per\second}. This is the memory-bound mirror of the
prefill result (Figure~\ref{fig:gqa-prefill-cb}): there the composite
approaches the MAC roofline, here the bandwidth roofline. Capped at
$16$\,K---the plain composite materializes the full $(M,S_{kv})$ scores in
TCM, so beyond that only the \textsf{softmax\_merge} recipe, which tiles
the softmax, stays within scratch.}
\label{fig:gqa-decode-stream}
\end{figure}
\paragraph{The compute-bound mirror: prefill.} Decode's \emph{production}
verdict---command form is latency-neutral at 64-way scale---is a property
of its regime, not of the
composite command. A decode step has $T_q{=}1$, so its score and context
products are skinny ($M{=}G\,T_q{=}8$): the MAC array is barely fed and
the kernel is bound by streaming the KV cache. Prefill is the opposite
corner. It processes a block of query positions at once, so $M{=}G\,T_q$
is large and tile-filling, the GEMMs carry real arithmetic intensity
($\sim$$M$ flops/byte, well above the roofline ridge), and the kernel is
\emph{compute-bound}. This is the regime the composite command was built
for (\S\ref{sec:gemm}): it streams the per-HW-tile
DMA$\rightleftarrows$compute pipeline so the MAC array stays fed, whereas
the primitive kernel's blocking \textsf{tl.dot} serializes each tile's
load and compute and starves the array between tiles. We run the same
three command forms on a single-rank compute-bound prefill (FlashAttention
$Q$-block $\times$ $S_{kv}$-tile, online softmax) and sweep the context
length (Figure~\ref{fig:gqa-prefill-cb}).
\begin{figure}[t]
\centering
\includegraphics[width=\linewidth]{gqa_prefill_compute_bound.png}
\caption{Compute-bound prefill, three command forms, swept over context
length ($M{=}8\,T_q$ tile-filling). \emph{Left:} end-to-end latency.
\emph{Right:} MAC utilization (achieved $\div$ the
\SI{8}{\tera\flop\per\second} per-PE peak). The hand-tiled primitive sits
flat at $\sim$\SI{68}{\percent}---its serial load$\to$dot path leaves the
MAC array idle between tiles regardless of context. The composite forms
climb with context (\SI{67}{}$\to$\SI{80}{\percent} plain,
\SI{67}{}$\to$\SI{83}{\percent} with the recipe) because a deeper $P\!\cdot
\!V$ reduction gives more HW tiles to pipeline, and they convert that into
wall-clock: at \num{1024} the recipe form is \SI{646.9}{} vs.\
\SI{794.1}{\micro\second} (\SI{19}{\percent} faster). The margin
\emph{grows} with context---the compute-bound mirror of the GEMM result of
\S\ref{sec:gemm}.}
\label{fig:gqa-prefill-cb}
\end{figure}
The two studies together state the composite command's value precisely. It
has two distinct benefits, and which one matters is set by the workload's
roofline position. The first is \emph{host-issue offload}: one macro
command in place of $O(n_{\text{tiles}})$ fine ones, which removes
PE\_CPU dispatch work and is regime-independent (it shows in the decode
command count). The second is \emph{MAC-array feeding}: the
scheduler-internal per-tile DMA$\rightleftarrows$compute pipeline, which
only converts to latency when the workload is compute-bound enough to have
a MAC array worth keeping busy (it shows in the prefill utilization). A
memory-bound decode exercises only the first; a compute-bound prefill
exercises both. The composite command is the single mechanism that
delivers each where it applies.
\subsection{Comprehensive Analysis}
\label{sec:gqa-analysis}
+204
View File
@@ -0,0 +1,204 @@
#!/usr/bin/env python3
"""SCRATCH EXPERIMENT (not production; do not commit).
Question: does charging the *primitive* decode kernel for per-HW-tile
(16x16x16) CPU dispatch flip the "composite gives no decode-latency
benefit" conclusion?
We monkeypatch TLContext.dot so that, in the primitive kernel, every
tl.dot whose (M,K,N) exceeds the HW GEMM tile (mac_m/mac_k/mac_n) is
split by the CPU into ceil(M/mac_m)*ceil(K/mac_k)*ceil(N/mac_n)
HW-tile-sized GemmCmds. Each tile GemmCmd is emitted through the normal
_emit() path, so it (a) charges PE_CPU dispatch overhead via
_charge_dispatch (PeCpuOverheadCmd), and (b) blocks like a normal
single-op GemmCmd on PE_GEMM at the cycle-accurate ceil-product latency.
We also inject mac_m/mac_k/mac_n into every pe_gemm topology node so BOTH
the primitive-tiled and the composite variants run on the *same*
cycle-accurate engine (fair comparison). Composite is left untouched:
the CPU emits ONE CompositeCmd, and PE_SCHEDULER tiles internally (no
per-HW-tile CPU dispatch).
Data correctness:
Inputs are ctx.zeros (q/k/v), so every matmul result is zeros and the
DataExecutor replay is trivial. To keep replay numerically correct
regardless, exactly ONE emitted tile per dot carries the *real* full
operands+output handles (so the DataExecutor computes the true (M,N)
result via the recorded handle shapes), while its timing fields
(m,k,n) are the HW-tile size so the engine charges exactly one tile of
cycle time. The remaining n_tiles-1 emitted tiles are timing-only
GemmCmds (16x16x16) writing to throwaway scratch. Net: n_tiles tiles
of engine time + n_tiles dispatch charges, and a correct final output.
Engine mode: enable_data=True (same as the production sweep's
_engine_latency_ns), op_log end-to-end latency.
"""
from __future__ import annotations
import json
from math import ceil
from pathlib import Path
# --- HW GEMM tile under test -------------------------------------------------
MAC_M = 16
MAC_K = 16
MAC_N = 16
# Legacy alt to also try: (8, 16, 32)
S_KV_LATENCY = (8192, 32_768, 65_536, 131_072)
ROOT = Path(__file__).resolve().parent
SWEEP_JSON = (
ROOT / "src" / "kernbench" / "benches" / "1H_milestone_output"
/ "gqa" / "long_ctx" / "sweep_decode_composite.json"
)
# --- mac-dim topology override -----------------------------------------------
def _topo_with_mac(mac_m: int, mac_k: int, mac_n: int):
"""Compiled topology with mac dims injected into every pe_gemm node."""
from kernbench.topology.builder import resolve_topology
handle = resolve_topology("topology.yaml")
g = handle.topology_obj
n = 0
for node in g.nodes.values():
if node.kind == "pe_gemm":
node.attrs["mac_m"] = mac_m
node.attrs["mac_k"] = mac_k
node.attrs["mac_n"] = mac_n
n += 1
print(f" injected mac=({mac_m},{mac_k},{mac_n}) into {n} pe_gemm nodes")
return handle
# --- tiling monkeypatch for TLContext.dot ------------------------------------
def _make_tiled_dot(orig_dot, mac_m: int, mac_k: int, mac_n: int):
from kernbench.common.pe_commands import GemmCmd
def tiled_dot(self, a, b):
if len(a.shape) < 2 or len(b.shape) < 2:
return orig_dot(self, a, b)
m, k = a.shape[-2], a.shape[-1]
k2, n = b.shape[-2], b.shape[-1]
if k != k2:
raise ValueError(f"dot shape mismatch: a.K={k} != b.K={k2}")
n_tiles = ceil(m / mac_m) * ceil(k / mac_k) * ceil(n / mac_n)
out_shape = (*a.shape[:-2], m, n)
out = self._make_compute_out(shape=out_shape, dtype=a.dtype)
self._await_pending(a, b)
if n_tiles <= 1:
self._emit(GemmCmd(a=a, b=b, out=out, m=m, k=k, n=n))
return out
# One real-data tile: full handles (so DataExecutor computes the
# true result), but timing fields = HW tile (one tile of cycles).
self._emit(GemmCmd(a=a, b=b, out=out, m=mac_m, k=mac_k, n=mac_n))
# Remaining timing-only tiles: throwaway scratch, 16x16x16.
scratch = self._make_compute_out(shape=(mac_m, mac_n), dtype=a.dtype)
for _ in range(n_tiles - 1):
self._emit(GemmCmd(a=a, b=b, out=scratch,
m=mac_m, k=mac_k, n=mac_n))
return out
return tiled_dot
# --- latency runner (replicates sweep's _engine_latency_ns) ------------------
def _engine_latency_ns(variant: str, S_kv: int, topo) -> float:
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_composite import ( # noqa: E501
_end_to_end_ns, _run_panel_fn,
)
from kernbench.runtime_api.bench_runner import run_bench
from kernbench.runtime_api.types import resolve_device
from kernbench.sim_engine.engine import GraphEngine
result = run_bench(
topology=topo, bench_fn=_run_panel_fn(variant, S_kv),
device=resolve_device(None),
engine_factory=lambda t, d: GraphEngine(
getattr(t, "topology_obj", t), enable_data=True,
),
)
if not result.completion.ok:
raise RuntimeError(
f"{variant}@{S_kv} failed: {result.completion}"
)
return _end_to_end_ns(result.engine.op_log)
def _emit_dispatch(variant: str, S_kv: int) -> int:
"""PE_CPU command count at the center rank (cube 6, pe 0)."""
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_composite import ( # noqa: E501
_emit_dispatch as prod_emit,
)
return prod_emit(variant, S_kv)[0]
def main() -> None:
import kernbench.triton_emu.tl_context as tlc
# Baseline (A): primitive UNTILED latencies from the production sweep
# (mac=0 / TFLOPS model). Read straight off the committed sweep JSON.
sweep = json.loads(SWEEP_JSON.read_text())
base_A = {}
for r in sweep["rows"]:
if r["variant"] == "primitive" and r["latency_ns"] is not None:
base_A[r["S_kv"]] = r["latency_ns"]
print(f"== mac tile = ({MAC_M},{MAC_K},{MAC_N}) ==")
# --- command-count sanity (emit-time, mac-independent) ---------------
orig_dot = tlc.TLContext.dot
print("\n[dispatch counts @ S_kv=131072]")
n_prim_untiled = _emit_dispatch("primitive", 131072)
tlc.TLContext.dot = _make_tiled_dot(orig_dot, MAC_M, MAC_K, MAC_N)
try:
n_prim_tiled = _emit_dispatch("primitive", 131072)
n_comp = None
finally:
tlc.TLContext.dot = orig_dot
n_comp = _emit_dispatch("composite", 131072)
print(f" primitive UNTILED PE_CPU cmds : {n_prim_untiled}")
print(f" primitive TILED PE_CPU cmds : {n_prim_tiled} "
f"(x{n_prim_tiled / max(n_prim_untiled,1):.0f})")
print(f" composite PE_CPU cmds : {n_comp}")
# --- latency sweep ----------------------------------------------------
rows = []
topo = _topo_with_mac(MAC_M, MAC_K, MAC_N)
for S_kv in S_KV_LATENCY:
A = base_A.get(S_kv)
# (C) composite on the mac engine (untouched dot path)
C = _engine_latency_ns("composite", S_kv, topo)
# (B) primitive TILED on the mac engine
tlc.TLContext.dot = _make_tiled_dot(orig_dot, MAC_M, MAC_K, MAC_N)
try:
B = _engine_latency_ns("primitive", S_kv, topo)
finally:
tlc.TLContext.dot = orig_dot
gap_pct = (B - C) / C * 100.0 if C else float("nan")
rows.append((S_kv, A, B, C, gap_pct))
print(f" S_kv={S_kv:>7}: A(untiled)={A!s:>12} "
f"B(tiled)={B:12.2f} C(comp)={C:12.2f} (B-C)/C={gap_pct:+6.1f}%")
# --- final table ------------------------------------------------------
print("\n==================== RESULT TABLE ====================")
print(f"{'S_kv':>8} | {'A untiled(ns)':>14} | {'B tiled(ns)':>14} | "
f"{'C comp(ns)':>14} | {'(B-C)/C':>9}")
print("-" * 72)
for S_kv, A, B, C, gap in rows:
a_s = f"{A:.2f}" if A is not None else "n/a"
print(f"{S_kv:>8} | {a_s:>14} | {B:>14.2f} | {C:>14.2f} | "
f"{gap:>+8.1f}%")
if __name__ == "__main__":
main()
@@ -0,0 +1,146 @@
"""Comparative figure for the Case-6 composite-command decode study.
Reads sweep_decode_composite.json (emitted by milestone-1h-gqa, sweep
``composite``) and writes one two-panel PNG into the bench-output dir:
gqa_decode_long_ctx_composite.png
Left — end-to-end decode latency (µs) vs context length, per command
form (primitive / composite / composite_extended).
Right — PE_CPU command count vs context length: the hand-tiled
primitive kernel issues O(n_tiles) commands (rises with
context), while the coarse composite forms issue O(1) and
*saturate* — PE_SCHEDULER absorbs the per-tile fan-out.
The x-axis is the global context length S_kv; each PE owns
S_local = S_kv/(C·P=64) tokens, so at the 1M production point every PE
runs Q·Kᵀ of (G·T_q, d_head)·(d_head, 16384) and P·V of
(G·T_q, 16384)·(16384, d_head).
Run (after the bench):
GQA_1H_RUN=1 GQA_1H_SWEEPS=composite python -m kernbench.cli.main run \\
--bench milestone-1h-gqa --topology topology.yaml
python scripts/paper/paper_plot_gqa_decode_long_ctx_composite.py
"""
from __future__ import annotations
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt # noqa: E402
_REPO_ROOT = Path(__file__).resolve().parents[2]
_FIG_DIR = (
_REPO_ROOT / "src" / "kernbench" / "benches"
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _FIG_DIR / "sweep_decode_composite.json"
_PAPER_FIG_DIR = (
_REPO_ROOT / "docs" / "report" / "1H-codesign-paper" / "figures"
)
_N_RANKS = 64 # C·P for the Case-6 64-way split.
# variant key → (display label, colour, marker)
_VARIANT_STYLE = {
"primitive": ("primitive (tl.dot, hand-tiled)", "#c0504d", "o"),
"composite": ("composite GEMM", "#3b6ea5", "s"),
"composite_extended": ("composite + softmax_merge", "#4f8a4f", "^"),
}
_ORDER = ("primitive", "composite", "composite_extended")
def _load() -> dict:
return json.loads(_SWEEP_JSON.read_text())
def _series(rows: list[dict], variant: str, key: str):
"""Sorted (S_kv, value) series for a variant, skipping null values
(latency is only measured over the tractable S_kv subset)."""
pts = sorted(
((r["S_kv"], r[key]) for r in rows
if r["variant"] == variant and r.get(key) is not None),
key=lambda t: t[0],
)
return [p[0] for p in pts], [p[1] for p in pts]
def _xticklabels(s_kvs: list[int]) -> list[str]:
out = []
for s in s_kvs:
if s >= 1 << 20:
out.append(f"{s // (1 << 20)}M")
else:
out.append(f"{s // 1024}K")
return out
def main() -> None:
sweep = _load()
rows = sweep["rows"]
s_kv_op = sweep["s_kv_opcount"]
s_kv_lat = sweep["s_kv_latency"]
fig, (ax_lat, ax_cmd) = plt.subplots(1, 2, figsize=(13.0, 4.8))
for v in _ORDER:
label, color, marker = _VARIANT_STYLE[v]
xs, lat = _series(rows, v, "latency_ns")
ax_lat.plot(xs, [y / 1e3 for y in lat], marker=marker,
color=color, label=label, lw=2)
xs, cmds = _series(rows, v, "pe_cpu_cmd_count")
ax_cmd.plot(xs, cmds, marker=marker, color=color, label=label, lw=2)
ax_lat.set_xticks(s_kv_lat)
ax_lat.set_xticklabels(_xticklabels(s_kv_lat))
ax_cmd.set_xticks(s_kv_op)
ax_cmd.set_xticklabels(_xticklabels(s_kv_op), fontsize=8)
for ax in (ax_lat, ax_cmd):
ax.set_xscale("log", base=2)
ax.set_xlabel(
r"context length $S_{kv}$ "
r"($S_{\mathrm{local}}=S_{kv}/64$ per PE)"
)
ax.grid(True, ls=":", alpha=0.5)
ax.legend(fontsize=9)
ax_lat.set_ylabel("end-to-end decode latency (µs)")
ax_lat.set_title(
"Case-6 decode latency per command form\n"
"(memory-bound: command form does not move the critical path)"
)
ax_cmd.set_ylabel("PE_CPU commands issued")
ax_cmd.set_title(
"PE_CPU command count per command form\n"
"(primitive O(n$_\\mathrm{tiles}$) rises; composite O(1) saturates)"
)
ax_cmd.axvline(1 << 20, color="#888", ls="--", lw=1, alpha=0.7)
ax_cmd.annotate("1M production\ncontext", xy=(1 << 20, 0),
xytext=(1 << 18, 0.6), fontsize=8,
textcoords=("data", "axes fraction"), ha="right",
color="#555")
fig.suptitle(
"Case-6 (Cube-SP × PE-SP) long-context decode — use of composite "
"commands\nLLaMA-3.1-70B single-KV-head group (8 cubes × 8 PEs), "
"$T_q{=}1$",
fontsize=11,
)
fig.tight_layout(rect=(0, 0, 1, 0.94))
out = _FIG_DIR / "gqa_decode_long_ctx_composite.png"
fig.savefig(out, dpi=150)
plt.close(fig)
print(f"wrote {out}")
# Mirror into the paper figures dir (derived artifact).
if _PAPER_FIG_DIR.is_dir():
dst = _PAPER_FIG_DIR / out.name
dst.write_bytes(out.read_bytes())
print(f"copied {dst}")
if __name__ == "__main__":
main()
@@ -0,0 +1,124 @@
"""Comparative figure for the memory-bound decode-streaming composite study.
Reads sweep_decode_streaming.json (emitted by milestone-1h-gqa, sweep
``decode_streaming``) and writes one two-panel PNG:
gqa_decode_streaming.png
Left — end-to-end single-rank decode latency (µs) vs per-rank context.
Right — achieved HBM bandwidth (GB/s) vs context, against the
256 GB/s per-rank roofline.
The memory-bound mirror of the compute-bound prefill figure. With T_q=1
the GEMMs are skinny (M=8) and the kernel is bound by streaming the KV
cache. Isolating a single rank (no inter-CUBE reduce) reveals what the
64-way Case-6 decode masks: the composite command still wins, not by
feeding the MAC array but by keeping the DMA pipeline full — its
scheduler-streamed concurrent tile DMAs extract ~230 GB/s (near the
256 GB/s roofline) while the primitive kernel's blocking tl.dot serializes
one tile DMA at a time and plateaus at ~166 GB/s. That bandwidth gap is a
~25-28 % latency win that grows nowhere near prefill's compute-bound
margin but is decidedly not zero.
Run (after the bench):
GQA_1H_RUN=1 GQA_1H_SWEEPS=decode_streaming python -m kernbench.cli.main \\
run --bench milestone-1h-gqa --topology topology.yaml
python scripts/paper/paper_plot_gqa_decode_streaming.py
"""
from __future__ import annotations
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt # noqa: E402
_REPO_ROOT = Path(__file__).resolve().parents[2]
_FIG_DIR = (
_REPO_ROOT / "src" / "kernbench" / "benches"
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _FIG_DIR / "sweep_decode_streaming.json"
_PAPER_FIG_DIR = (
_REPO_ROOT / "docs" / "report" / "1H-codesign-paper" / "figures"
)
# Per-rank HBM roofline: 8 pseudo-channels × 32 GB/s (topology.yaml
# hbm_ctrl.num_pcs / pc_bw_gbs; = pe_dma_to_noc_bw_gbs).
_PEAK_HBM_GBS = 256.0
_VARIANT_STYLE = {
"primitive": ("primitive (tl.dot, hand-tiled)", "#c0504d", "o"),
"composite": ("composite GEMM", "#3b6ea5", "s"),
"composite_extended": ("composite + softmax_merge", "#4f8a4f", "^"),
}
_ORDER = ("primitive", "composite", "composite_extended")
def _ctx_label(c: int) -> str:
return f"{c // 1024}K" if c >= 1024 else str(c)
def _series(rows, variant, key):
pts = sorted(((r["s_kv"], r[key]) for r in rows
if r["variant"] == variant), key=lambda t: t[0])
return [p[0] for p in pts], [p[1] for p in pts]
def main() -> None:
sweep = json.loads(_SWEEP_JSON.read_text())
rows = sweep["rows"]
ctxs = sweep["s_kv_points"]
fig, (ax_lat, ax_bw) = plt.subplots(1, 2, figsize=(13.0, 4.8))
for v in _ORDER:
label, color, marker = _VARIANT_STYLE[v]
xs, lat = _series(rows, v, "latency_ns")
ax_lat.plot(xs, [y / 1e3 for y in lat], marker=marker,
color=color, label=label, lw=2)
xs, bw = _series(rows, v, "achieved_bw_gbs")
ax_bw.plot(xs, bw, marker=marker, color=color, label=label, lw=2)
for ax in (ax_lat, ax_bw):
ax.set_xscale("log", base=2)
ax.set_xticks(ctxs)
ax.set_xticklabels([_ctx_label(c) for c in ctxs])
ax.set_xlabel(r"per-rank context length $S_{kv}$ ($T_q{=}1$)")
ax.grid(True, ls=":", alpha=0.5)
ax.legend(fontsize=9)
ax_lat.set_ylabel("end-to-end decode latency (µs)")
ax_lat.set_title("Single-rank memory-bound decode latency per command form")
ax_bw.set_ylabel("achieved HBM bandwidth (GB/s)")
ax_bw.set_title(
"HBM bandwidth — composite keeps the DMA pipe full; primitive plateaus"
)
ax_bw.axhline(_PEAK_HBM_GBS, color="#888", ls="--", lw=1, alpha=0.7)
ax_bw.text(ctxs[0], _PEAK_HBM_GBS - 8, "256 GB/s roofline",
fontsize=8, color="#555", va="top")
ax_bw.set_ylim(0, _PEAK_HBM_GBS * 1.08)
fig.suptitle(
"Memory-bound decode streaming — use of composite commands\n"
"single-rank, GQA single-KV-head group ($h_q{=}8$, $d_{\\text{head}}"
"{=}128$); $M{=}8$ skinny, KV-streaming-bound",
fontsize=11,
)
fig.tight_layout(rect=(0, 0, 1, 0.92))
out = _FIG_DIR / "gqa_decode_streaming.png"
fig.savefig(out, dpi=150)
plt.close(fig)
print(f"wrote {out}")
if _PAPER_FIG_DIR.is_dir():
dst = _PAPER_FIG_DIR / out.name
dst.write_bytes(out.read_bytes())
print(f"copied {dst}")
if __name__ == "__main__":
main()
@@ -0,0 +1,111 @@
"""Comparative figure for the compute-bound prefill composite study.
Reads sweep_prefill_compute_bound.json (emitted by milestone-1h-gqa,
sweep ``prefill_cb``) and writes one two-panel PNG:
gqa_prefill_compute_bound.png
Left — end-to-end prefill latency (µs) vs context length.
Right — MAC utilization (achieved / 8 TFLOP·s⁻¹ per-PE peak) vs context.
Unlike memory-bound decode (where command form is latency-neutral), in
compute-bound prefill the composite command keeps the MAC array fed by
streaming DMA↔compute per HW tile, so it wins on both latency and
utilization — and the margin grows with context (deeper P·V reduction =
more tiles to pipeline).
Run (after the bench):
GQA_1H_RUN=1 GQA_1H_SWEEPS=prefill_cb python -m kernbench.cli.main run \\
--bench milestone-1h-gqa --topology topology.yaml
python scripts/paper/paper_plot_gqa_prefill_compute_bound.py
"""
from __future__ import annotations
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt # noqa: E402
_REPO_ROOT = Path(__file__).resolve().parents[2]
_FIG_DIR = (
_REPO_ROOT / "src" / "kernbench" / "benches"
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _FIG_DIR / "sweep_prefill_compute_bound.json"
_PAPER_FIG_DIR = (
_REPO_ROOT / "docs" / "report" / "1H-codesign-paper" / "figures"
)
_VARIANT_STYLE = {
"primitive": ("primitive (tl.dot, hand-tiled)", "#c0504d", "o"),
"composite": ("composite GEMM", "#3b6ea5", "s"),
"composite_extended": ("composite + softmax_merge", "#4f8a4f", "^"),
}
_ORDER = ("primitive", "composite", "composite_extended")
def _ctx_label(c: int) -> str:
return f"{c // 1024}K" if c >= 1024 else str(c)
def _series(rows, variant, key):
pts = sorted(((r["ctx_len"], r[key]) for r in rows
if r["variant"] == variant), key=lambda t: t[0])
return [p[0] for p in pts], [p[1] for p in pts]
def main() -> None:
sweep = json.loads(_SWEEP_JSON.read_text())
rows = sweep["rows"]
ctxs = sweep["ctx_points"]
fig, (ax_lat, ax_util) = plt.subplots(1, 2, figsize=(13.0, 4.8))
for v in _ORDER:
label, color, marker = _VARIANT_STYLE[v]
xs, lat = _series(rows, v, "latency_ns")
ax_lat.plot(xs, [y / 1e3 for y in lat], marker=marker,
color=color, label=label, lw=2)
xs, util = _series(rows, v, "mac_util")
ax_util.plot(xs, [u * 100 for u in util], marker=marker,
color=color, label=label, lw=2)
for ax in (ax_lat, ax_util):
ax.set_xscale("log", base=2)
ax.set_xticks(ctxs)
ax.set_xticklabels([_ctx_label(c) for c in ctxs])
ax.set_xlabel(r"context length (= $T_q$ = $S_{kv}$)")
ax.grid(True, ls=":", alpha=0.5)
ax.legend(fontsize=9)
ax_lat.set_ylabel("end-to-end prefill latency (µs)")
ax_lat.set_title("Compute-bound prefill latency per command form")
ax_util.set_ylabel("MAC utilization (% of 8 TFLOP·s⁻¹ peak)")
ax_util.set_title(
"MAC utilization — composite keeps the array fed; primitive starves"
)
ax_util.axhline(100, color="#888", ls="--", lw=1, alpha=0.6)
fig.suptitle(
"Compute-bound prefill attention — use of composite commands\n"
"single-rank, GQA single-KV-head group ($h_q{=}8$, $d_{\\text{head}}"
"{=}128$); $M{=}8T_q$ tile-filling",
fontsize=11,
)
fig.tight_layout(rect=(0, 0, 1, 0.92))
out = _FIG_DIR / "gqa_prefill_compute_bound.png"
fig.savefig(out, dpi=150)
plt.close(fig)
print(f"wrote {out}")
if _PAPER_FIG_DIR.is_dir():
dst = _PAPER_FIG_DIR / out.name
dst.write_bytes(out.read_bytes())
print(f"copied {dst}")
if __name__ == "__main__":
main()
@@ -243,6 +243,12 @@ _FLOW_DATA = "#6699CC" # blue flit (data)
_FLOW_CTRL = "#ED7D65" # orange flit (control)
_FLOW_CRED = "#7BB661" # green (credit / ack)
# Font-size multiplier for the flow helpers. The stacked figure renders at
# full text width (a two-column figure*), so its fonts are scaled up here so
# they stay legible after the figure is fit to the page. Default 1.0 leaves
# the 2×2 flow figure unchanged.
_FS = 1.0
def _flow_swim(ax, x, y, w, h, label, sub=None):
ax.add_patch(mpatches.FancyBboxPatch(
@@ -251,11 +257,11 @@ def _flow_swim(ax, x, y, w, h, label, sub=None):
facecolor="white", edgecolor="#3a3a3a", linewidth=1.6))
ax.text(x + w / 2, y + h / 2 + (0.20 if sub else 0),
label, ha="center", va="center",
fontsize=10.5, fontweight="bold", color="#111")
fontsize=10.5 * _FS, fontweight="bold", color="#111")
if sub:
ax.text(x + w / 2, y + h / 2 - 0.35, sub,
ha="center", va="center",
fontsize=8, color="#555", fontstyle="italic")
fontsize=8 * _FS, color="#555", fontstyle="italic")
def _flow_flit(ax, x, y, color, letter="", size=0.32):
@@ -265,7 +271,7 @@ def _flow_flit(ax, x, y, color, letter="", size=0.32):
facecolor=color, edgecolor="#222", linewidth=0.8))
if letter:
ax.text(x, y, letter, ha="center", va="center",
fontsize=8.5, fontweight="bold", color="white")
fontsize=8.5 * _FS, fontweight="bold", color="white")
def _flow_wire(ax, x1, x2, y, color="#888"):
@@ -280,7 +286,7 @@ def _flow_dot(ax, x, y, color):
def _flow_label(ax, x, y, text, color="#222", fontsize=9):
ax.text(x, y, text, ha="center", va="center",
fontsize=fontsize, color=color, fontweight="bold")
fontsize=fontsize * _FS, color=color, fontweight="bold")
def _flow_lead(ax, x_tail, y_tail, x_head, y_head, color="#777"):
@@ -297,7 +303,7 @@ def _flow_setup(ax, kind):
facecolor="white", edgecolor=_COLOR[kind],
linewidth=0.9, alpha=0.85, zorder=0))
ax.text(10, 0.65, _TITLE[kind],
ha="center", va="center", fontsize=11.5,
ha="center", va="center", fontsize=11.5 * _FS,
fontweight="bold",
color=("#1B5E20" if kind == "ipcq" else "white"),
bbox=dict(
@@ -449,7 +455,7 @@ def _flow_legend(fig, leg_y=0.02):
facecolor=color, edgecolor="#222", linewidth=0.6,
transform=fig.transFigure))
fig.text(lx + 0.025, y + 0.011, text,
ha="left", va="center", fontsize=10,
ha="left", va="center", fontsize=10 * _FS,
fontweight="bold", color="#222")
lx += 0.32
@@ -461,21 +467,27 @@ def _flow_legend(fig, leg_y=0.02):
def _plot_architecture_flow() -> Path:
"""4 panels arranged 2×2 — each a horizontal Sender → NoC → Receiver
flow with small flit-style packets, dot-marker annotations and a
bottom name plate. Aesthetic mirrors latency_model.png."""
fig, axes = plt.subplots(2, 2, figsize=(18.0, 11.5))
_draw_flow_case(axes[0, 0], "doorbell")
_draw_flow_case(axes[0, 1], "hmq")
_draw_flow_case(axes[1, 0], "rdma")
_draw_flow_case(axes[1, 1], "ipcq")
fig.suptitle(
"Per-send architecture = $\\Sigma$ control overhead + "
"data-path DMA + receiver wake",
fontsize=14.5, fontweight="bold", y=0.995)
_flow_legend(fig, leg_y=0.015)
fig.tight_layout(rect=(0, 0.07, 1, 0.96))
out = _OUT_DIR / "ipcq_alternatives_architecture_flow.png"
fig.savefig(out, dpi=150, bbox_inches="tight")
plt.close(fig)
bottom name plate. Aesthetic mirrors latency_model.png. Fonts scaled
up via _FS so the figure stays legible when fit to the page."""
global _FS
_FS = 1.5
try:
fig, axes = plt.subplots(2, 2, figsize=(18.0, 11.5))
_draw_flow_case(axes[0, 0], "doorbell")
_draw_flow_case(axes[0, 1], "hmq")
_draw_flow_case(axes[1, 0], "rdma")
_draw_flow_case(axes[1, 1], "ipcq")
fig.suptitle(
"Per-send architecture = $\\Sigma$ control overhead + "
"data-path DMA + receiver wake",
fontsize=14.5 * _FS, fontweight="bold", y=0.995)
_flow_legend(fig, leg_y=0.015)
fig.tight_layout(rect=(0, 0.07, 1, 0.96))
out = _OUT_DIR / "ipcq_alternatives_architecture_flow.png"
fig.savefig(out, dpi=150, bbox_inches="tight")
plt.close(fig)
finally:
_FS = 1.0
return out
@@ -483,19 +495,25 @@ def _plot_architecture_flow() -> Path:
def _plot_architecture_stacked() -> Path:
"""Same per-case flow content as the 2×2 figure, but stacked one
case per row (4 rows × 1 column). Wider panels give each case more
horizontal room."""
fig, axes = plt.subplots(4, 1, figsize=(16.0, 19.0))
for ax, kind in zip(axes, ["doorbell", "hmq", "rdma", "ipcq"]):
_draw_flow_case(ax, kind)
fig.suptitle(
"Per-send architecture = $\\Sigma$ control overhead + "
"data-path DMA + receiver wake",
fontsize=14.5, fontweight="bold", y=0.997)
_flow_legend(fig, leg_y=0.010)
fig.tight_layout(rect=(0, 0.05, 1, 0.97))
out = _OUT_DIR / "ipcq_alternatives_architecture_stacked.png"
fig.savefig(out, dpi=150, bbox_inches="tight")
plt.close(fig)
horizontal room. Rendered at full text width (two-column figure*), so
fonts are scaled up via _FS to stay legible after page fitting."""
global _FS
_FS = 1.9
try:
fig, axes = plt.subplots(4, 1, figsize=(13.0, 15.0))
for ax, kind in zip(axes, ["doorbell", "hmq", "rdma", "ipcq"]):
_draw_flow_case(ax, kind)
fig.suptitle(
"Per-send architecture = $\\Sigma$ control overhead + "
"data-path DMA + receiver wake",
fontsize=14.5 * _FS, fontweight="bold", y=0.997)
_flow_legend(fig, leg_y=0.010)
fig.tight_layout(rect=(0, 0.06, 1, 0.97))
out = _OUT_DIR / "ipcq_alternatives_architecture_stacked.png"
fig.savefig(out, dpi=150, bbox_inches="tight")
plt.close(fig)
finally:
_FS = 1.0
return out
Binary file not shown.

Before

Width:  |  Height:  |  Size: 253 KiB

After

Width:  |  Height:  |  Size: 252 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 230 KiB

After

Width:  |  Height:  |  Size: 305 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 260 KiB

After

Width:  |  Height:  |  Size: 381 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 153 KiB

After

Width:  |  Height:  |  Size: 153 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 189 KiB

After

Width:  |  Height:  |  Size: 189 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 160 KiB

After

Width:  |  Height:  |  Size: 160 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 87 KiB

After

Width:  |  Height:  |  Size: 86 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 209 KiB

After

Width:  |  Height:  |  Size: 209 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 136 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 136 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 145 KiB

@@ -0,0 +1,258 @@
{
"version": 2,
"variants": [
"primitive",
"composite",
"composite_extended"
],
"s_kv_opcount": [
8192,
65536,
131072,
262144,
524288,
1048576
],
"s_kv_latency": [
8192,
32768,
65536,
131072
],
"rows": [
{
"variant": "primitive",
"S_kv": 8192,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 96,
"pe_cpu_dispatch_cycles": 930,
"latency_ns": 30646.00300000079
},
{
"variant": "composite",
"S_kv": 8192,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 94,
"pe_cpu_dispatch_cycles": 978,
"latency_ns": 30566.1840000007
},
{
"variant": "composite_extended",
"S_kv": 8192,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 98,
"pe_cpu_dispatch_cycles": 1032,
"latency_ns": 30483.359500000362
},
{
"variant": "primitive",
"S_kv": 65536,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 96,
"pe_cpu_dispatch_cycles": 930,
"latency_ns": 231579.37900000165
},
{
"variant": "composite",
"S_kv": 65536,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 94,
"pe_cpu_dispatch_cycles": 978,
"latency_ns": 231270.18400000152
},
{
"variant": "composite_extended",
"S_kv": 65536,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 98,
"pe_cpu_dispatch_cycles": 1032,
"latency_ns": 231211.1740000015
},
{
"variant": "primitive",
"S_kv": 131072,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 118,
"pe_cpu_dispatch_cycles": 1145,
"latency_ns": 461119.51900000183
},
{
"variant": "composite",
"S_kv": 131072,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 94,
"pe_cpu_dispatch_cycles": 978,
"latency_ns": 460651.3170000027
},
{
"variant": "composite_extended",
"S_kv": 131072,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 98,
"pe_cpu_dispatch_cycles": 1032,
"latency_ns": 460552.9805000025
},
{
"variant": "primitive",
"S_kv": 262144,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 162,
"pe_cpu_dispatch_cycles": 1575,
"latency_ns": null
},
{
"variant": "composite",
"S_kv": 262144,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 94,
"pe_cpu_dispatch_cycles": 978,
"latency_ns": null
},
{
"variant": "composite_extended",
"S_kv": 262144,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 98,
"pe_cpu_dispatch_cycles": 1032,
"latency_ns": null
},
{
"variant": "primitive",
"S_kv": 524288,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 250,
"pe_cpu_dispatch_cycles": 2435,
"latency_ns": null
},
{
"variant": "composite",
"S_kv": 524288,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 94,
"pe_cpu_dispatch_cycles": 978,
"latency_ns": null
},
{
"variant": "composite_extended",
"S_kv": 524288,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 98,
"pe_cpu_dispatch_cycles": 1032,
"latency_ns": null
},
{
"variant": "primitive",
"S_kv": 1048576,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 426,
"pe_cpu_dispatch_cycles": 4155,
"latency_ns": null
},
{
"variant": "composite",
"S_kv": 1048576,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 94,
"pe_cpu_dispatch_cycles": 978,
"latency_ns": null
},
{
"variant": "composite_extended",
"S_kv": 1048576,
"C": 8,
"P": 8,
"T_q": 1,
"d_head": 128,
"h_q": 8,
"h_kv": 1,
"pe_cpu_cmd_count": 98,
"pe_cpu_dispatch_cycles": 1032,
"latency_ns": null
}
]
}
@@ -0,0 +1,172 @@
{
"version": 1,
"variants": [
"primitive",
"composite",
"composite_extended"
],
"s_kv_points": [
2048,
4096,
8192,
16384
],
"rows": [
{
"variant": "primitive",
"s_kv": 2048,
"M": 8,
"latency_ns": 6188.437999999816,
"gemm_busy_ns": 1048.576000000001,
"dma_busy_ns": 6322.0,
"kv_bytes": 1048576.0,
"achieved_bw_gbs": 169.44114169036374,
"achieved_tflops": 1.3555291335229098,
"mac_util": 0.16944114169036373,
"dma_occupancy": 1.021582505957107
},
{
"variant": "composite",
"s_kv": 2048,
"M": 8,
"latency_ns": 4836.115999999989,
"gemm_busy_ns": 6629.632000000123,
"dma_busy_ns": 267480.720000013,
"kv_bytes": 1048576.0,
"achieved_bw_gbs": 216.82192900253062,
"achieved_tflops": 1.734575432020245,
"mac_util": 0.2168219290025306,
"dma_occupancy": 55.30899589671001
},
{
"variant": "composite_extended",
"s_kv": 2048,
"M": 8,
"latency_ns": 4737.751999999986,
"gemm_busy_ns": 8481.408000000116,
"dma_busy_ns": 251398.7400000122,
"kv_bytes": 1048576.0,
"achieved_bw_gbs": 221.32353065335693,
"achieved_tflops": 1.7705882452268553,
"mac_util": 0.22132353065335691,
"dma_occupancy": 53.06287454472352
},
{
"variant": "primitive",
"s_kv": 4096,
"M": 8,
"latency_ns": 12510.006000000496,
"gemm_busy_ns": 2097.152000000002,
"dma_busy_ns": 12602.0,
"kv_bytes": 2097152.0,
"achieved_bw_gbs": 167.6379691584414,
"achieved_tflops": 1.3411037532675312,
"mac_util": 0.1676379691584414,
"dma_occupancy": 1.0073536335633653
},
{
"variant": "composite",
"s_kv": 4096,
"M": 8,
"latency_ns": 9360.883999999554,
"gemm_busy_ns": 19520.799999939867,
"dma_busy_ns": 1059043.5999999335,
"kv_bytes": 2097152.0,
"achieved_bw_gbs": 224.03354213128802,
"achieved_tflops": 1.792268337050304,
"mac_util": 0.224033542131288,
"dma_occupancy": 113.13499878857424
},
{
"variant": "composite_extended",
"s_kv": 4096,
"M": 8,
"latency_ns": 9198.135999999562,
"gemm_busy_ns": 39531.583999941715,
"dma_busy_ns": 1026582.7399999356,
"kv_bytes": 2097152.0,
"achieved_bw_gbs": 227.99749862364504,
"achieved_tflops": 1.8239799889891604,
"mac_util": 0.22799749862364505,
"dma_occupancy": 111.60769312390951
},
{
"variant": "primitive",
"s_kv": 8192,
"M": 8,
"latency_ns": 25153.141999997388,
"gemm_busy_ns": 4194.304000000004,
"dma_busy_ns": 25162.0,
"kv_bytes": 4194304.0,
"achieved_bw_gbs": 166.7506985807354,
"achieved_tflops": 1.334005588645883,
"mac_util": 0.16675069858073538,
"dma_occupancy": 1.0003521627637062
},
{
"variant": "composite",
"s_kv": 8192,
"M": 8,
"latency_ns": 18419.187999999034,
"gemm_busy_ns": 64177.11999976149,
"dma_busy_ns": 4214541.840000685,
"kv_bytes": 4194304.0,
"achieved_bw_gbs": 227.713838416776,
"achieved_tflops": 1.821710707334208,
"mac_util": 0.227713838416776,
"dma_occupancy": 228.81257523409317
},
{
"variant": "composite_extended",
"s_kv": 8192,
"M": 8,
"latency_ns": 18128.439999999027,
"gemm_busy_ns": 169650.68799976518,
"dma_busy_ns": 4149323.2200006745,
"kv_bytes": 4194304.0,
"achieved_bw_gbs": 231.365964197704,
"achieved_tflops": 1.850927713581632,
"mac_util": 0.231365964197704,
"dma_occupancy": 228.88473691067168
},
{
"variant": "primitive",
"s_kv": 16384,
"M": 8,
"latency_ns": 50439.414000008954,
"gemm_busy_ns": 8388.607999999076,
"dma_busy_ns": 50282.0,
"kv_bytes": 8388608.0,
"achieved_bw_gbs": 166.31057609032712,
"achieved_tflops": 1.330484608722617,
"mac_util": 0.16631057609032712,
"dma_occupancy": 0.9968791469304357
},
{
"variant": "composite",
"s_kv": 16384,
"M": 8,
"latency_ns": 36535.79600000568,
"gemm_busy_ns": 228978.46400288073,
"dma_busy_ns": 16815028.239995122,
"kv_bytes": 8388608.0,
"achieved_bw_gbs": 229.59970545047648,
"achieved_tflops": 1.8367976436038118,
"mac_util": 0.22959970545047648,
"dma_occupancy": 460.2343477063565
},
{
"variant": "composite_extended",
"s_kv": 16384,
"M": 8,
"latency_ns": 35989.048000005685,
"gemm_busy_ns": 701990.3680028584,
"dma_busy_ns": 16684294.09999516,
"kv_bytes": 8388608.0,
"achieved_bw_gbs": 233.0877993771515,
"achieved_tflops": 1.864702395017212,
"mac_util": 0.2330877993771515,
"dma_occupancy": 463.5936493789034
}
]
}
@@ -0,0 +1,105 @@
{
"version": 1,
"variants": [
"primitive",
"composite",
"composite_extended"
],
"ctx_points": [
256,
512,
1024
],
"rows": [
{
"variant": "primitive",
"ctx_len": 256,
"M": 2048,
"latency_ns": 48857.361999999725,
"gemm_busy_ns": 33554.43200000003,
"dma_busy_ns": 20880.000000000004,
"achieved_tflops": 5.4942683151825005,
"mac_util": 0.6867835393978126
},
{
"variant": "composite",
"ctx_len": 256,
"M": 2048,
"latency_ns": 49887.505999999594,
"gemm_busy_ns": 34531.32799999314,
"dma_busy_ns": 1095773.440000072,
"achieved_tflops": 5.380815308746887,
"mac_util": 0.6726019135933609
},
{
"variant": "composite_extended",
"ctx_len": 256,
"M": 2048,
"latency_ns": 50280.91399999841,
"gemm_busy_ns": 129096.19199996963,
"dma_busy_ns": 633563.5200000411,
"achieved_tflops": 5.338714725830331,
"mac_util": 0.6673393407287914
},
{
"variant": "primitive",
"ctx_len": 512,
"M": 4096,
"latency_ns": 197496.81800000605,
"gemm_busy_ns": 134217.72800000012,
"dma_busy_ns": 66880.0,
"achieved_tflops": 5.436755056985105,
"mac_util": 0.6795943821231382
},
{
"variant": "composite",
"ctx_len": 512,
"M": 4096,
"latency_ns": 177341.23400000148,
"gemm_busy_ns": 141168.63999995415,
"dma_busy_ns": 8567063.039998509,
"achieved_tflops": 6.054665346469796,
"mac_util": 0.7568331683087245
},
{
"variant": "composite_extended",
"ctx_len": 512,
"M": 4096,
"latency_ns": 174702.7699999963,
"gemm_busy_ns": 1076198.3999996716,
"dma_busy_ns": 6594394.87999886,
"achieved_tflops": 6.146106464139193,
"mac_util": 0.7682633080173992
},
{
"variant": "primitive",
"ctx_len": 1024,
"M": 8192,
"latency_ns": 794090.4820000405,
"gemm_busy_ns": 536870.9120000004,
"dma_busy_ns": 234240.0,
"achieved_tflops": 5.4086623544239645,
"mac_util": 0.6760827943029956
},
{
"variant": "composite",
"ctx_len": 1024,
"M": 8192,
"latency_ns": 668110.7059999453,
"gemm_busy_ns": 873011.1999923651,
"dma_busy_ns": 67794150.3999821,
"achieved_tflops": 6.428526376585786,
"mac_util": 0.8035657970732233
},
{
"variant": "composite_extended",
"ctx_len": 1024,
"M": 8192,
"latency_ns": 646937.2019999702,
"gemm_busy_ns": 8870907.903995967,
"dma_busy_ns": 59655820.79998447,
"achieved_tflops": 6.638924586068552,
"mac_util": 0.829865573258569
}
]
}
@@ -36,26 +36,14 @@ Topology / SFR:
"""
from __future__ import annotations
from kernbench.benches.gqa_helpers.long_ctx._gqa_mlo_reduce import (
_ROOT_CUBE,
_merge_running,
reduce_mlo,
)
TILE_S_KV = 1024 # ADR-0063 §A.2 S_kv-axis tile sweep (per-tile width).
# lrab geometry for the C=8 single-KV-head group (4×2 cube sub-mesh).
_SUB_W = 4
_SUB_H = 2
_ROOT_COL = _SUB_W // 2 # 2
_ROOT_ROW = _SUB_H // 2 # 1
_ROOT_CUBE = _ROOT_ROW * _SUB_W + _ROOT_COL # 6
def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl):
"""Online-softmax merge of two partial ``(m, , O)`` triples."""
m_new = tl.maximum(m_local, m_other)
scale_old = tl.exp(m_local - m_new)
scale_new = tl.exp(m_other - m_new)
l_new = l_local * scale_old + l_other * scale_new
O_new = O_local * scale_old + O_other * scale_new
return m_new, l_new, O_new
def gqa_attention_decode_long_ctx_cube_sp_pe_sp_kernel(
q_ptr: int,
@@ -119,175 +107,8 @@ def gqa_attention_decode_long_ctx_cube_sp_pe_sp_kernel(
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
# ── Intra-CUBE reduce: row chain (intra_W) + col bridge (intra_N) ──
PE_GRID_COLS = 4
pe_col = pe_id % PE_GRID_COLS
pe_row = pe_id // PE_GRID_COLS
pe_cols_used = min(PE_GRID_COLS, P)
pe_rows_used = (P + PE_GRID_COLS - 1) // PE_GRID_COLS
if pe_cols_used > 1:
if pe_col < pe_cols_used - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_E", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="intra_E", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_col > 0:
tl.send(dir="intra_W", src=m_local)
tl.send(dir="intra_W", src=l_local)
tl.send(dir="intra_W", src=O_local)
if pe_col == 0 and pe_rows_used > 1:
if pe_row < pe_rows_used - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_S", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="intra_S", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_row > 0:
tl.send(dir="intra_N", src=m_local)
tl.send(dir="intra_N", src=l_local)
tl.send(dir="intra_N", src=O_local)
# ── Inter-CUBE lrab-adapted center-root reduce (ADR-0060 §4.2) ──
# Only PE 0 of each CUBE participates. Adapts Phases 1-2 of
# lrab_hierarchical_allreduce.py: bidirectional row reduce converges
# at root_col; bidirectional col reduce on root_col converges at
# root_row. Plain ``+`` replaced by log-sum-exp ``_merge_running``.
if pe_id == 0:
row = cube_id // _SUB_W
col = cube_id % _SUB_W
# Phase 1: row reduce — converge at col == _ROOT_COL.
if col == 0:
tl.send(dir="E", src=m_local)
tl.send(dir="E", src=l_local)
tl.send(dir="E", src=O_local)
elif 0 < col < _ROOT_COL:
with tl.scratch_scope():
m_other = tl.recv(dir="W", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="W", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="W", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="E", src=m_local)
tl.send(dir="E", src=l_local)
tl.send(dir="E", src=O_local)
elif col == _ROOT_COL:
with tl.scratch_scope():
m_other = tl.recv(dir="W", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="W", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="W", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
with tl.scratch_scope():
m_other = tl.recv(dir="E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="E", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="E", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
elif _ROOT_COL < col < _SUB_W - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="E", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="E", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="W", src=m_local)
tl.send(dir="W", src=l_local)
tl.send(dir="W", src=O_local)
elif col == _SUB_W - 1:
tl.send(dir="W", src=m_local)
tl.send(dir="W", src=l_local)
tl.send(dir="W", src=O_local)
# Phase 2: col reduce on col == _ROOT_COL — converge at row == _ROOT_ROW.
if col == _ROOT_COL:
if row == 0:
tl.send(dir="S", src=m_local)
tl.send(dir="S", src=l_local)
tl.send(dir="S", src=O_local)
elif 0 < row < _ROOT_ROW:
with tl.scratch_scope():
m_other = tl.recv(dir="N", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="N", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="N", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="S", src=m_local)
tl.send(dir="S", src=l_local)
tl.send(dir="S", src=O_local)
elif row == _ROOT_ROW:
with tl.scratch_scope():
m_other = tl.recv(dir="N", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="N", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="N", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if _SUB_H - 1 > _ROOT_ROW:
with tl.scratch_scope():
m_other = tl.recv(dir="S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="S", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="S", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
elif _ROOT_ROW < row < _SUB_H - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="S", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="S", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="N", src=m_local)
tl.send(dir="N", src=l_local)
tl.send(dir="N", src=O_local)
elif row == _SUB_H - 1 and _SUB_H - 1 > _ROOT_ROW:
tl.send(dir="N", src=m_local)
tl.send(dir="N", src=l_local)
tl.send(dir="N", src=O_local)
# ── Two-level (m, , O) reduce-to-root (shared helper) ──
reduce_mlo(pe_id, cube_id, m_local, l_local, O_local, P, tl=tl)
# ── Final normalise + store (root only: PE 0 of CUBE 6) ──
if pe_id == 0 and cube_id == _ROOT_CUBE:
@@ -0,0 +1,71 @@
"""GQA decode kernel — Case 6, **composite-GEMM** command form.
Identical placement and (m, , O) reduce as the primitive Case-6 kernel
(``_gqa_attention_decode_long_ctx_cube_sp_pe_sp``); the difference is the
command *granularity* of the local attention. The primitive kernel walks
its KV slice in ``TILE_S_KV``-wide tiles, issuing per-tile loads + GEMMs
+ a manual online-softmax merge — O(n_tiles) PE_CPU commands. This kernel
instead issues **one coarse composite GEMM over the whole ``S_local``**
for each matrix product, passing K and V as HBM refs so PE_SCHEDULER
streams and tiles them (the DMA fan-out moves off PE_CPU). The
online-softmax stays primitive but now runs once over the full score row.
So the kernel issues O(1) coarse commands; the scheduler expands each
composite into the same per-tile DMA + MAC work the primitive kernel
issued by hand. Fewer, coarser PE_CPU commands ⇒ lower dispatch cost
under the ADR-0064 Rev2 structural model (the CPU-offload win).
"""
from __future__ import annotations
from kernbench.benches.gqa_helpers.long_ctx._gqa_mlo_reduce import (
_ROOT_CUBE,
reduce_mlo,
)
def gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_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:
"""Case-6 decode (Cube-SP × PE-SP) — coarse composite-GEMM command form."""
G = h_q // h_kv
n_ranks = C * P
S_local = S_kv // n_ranks
pe_id = tl.program_id(axis=0)
cube_id = tl.program_id(axis=1)
# ── Local attention as two coarse composite GEMMs ──
# K, V are HBM refs: PE_SCHEDULER streams + tiles them over S_local
# (the per-tile DMA fan-out the primitive kernel issued by hand).
Q = tl.load(q_ptr, shape=(G * T_q, d_head), dtype="f16")
K_T = tl.ref(k_ptr, shape=(d_head, S_local), dtype="f16")
V = tl.ref(v_ptr, shape=(S_local, d_head), dtype="f16")
scores = tl.composite(op="gemm", a=Q, b=K_T) # Q·Kᵀ, one command
m_local = tl.max(scores, axis=-1)
centered = scores - m_local
exp_scores = tl.exp(centered)
l_local = tl.sum(exp_scores, axis=-1)
# P·V into a pre-allocated TCM output (explicit out — the composite
# output must be a real TCM handle, not auto-allocated scratch).
O_local = tl.zeros((G * T_q, d_head), dtype="f16")
tl.composite(op="gemm", a=exp_scores, b=V, out=O_local) # P·V, one command
# ── Two-level (m, , O) reduce-to-root (shared helper) ──
reduce_mlo(pe_id, cube_id, m_local, l_local, O_local, P, tl=tl)
# ── Final normalise + store (root only: PE 0 of CUBE 6) ──
if pe_id == 0 and cube_id == _ROOT_CUBE:
O_final = O_local / l_local
tl.store(o_ptr, O_final)
@@ -0,0 +1,94 @@
"""GQA decode kernel — Case 6, **composite_extended** command form.
Same placement and (m, , O) reduce as the primitive / composite-GEMM
Case-6 kernels; the local attention is the opt2 **two-composite** form
(ADR-0060 §8 item 4 / ADR-0065), issued coarsely:
establish a small first KV slice computes the running (m, , O) with
primitives (kernbench has no scratch-backed ``-inf``
initializer, so the recipe — which *merges* into an existing
accumulator — needs a seed).
#1 Q·Kᵀ one coarse composite GEMM over the remaining S_local
(K passed as an HBM ref → PE_SCHEDULER streams + tiles it).
#2 softmax → one ``softmax_merge`` recipe composite (online-softmax
+ P·V merge of (m, , O)) whose head GEMM is P·V over the same
remaining slice (V as an HBM ref), with an ``add``
epilogue folding the P·V contribution into ``O``.
So the bulk of the KV slice is two coarse PE_CPU commands, and the recipe
additionally folds the per-tile online merge + P·V into one descriptor —
the fewest / cheapest commands of the three forms (the headline
CPU-offload win, ADR-0064 Rev2).
Scope: runnable in op_log mode (latency / dispatch only). Full data-mode
numeric parity of the recipe's 8 MATH ops is a separate follow-up
(DDD-0065 / the P5-numerics note).
"""
from __future__ import annotations
from kernbench.benches.gqa_helpers.long_ctx._gqa_mlo_reduce import (
_ROOT_CUBE,
reduce_mlo,
)
# Small primitive seed slice that establishes the running (m, , O) the
# recipe then merges into. One scheduler K-tile wide (ADR-0064 TILE_K).
_SEED_S_KV = 64
def gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext_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:
"""Case-6 decode (Cube-SP × PE-SP) — softmax_merge recipe command form."""
G = h_q // h_kv
n_ranks = C * P
S_local = S_kv // n_ranks
pe_id = tl.program_id(axis=0)
cube_id = tl.program_id(axis=1)
KV_ROW_BYTES = d_head * 2 # f16
# ── Establish running (m, , O) on a small seed slice (primitive) ──
Q = tl.load(q_ptr, shape=(G * T_q, d_head), dtype="f16")
seed = min(_SEED_S_KV, S_local)
K_T0 = tl.load(k_ptr, shape=(d_head, seed), dtype="f16")
V0 = tl.load(v_ptr, shape=(seed, 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)
# ── Remaining slice: one Q·Kᵀ composite + one softmax_merge recipe ──
rest = S_local - seed
if rest > 0:
K_T1 = tl.ref(k_ptr + seed * KV_ROW_BYTES,
shape=(d_head, rest), dtype="f16")
V1 = tl.ref(v_ptr + seed * KV_ROW_BYTES,
shape=(rest, d_head), dtype="f16")
scores1 = tl.composite(op="gemm", a=Q, b=K_T1)
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}],
)
# ── Two-level (m, , O) reduce-to-root (shared helper) ──
reduce_mlo(pe_id, cube_id, m_local, l_local, O_local, P, tl=tl)
# ── Final normalise + store (root only: PE 0 of CUBE 6) ──
if pe_id == 0 and cube_id == _ROOT_CUBE:
O_final = O_local / l_local
tl.store(o_ptr, O_final)
@@ -0,0 +1,129 @@
"""GQA decode kernel — Case 6, **primitive-TILED** (16×16×16 MAC blocking).
Same Case-6 placement and (m, , O) reduce as the primitive baseline
(``_gqa_attention_decode_long_ctx_cube_sp_pe_sp``); the only difference is
that each local-attention matmul is *hand-blocked into 16×16×16 GemmCmds*
(mac=16) instead of one coarse ``tl.dot`` per tile. This models a kernel
that issues the MAC-array fan-out from PE_CPU itself: the per-block GemmCmd
count is ``ceil(M/16)·ceil(K/16)·ceil(N/16)`` per matmul, charging the full
ADR-0064 dispatch cost for every block — the "dispatch explosion" the
composite form offloads to PE_SCHEDULER.
FLOPs are conserved (each 16³ GemmCmd carries the TFLOPS-model compute of
its block; the blocks sum to the full matmul), so end-to-end compute time
is unchanged vs the coarse primitive — only the PE_CPU command count and
its dispatch cycles grow. Inputs are zero (decode bench convention), so the
blocked accumulation is identically zero; the kernel returns a single
zeroed ``(M, N)`` output handle that the downstream softmax consumes — no
per-block accumulation handle needed for this zero-input study.
"""
from __future__ import annotations
from math import ceil
from kernbench.common.pe_commands import GemmCmd
from kernbench.benches.gqa_helpers.long_ctx._gqa_mlo_reduce import (
_ROOT_CUBE,
_merge_running,
reduce_mlo,
)
TILE_S_KV = 1024 # ADR-0063 §A.2 S_kv-axis tile sweep (per-tile width).
MAC = 16 # 16×16×16 MAC-array blocking granularity.
def _blocked_dot(A, B, *, tl):
"""``A @ B`` issued as one GemmCmd per 16×16×16 block.
``A``:(M, K), ``B``:(K, N) → out:(M, N). Emits
``ceil(M/16)·ceil(K/16)·ceil(N/16)`` GemmCmds (each charged the full
ADR-0064 dispatch cost via ``tl.dot``'s emit path). **All blocks write
the same ``(M, N)`` output handle** (``out``), and that handle is
returned — so downstream softmax ops depend on it exactly like the
coarse ``tl.dot`` path (the engine tracks the producer by output handle
id, and ``GemmCmd`` is a blocking PE_CPU command, so the block GEMMs
serialize on the critical path ahead of the consuming ``tl.max``/
``tl.exp``). K is innermost so each (mi, ni) output tile accumulates
across the K blocks into the same handle.
"""
M, K = A.shape[-2], A.shape[-1]
K2, N = B.shape[-2], B.shape[-1]
assert K == K2, f"blocked_dot shape mismatch K={K} != {K2}"
out = tl._make_compute_out(shape=(M, N), dtype="f16")
# One GemmCmd per 16³ block, all writing the shared `out` handle. A/B
# reuse the full operand handles at block dims (m,k,n = 16, or the
# ragged tail) — operands are TCM-resident (pinned), so this charges
# dispatch + the block's TFLOPS-model compute. K innermost = accumulate
# the K blocks into each (mi, ni) output tile of the shared handle.
for _mi in range(ceil(M / MAC)):
bm = min(MAC, M - _mi * MAC)
for _ni in range(ceil(N / MAC)):
bn = min(MAC, N - _ni * MAC)
for _ki in range(ceil(K / MAC)):
bk = min(MAC, K - _ki * MAC)
tl._emit(GemmCmd(a=A, b=B, out=out, m=bm, k=bk, n=bn))
return out
def gqa_attention_decode_long_ctx_cube_sp_pe_sp_tiled_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:
"""Case-6 decode, primitive-TILED (16×16×16 GemmCmd blocking)."""
G = h_q // h_kv
n_ranks = C * P
S_local = S_kv // n_ranks
pe_id = tl.program_id(axis=0)
cube_id = tl.program_id(axis=1)
Q = tl.load(q_ptr, shape=(G * T_q, d_head), dtype="f16")
n_tiles = (S_local + TILE_S_KV - 1) // TILE_S_KV
KV_ROW_BYTES = d_head * 2 # f16
tile_s0 = min(TILE_S_KV, S_local)
K_T = tl.load(k_ptr, shape=(d_head, tile_s0), dtype="f16")
V = tl.load(v_ptr, shape=(tile_s0, d_head), dtype="f16")
scores = _blocked_dot(Q, K_T, tl=tl)
m_local = tl.max(scores, axis=-1)
centered = scores - m_local
exp_scores = tl.exp(centered)
l_local = tl.sum(exp_scores, axis=-1)
O_local = _blocked_dot(exp_scores, V, tl=tl)
for tile_idx in range(1, n_tiles):
tile_start = tile_idx * TILE_S_KV
tile_s = min(TILE_S_KV, S_local - tile_start)
with tl.scratch_scope():
K_T_t = tl.load(k_ptr + tile_start * KV_ROW_BYTES,
shape=(d_head, tile_s), dtype="f16")
V_t = tl.load(v_ptr + tile_start * KV_ROW_BYTES,
shape=(tile_s, d_head), dtype="f16")
scores_t = _blocked_dot(Q, K_T_t, tl=tl)
m_tile = tl.max(scores_t, axis=-1)
centered_t = scores_t - m_tile
exp_scores_t = tl.exp(centered_t)
l_tile = tl.sum(exp_scores_t, axis=-1)
O_tile = _blocked_dot(exp_scores_t, V_t, tl=tl)
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_tile, l_tile, O_tile, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
reduce_mlo(pe_id, cube_id, m_local, l_local, O_local, P, tl=tl)
if pe_id == 0 and cube_id == _ROOT_CUBE:
O_final = O_local / l_local
tl.store(o_ptr, O_final)
@@ -0,0 +1,214 @@
"""Shared (m, , O) reduce for the Case-6 long-context decode kernels.
Extracted from the Cube-SP × PE-SP decode kernel so the three
command-form variants (primitive / composite / composite_extended)
share one identical reduce. The local attention differs per variant;
the reduce does not. Behavior is byte-equal to the pre-extraction inline
reduce (guarded by tests/attention/test_gqa_decode_long_ctx_composite.py).
The reduce is two-level:
• intra-CUBE 8-way (row chain along intra_W + col bridge along
intra_N) over the 2×4 PE grid;
• inter-CUBE 2-phase lrab over the 4×2 CUBE sub-mesh (ADR-0060 §4.2
lrab-adapted center-root reduce): Phase 1 row reduce converges at
root_col=2; Phase 2 col reduce on root_col converges at root_row=1;
root cube id = 6.
Plain ``+`` is replaced by log-sum-exp ``_merge_running``.
"""
from __future__ import annotations
# lrab geometry for the C=8 single-KV-head group (4×2 cube sub-mesh).
_SUB_W = 4
_SUB_H = 2
_ROOT_COL = _SUB_W // 2 # 2
_ROOT_ROW = _SUB_H // 2 # 1
_ROOT_CUBE = _ROOT_ROW * _SUB_W + _ROOT_COL # 6
PE_GRID_COLS = 4
def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl):
"""Online-softmax merge of two partial ``(m, , O)`` triples."""
m_new = tl.maximum(m_local, m_other)
scale_old = tl.exp(m_local - m_new)
scale_new = tl.exp(m_other - m_new)
l_new = l_local * scale_old + l_other * scale_new
O_new = O_local * scale_old + O_other * scale_new
return m_new, l_new, O_new
def reduce_mlo(pe_id, cube_id, m_local, l_local, O_local, P, *, tl):
"""Two-level (m, , O) reduce-to-root (PE 0 of CUBE 6). In place.
Updates ``m_local``/``l_local``/``O_local`` in place via ``copy_to``;
after this call only PE 0 of CUBE ``_ROOT_CUBE`` holds the full result.
"""
# ── Intra-CUBE reduce: row chain (intra_W) + col bridge (intra_N) ──
pe_col = pe_id % PE_GRID_COLS
pe_row = pe_id // PE_GRID_COLS
pe_cols_used = min(PE_GRID_COLS, P)
pe_rows_used = (P + PE_GRID_COLS - 1) // PE_GRID_COLS
if pe_cols_used > 1:
if pe_col < pe_cols_used - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_E", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="intra_E", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_col > 0:
tl.send(dir="intra_W", src=m_local)
tl.send(dir="intra_W", src=l_local)
tl.send(dir="intra_W", src=O_local)
if pe_col == 0 and pe_rows_used > 1:
if pe_row < pe_rows_used - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_S", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="intra_S", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_row > 0:
tl.send(dir="intra_N", src=m_local)
tl.send(dir="intra_N", src=l_local)
tl.send(dir="intra_N", src=O_local)
# ── Inter-CUBE lrab-adapted center-root reduce (ADR-0060 §4.2) ──
# Only PE 0 of each CUBE participates. Adapts Phases 1-2 of
# lrab_hierarchical_allreduce.py: bidirectional row reduce converges
# at root_col; bidirectional col reduce on root_col converges at
# root_row. Plain ``+`` replaced by log-sum-exp ``_merge_running``.
if pe_id == 0:
row = cube_id // _SUB_W
col = cube_id % _SUB_W
# Phase 1: row reduce — converge at col == _ROOT_COL.
if col == 0:
tl.send(dir="E", src=m_local)
tl.send(dir="E", src=l_local)
tl.send(dir="E", src=O_local)
elif 0 < col < _ROOT_COL:
with tl.scratch_scope():
m_other = tl.recv(dir="W", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="W", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="W", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="E", src=m_local)
tl.send(dir="E", src=l_local)
tl.send(dir="E", src=O_local)
elif col == _ROOT_COL:
with tl.scratch_scope():
m_other = tl.recv(dir="W", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="W", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="W", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
with tl.scratch_scope():
m_other = tl.recv(dir="E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="E", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="E", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
elif _ROOT_COL < col < _SUB_W - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="E", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="E", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="W", src=m_local)
tl.send(dir="W", src=l_local)
tl.send(dir="W", src=O_local)
elif col == _SUB_W - 1:
tl.send(dir="W", src=m_local)
tl.send(dir="W", src=l_local)
tl.send(dir="W", src=O_local)
# Phase 2: col reduce on col == _ROOT_COL — converge at row == _ROOT_ROW.
if col == _ROOT_COL:
if row == 0:
tl.send(dir="S", src=m_local)
tl.send(dir="S", src=l_local)
tl.send(dir="S", src=O_local)
elif 0 < row < _ROOT_ROW:
with tl.scratch_scope():
m_other = tl.recv(dir="N", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="N", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="N", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="S", src=m_local)
tl.send(dir="S", src=l_local)
tl.send(dir="S", src=O_local)
elif row == _ROOT_ROW:
with tl.scratch_scope():
m_other = tl.recv(dir="N", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="N", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="N", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if _SUB_H - 1 > _ROOT_ROW:
with tl.scratch_scope():
m_other = tl.recv(dir="S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="S", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="S", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
elif _ROOT_ROW < row < _SUB_H - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="S", shape=l_local.shape, dtype="f16")
O_other = tl.recv(dir="S", shape=O_local.shape, dtype="f16")
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_other, l_other, O_other, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
tl.send(dir="N", src=m_local)
tl.send(dir="N", src=l_local)
tl.send(dir="N", src=O_local)
elif row == _SUB_H - 1 and _SUB_H - 1 > _ROOT_ROW:
tl.send(dir="N", src=m_local)
tl.send(dir="N", src=l_local)
tl.send(dir="N", src=O_local)
@@ -0,0 +1,161 @@
"""Compute-bound prefill attention — 3 command-form variants (single-rank).
Companion to the memory-bound decode study
(``_gqa_attention_decode_long_ctx_cube_sp_pe_sp*``). Prefill processes a
block of T_q query positions at once, so the score / context GEMMs have a
large M = G·T_q and high arithmetic intensity (~M flops/byte) — the
workload is **compute-bound** (above the roofline ridge), unlike T_q=1
decode (M=8, memory-bound). This is the regime where the composite
command's value shows: it streams DMA↔compute per HW tile to keep the MAC
array fed, while the primitive kernel serializes load→dot and starves it.
Single-rank (C=P=1): no cross-device reduce — the focus is the per-PE
GEMM-issue mechanism. FlashAttention 2-D tiling (Q-block × S_kv-tile,
online softmax) bounds the TCM scratch.
Three forms, differing only in how each Q-block's local attention is
issued (placement/softmax identical):
primitive tl.dot per S_kv-tile + primitive online merge.
composite one coarse composite GEMM per Q-block over the whole
S_kv (K, V as HBM refs → PE_SCHEDULER tiles/streams).
composite_extended Q·Kᵀ composite + softmax_merge recipe composite.
"""
from __future__ import annotations
from kernbench.benches.gqa_helpers.long_ctx._gqa_mlo_reduce import _merge_running
M_BLOCK = 128 # query rows per FlashAttention Q-block (tile-filling).
TILE_S_KV = 256 # primitive per-tile S_kv width.
_SEED_S_KV = 64 # composite_extended recipe seed slice.
def _qblock_bounds(T_q: int, h_q: int, h_kv: int):
G = h_q // h_kv
M_total = G * T_q
n_qblocks = (M_total + M_BLOCK - 1) // M_BLOCK
return G, M_total, n_qblocks
# ── primitive: hand-tiled monolithic dots + online merge ─────────────
def gqa_prefill_primitive_kernel(
q_ptr, k_ptr, v_ptr, o_ptr,
T_q, S_kv, h_q, h_kv, d_head, C, P, *, tl,
) -> None:
G, M_total, n_qblocks = _qblock_bounds(T_q, h_q, h_kv)
ROW = d_head * 2 # f16
n_tiles = (S_kv + TILE_S_KV - 1) // TILE_S_KV
for qb in range(n_qblocks):
# Per-Q-block scratch_scope: this block's running (m,,O) plus its
# tile-0 transients are freed before the next block, so scratch
# stays O(one block) rather than accumulating across all blocks.
with tl.scratch_scope():
m_blk = min(M_BLOCK, M_total - qb * M_BLOCK)
q_off = qb * M_BLOCK * ROW
Q = tl.load(q_ptr + q_off, shape=(m_blk, d_head), dtype="f16")
tile_s0 = min(TILE_S_KV, S_kv)
K_T = tl.load(k_ptr, shape=(d_head, tile_s0), dtype="f16")
V = tl.load(v_ptr, shape=(tile_s0, d_head), dtype="f16")
scores = tl.dot(Q, K_T)
m_local = tl.max(scores, axis=-1)
exp_s = tl.exp(scores - m_local)
l_local = tl.sum(exp_s, axis=-1)
O_local = tl.dot(exp_s, V)
for ti in range(1, n_tiles):
start = ti * TILE_S_KV
tile_s = min(TILE_S_KV, S_kv - start)
with tl.scratch_scope():
K_T_t = tl.load(k_ptr + start * ROW,
shape=(d_head, tile_s), dtype="f16")
V_t = tl.load(v_ptr + start * ROW,
shape=(tile_s, d_head), dtype="f16")
scores_t = tl.dot(Q, K_T_t)
m_tile = tl.max(scores_t, axis=-1)
exp_t = tl.exp(scores_t - m_tile)
l_tile = tl.sum(exp_t, axis=-1)
O_tile = tl.dot(exp_t, V_t)
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_tile, l_tile, O_tile, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
O_final = O_local / l_local
tl.store(o_ptr + q_off, O_final)
# ── composite: one coarse composite GEMM per Q-block ─────────────────
def gqa_prefill_composite_kernel(
q_ptr, k_ptr, v_ptr, o_ptr,
T_q, S_kv, h_q, h_kv, d_head, C, P, *, tl,
) -> None:
G, M_total, n_qblocks = _qblock_bounds(T_q, h_q, h_kv)
ROW = d_head * 2
for qb in range(n_qblocks):
with tl.scratch_scope():
m_blk = min(M_BLOCK, M_total - qb * M_BLOCK)
q_off = qb * M_BLOCK * ROW
Q = tl.load(q_ptr + q_off, shape=(m_blk, d_head), dtype="f16")
K_T = tl.ref(k_ptr, shape=(d_head, S_kv), dtype="f16")
V = tl.ref(v_ptr, shape=(S_kv, d_head), dtype="f16")
scores = tl.composite(op="gemm", a=Q, b=K_T) # Q·Kᵀ, one command
m_local = tl.max(scores, axis=-1)
exp_s = tl.exp(scores - m_local)
l_local = tl.sum(exp_s, axis=-1)
O_local = tl.zeros((m_blk, d_head), dtype="f16")
tl.composite(op="gemm", a=exp_s, b=V, out=O_local) # P·V, one command
O_final = O_local / l_local
tl.store(o_ptr + q_off, O_final)
# ── composite_extended: Q·Kᵀ composite + softmax_merge recipe ────────
def gqa_prefill_composite_ext_kernel(
q_ptr, k_ptr, v_ptr, o_ptr,
T_q, S_kv, h_q, h_kv, d_head, C, P, *, tl,
) -> None:
G, M_total, n_qblocks = _qblock_bounds(T_q, h_q, h_kv)
ROW = d_head * 2
for qb in range(n_qblocks):
with tl.scratch_scope():
m_blk = min(M_BLOCK, M_total - qb * M_BLOCK)
q_off = qb * M_BLOCK * ROW
Q = tl.load(q_ptr + q_off, shape=(m_blk, d_head), dtype="f16")
seed = min(_SEED_S_KV, S_kv)
K_T0 = tl.load(k_ptr, shape=(d_head, seed), dtype="f16")
V0 = tl.load(v_ptr, shape=(seed, 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)
rest = S_kv - seed
if rest > 0:
K_T1 = tl.ref(k_ptr + seed * ROW,
shape=(d_head, rest), dtype="f16")
V1 = tl.ref(v_ptr + seed * ROW,
shape=(rest, d_head), dtype="f16")
scores1 = tl.composite(op="gemm", a=Q, b=K_T1)
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}],
)
O_final = O_local / l_local
tl.store(o_ptr + q_off, O_final)
@@ -0,0 +1,182 @@
"""milestone-1h-gqa decode: Case-6 composite-command study.
Three command-form variants of the Case-6 (Cube-SP × PE-SP) long-context
decode kernel, swept over S_kv so the local-attention tile loop runs
multiple tiles (S_local = S_kv/(C·P) > TILE_S_KV = 1024):
primitive tl.dot + primitive online-softmax merge (baseline).
composite per-tile GEMMs as tl.composite GEMM commands.
composite_extended per-tile attention as a Q·Kᵀ composite + a
softmax_merge recipe composite (ADR-0065).
The placement, DP policy, and (m, , O) reduce are identical across
variants; only the per-tile command form differs. The sweep writes
per-(variant, S_kv) latency, engine occupancy, and PE_CPU command count
to sweep_decode_composite.json so the comparative plot
(paper_plot_gqa_decode_long_ctx_composite.py) can read off the
CPU-offload win.
Run in op_log mode (enable_data=False): latency / dispatch only — recipe
data-mode numeric parity is a separate follow-up (DDD-0065).
"""
from __future__ import annotations
import json
from pathlib import Path
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp import ( # noqa: E501
gqa_attention_decode_long_ctx_cube_sp_pe_sp_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite import ( # noqa: E501
gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext import ( # noqa: E501
gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext_kernel,
)
from kernbench.benches.gqa_helpers.shared._gqa_panel_helpers import _ccl_cfg
from kernbench.ccl.sfr_config import configure_sfr_intercube_multisip
from kernbench.policy.placement.dp import DPPolicy
_OUTPUT_DIR = (
Path(__file__).resolve().parents[2]
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _OUTPUT_DIR / "sweep_decode_composite.json"
# ── Variant + S_kv registry ──────────────────────────────────────────
# Each kernel implements the same Case-6 placement; only the per-tile
# command form differs (see module docstring).
_VARIANT_KERNELS = {
"primitive": gqa_attention_decode_long_ctx_cube_sp_pe_sp_kernel,
"composite": gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_kernel,
"composite_extended":
gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext_kernel,
}
_VARIANTS = ("primitive", "composite", "composite_extended")
# Op-count (PE_CPU dispatch) is computed at emit time (exact, instant), so
# it spans up to the 1M production point (S_local = 1M/64 = 16384, 16
# tiles): the per-PE work is Q·Kᵀ (8,128)·(128,16384) and P·V
# (8,16384)·(16384,128). End-to-end latency needs the data-mode engine,
# whose cost scales with S_kv, so it is swept only over a tractable range.
_S_KV_OPCOUNT = (8192, 65_536, 131_072, 262_144, 524_288, 1_048_576)
_S_KV_LATENCY = (8192, 32_768, 65_536, 131_072)
# LLaMA-3.1-70B single-KV-head group (8 cubes × 8 PEs), one decode step.
_BASE_PARAMS = dict(C=8, P=8, T_q=1, d_head=128, h_q=8, h_kv=1)
# ── Per-(variant, S_kv) runner ───────────────────────────────────────
def _run_panel_fn(variant: str, S_kv: int):
kernel = _VARIANT_KERNELS[variant]
p = _BASE_PARAMS
panel = f"decode_long_ctx_composite_{variant}_s{S_kv}"
def _bench_fn(ctx):
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
dp_kv = DPPolicy(cube="row_wise", pe="row_wise",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name=f"{panel}_q")
k = ctx.zeros((S_kv, p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name=f"{panel}_k")
v = ctx.zeros((S_kv, p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name=f"{panel}_v")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name=f"{panel}_o")
ctx.launch(panel, kernel, q, k, v, o,
p["T_q"], S_kv, p["h_q"], p["h_kv"], p["d_head"],
p["C"], p["P"], _auto_dim_remap=False)
return _bench_fn
def _end_to_end_ns(op_log) -> float:
if not op_log:
return 0.0
return max(r.t_end for r in op_log) - min(r.t_start for r in op_log)
def _emit_dispatch(variant: str, S_kv: int) -> tuple[int, float]:
"""PE_CPU dispatch (# commands, summed cycles) the kernel emits at the
lrab center rank (cube 6, pe 0) — computed at command-emit time, so it
is exact and S_kv-independent for the composite forms (no engine)."""
from kernbench.common.pe_commands import PeCpuOverheadCmd
from kernbench.common.pe_cost_model import DEFAULT_PE_COST_MODEL
from kernbench.triton_emu.tl_context import TLContext, run_kernel
p = _BASE_PARAMS
tl = TLContext(
pe_id=0, num_programs=p["P"], cost_model=DEFAULT_PE_COST_MODEL,
cube_id=6, num_cubes=p["C"], scratch_base=1 << 61, scratch_size=1 << 20,
)
run_kernel(
_VARIANT_KERNELS[variant], tl,
0x1000, 0x2000, 0x3000, 0x4000,
p["T_q"], S_kv, p["h_q"], p["h_kv"], p["d_head"], p["C"], p["P"],
)
cmds = [c for c in tl.commands if isinstance(c, PeCpuOverheadCmd)]
return len(cmds), sum(c.cycles for c in cmds)
def _engine_latency_ns(variant: str, S_kv: int, topology: str) -> float:
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 = resolve_topology(topology)
result = run_bench(
topology=topo, bench_fn=_run_panel_fn(variant, S_kv),
device=resolve_device(None),
engine_factory=lambda t, d: GraphEngine(
getattr(t, "topology_obj", t), enable_data=True,
),
)
if not result.completion.ok:
raise RuntimeError(
f"gqa-decode-composite {variant}@{S_kv} failed: {result.completion}"
)
return _end_to_end_ns(result.engine.op_log)
# ── Sweep entry (called by the umbrella milestone_1h_gqa) ────────────
def run_sweep(topology: str = "topology.yaml") -> int:
"""Emit-level dispatch over the full S_kv range (to 1M) + engine
latency over the tractable range; write sweep.json. Returns row count."""
_OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
rows = []
for S_kv in _S_KV_OPCOUNT:
for variant in _VARIANTS:
n_cmds, cycles = _emit_dispatch(variant, S_kv)
latency = (
_engine_latency_ns(variant, S_kv, topology)
if S_kv in _S_KV_LATENCY else None
)
rows.append({
"variant": variant,
"S_kv": S_kv,
**_BASE_PARAMS,
"pe_cpu_cmd_count": n_cmds,
"pe_cpu_dispatch_cycles": cycles,
"latency_ns": latency,
})
sweep = {
"version": 2,
"variants": list(_VARIANTS),
"s_kv_opcount": list(_S_KV_OPCOUNT),
"s_kv_latency": list(_S_KV_LATENCY),
"rows": rows,
}
_SWEEP_JSON.write_text(json.dumps(sweep, indent=2))
print(f" gqa-decode-composite: {len(rows)} rows -> {_SWEEP_JSON}")
return len(rows)
@@ -0,0 +1,168 @@
"""milestone-1h-gqa: memory-bound decode streaming composite-command study.
Single-rank companion to the compute-bound prefill study
(``gqa_prefill_compute_bound``). Same three command-form kernels
(``_gqa_prefill_compute_bound``), same single-rank (C=P=1) harness, but
driven with **T_q=1** (decode) and swept over a large S_kv. With T_q=1 the
score / context GEMMs are skinny (M = G·T_q = 8), so the MAC array is
barely fed and the kernel is **memory-bound** — streaming the KV cache out
of HBM dominates. This is the regime mirror of prefill: command form is
*latency-neutral* here, because the bottleneck is data movement, not
issue.
Running single-rank (no cross-CUBE (m,,O) reduce) isolates the
local-attention streaming vs. issue trade-off cleanly — unlike the 64-way
Case-6 decode sweep, whose small-S_kv end-to-end latency is masked by the
inter-CUBE reduce tail. S_kv here is the *per-rank* context, so S_kv=64K
single-rank is the per-PE load of a 4M-token, 64-way-sharded decode.
Records per (variant, S_kv) the end-to-end latency, GEMM/DMA engine busy
time, MAC utilization (achieved ÷ peak), and DMA occupancy (dma_busy ÷
e2e), so the comparative plot can show all three command forms landing on
the same latency curve while the MAC array stays floored and the DMA
channel stays saturated.
Runs in data mode (engine latency). Gated via the umbrella
``GQA_1H_SWEEPS=decode_streaming``.
"""
from __future__ import annotations
import json
from pathlib import Path
from kernbench.benches.gqa_helpers.long_ctx._gqa_prefill_compute_bound import (
gqa_prefill_composite_ext_kernel,
gqa_prefill_composite_kernel,
gqa_prefill_primitive_kernel,
)
from kernbench.policy.placement.dp import DPPolicy
_OUTPUT_DIR = (
Path(__file__).resolve().parents[2]
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _OUTPUT_DIR / "sweep_decode_streaming.json"
# Same kernels as the prefill study — they are general attention kernels;
# the regime is set by T_q (=1 here → memory-bound).
_VARIANT_KERNELS = {
"primitive": gqa_prefill_primitive_kernel,
"composite": gqa_prefill_composite_kernel,
"composite_extended": gqa_prefill_composite_ext_kernel,
}
_VARIANTS = ("primitive", "composite", "composite_extended")
# Per-rank context length. T_q=1, so M = G·T_q = 8 (skinny, memory-bound)
# at every point. Capped at 16K: the *plain* composite materializes the
# full (M, S_kv) scores in TCM scratch (~48·S_kv B incl. the exp transients),
# which overflows the 1 MB kernel scratch beyond ~21K — itself a sign that
# the recipe (composite_extended), which tiles the softmax, is what long
# context actually needs. S_kv is per-rank, so 16K single-rank already
# equals the per-PE load of a 1M-token, 64-way-sharded decode.
_S_KV_POINTS = (2048, 4096, 8192, 16384)
_T_Q = 1
_H_Q, _H_KV, _D_HEAD = 8, 1, 128
_PEAK_TFLOPS = 8.0 # per-PE f16 GEMM peak (topology.yaml pe_gemm.peak_tflops_f16)
def _run_panel_fn(variant: str, s_kv: int):
kernel = _VARIANT_KERNELS[variant]
panel = f"decode_stream_{variant}_s{s_kv}"
def _bench_fn(ctx):
dp = DPPolicy(cube="replicate", pe="replicate",
num_cubes=1, num_pes=1)
q = ctx.zeros((_T_Q, _H_Q * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_q")
k = ctx.zeros((s_kv, _H_KV * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_k")
v = ctx.zeros((s_kv, _H_KV * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_v")
o = ctx.empty((_T_Q, _H_Q * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_o")
ctx.launch(panel, kernel, q, k, v, o,
_T_Q, s_kv, _H_Q, _H_KV, _D_HEAD, 1, 1,
_auto_dim_remap=False)
return _bench_fn
def _end_to_end_ns(op_log) -> float:
if not op_log:
return 0.0
return max(r.t_end for r in op_log) - min(r.t_start for r in op_log)
def _engine_busy_ns(op_log, suffix: str) -> float:
return sum(r.t_end - r.t_start
for r in op_log if r.component_id.endswith("." + suffix))
def _run_panel(variant: str, s_kv: int, topology: str) -> dict:
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 = resolve_topology(topology)
result = run_bench(
topology=topo, bench_fn=_run_panel_fn(variant, s_kv),
device=resolve_device(None),
engine_factory=lambda t, d: GraphEngine(
getattr(t, "topology_obj", t), enable_data=True,
),
)
if not result.completion.ok:
raise RuntimeError(
f"gqa-decode-streaming {variant}@{s_kv} failed: {result.completion}"
)
op_log = result.engine.op_log
e2e = _end_to_end_ns(op_log)
gemm = _engine_busy_ns(op_log, "pe_gemm")
dma = _engine_busy_ns(op_log, "pe_dma")
G = _H_Q // _H_KV
M = G * _T_Q
# Useful attention flops (Q·Kᵀ + P·V), single rank.
useful_flops = 4.0 * M * _D_HEAD * s_kv
# KV-cache bytes streamed from HBM (K + V, f16) — the memory-bound
# denominator. achieved_bw = bytes / e2e (GB/s) measures how close the
# command form gets to the HBM roofline (the memory-bound mirror of
# MAC utilization in the compute-bound prefill study).
kv_bytes = 2.0 * s_kv * _D_HEAD * 2.0
return {
"variant": variant,
"s_kv": s_kv,
"M": M,
"latency_ns": e2e,
"gemm_busy_ns": gemm,
"dma_busy_ns": dma,
"kv_bytes": kv_bytes,
"achieved_bw_gbs": (kv_bytes / e2e) if e2e > 0 else 0.0,
"achieved_tflops": (useful_flops / e2e / 1e3) if e2e > 0 else 0.0,
"mac_util": (useful_flops / e2e / 1e3 / _PEAK_TFLOPS) if e2e > 0 else 0.0,
"dma_occupancy": (dma / e2e) if e2e > 0 else 0.0,
}
def run_sweep(topology: str = "topology.yaml") -> int:
"""Drive all (variant, S_kv) decode-streaming panels; write sweep.json."""
_OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
rows = [
_run_panel(variant, s_kv, topology)
for s_kv in _S_KV_POINTS
for variant in _VARIANTS
]
sweep = {
"version": 1,
"variants": list(_VARIANTS),
"s_kv_points": list(_S_KV_POINTS),
"rows": rows,
}
_SWEEP_JSON.write_text(json.dumps(sweep, indent=2))
print(f" gqa-decode-streaming: {len(rows)} rows -> {_SWEEP_JSON}")
return len(rows)
if __name__ == "__main__":
run_sweep()
@@ -0,0 +1,142 @@
"""milestone-1h-gqa: compute-bound prefill composite-command study.
Three command-form variants of a single-rank compute-bound prefill
attention kernel (``_gqa_prefill_compute_bound``), swept over context
length S_kv = T_q. Unlike the memory-bound decode study, prefill has a
large M = G·T_q, so the score / context GEMMs are compute-bound — the
regime where the composite command keeps the MAC array fed (DMA↔compute
pipelining) and the hand-tiled primitive starves it on the serial
load→dot path.
Records per (variant, context) the end-to-end latency, the GEMM-engine
busy time, and the MAC occupancy (gemm_busy / e2e) so the comparative
plot can show composite winning on both latency and utilization.
Runs in data mode (engine latency). Gated via the umbrella
``GQA_1H_SWEEPS=prefill_cb``.
"""
from __future__ import annotations
import json
from pathlib import Path
from kernbench.benches.gqa_helpers.long_ctx._gqa_prefill_compute_bound import (
gqa_prefill_composite_ext_kernel,
gqa_prefill_composite_kernel,
gqa_prefill_primitive_kernel,
)
from kernbench.policy.placement.dp import DPPolicy
_OUTPUT_DIR = (
Path(__file__).resolve().parents[2]
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _OUTPUT_DIR / "sweep_prefill_compute_bound.json"
_VARIANT_KERNELS = {
"primitive": gqa_prefill_primitive_kernel,
"composite": gqa_prefill_composite_kernel,
"composite_extended": gqa_prefill_composite_ext_kernel,
}
_VARIANTS = ("primitive", "composite", "composite_extended")
# Context length S_kv = T_q (prefill processes T_q tokens against S_kv=T_q
# keys). M = G·T_q = 8·T_q is tile-filling/compute-bound at every point.
_CTX_POINTS = (256, 512, 1024)
_H_Q, _H_KV, _D_HEAD = 8, 1, 128
_PEAK_TFLOPS = 8.0 # per-PE f16 GEMM peak (topology.yaml pe_gemm.peak_tflops_f16)
def _run_panel_fn(variant: str, ctx_len: int):
kernel = _VARIANT_KERNELS[variant]
panel = f"prefill_cb_{variant}_c{ctx_len}"
def _bench_fn(ctx):
dp = DPPolicy(cube="replicate", pe="replicate",
num_cubes=1, num_pes=1)
q = ctx.zeros((ctx_len, _H_Q * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_q")
k = ctx.zeros((ctx_len, _H_KV * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_k")
v = ctx.zeros((ctx_len, _H_KV * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_v")
o = ctx.empty((ctx_len, _H_Q * _D_HEAD),
dtype="f16", dp=dp, name=f"{panel}_o")
ctx.launch(panel, kernel, q, k, v, o,
ctx_len, ctx_len, _H_Q, _H_KV, _D_HEAD, 1, 1,
_auto_dim_remap=False)
return _bench_fn
def _end_to_end_ns(op_log) -> float:
if not op_log:
return 0.0
return max(r.t_end for r in op_log) - min(r.t_start for r in op_log)
def _engine_busy_ns(op_log, suffix: str) -> float:
return sum(r.t_end - r.t_start
for r in op_log if r.component_id.endswith("." + suffix))
def _run_panel(variant: str, ctx_len: int, topology: str) -> dict:
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 = resolve_topology(topology)
result = run_bench(
topology=topo, bench_fn=_run_panel_fn(variant, ctx_len),
device=resolve_device(None),
engine_factory=lambda t, d: GraphEngine(
getattr(t, "topology_obj", t), enable_data=True,
),
)
if not result.completion.ok:
raise RuntimeError(
f"gqa-prefill-cb {variant}@{ctx_len} failed: {result.completion}"
)
op_log = result.engine.op_log
e2e = _end_to_end_ns(op_log)
gemm = _engine_busy_ns(op_log, "pe_gemm")
dma = _engine_busy_ns(op_log, "pe_dma")
# Useful attention flops (Q·Kᵀ + P·V), single rank.
G = _H_Q // _H_KV
M = G * ctx_len
useful_flops = 4.0 * M * _D_HEAD * ctx_len
# MAC utilization = achieved / peak. ``achieved_tflops`` uses wall-clock
# (useful_flops / e2e) so it is bounded by peak even when the composite
# path overlaps many tile GEMMs (gemm_busy is a sum over overlapping ops
# and is kept only for diagnostics).
return {
"variant": variant,
"ctx_len": ctx_len,
"M": M,
"latency_ns": e2e,
"gemm_busy_ns": gemm,
"dma_busy_ns": dma,
"achieved_tflops": (useful_flops / e2e / 1e3) if e2e > 0 else 0.0,
"mac_util": (useful_flops / e2e / 1e3 / _PEAK_TFLOPS) if e2e > 0 else 0.0,
}
def run_sweep(topology: str = "topology.yaml") -> int:
"""Drive all (variant, context) prefill panels; write sweep.json."""
_OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
rows = [
_run_panel(variant, ctx_len, topology)
for ctx_len in _CTX_POINTS
for variant in _VARIANTS
]
sweep = {
"version": 1,
"variants": list(_VARIANTS),
"ctx_points": list(_CTX_POINTS),
"rows": rows,
}
_SWEEP_JSON.write_text(json.dumps(sweep, indent=2))
print(f" gqa-prefill-cb: {len(rows)} rows -> {_SWEEP_JSON}")
return len(rows)
+17 -4
View File
@@ -8,12 +8,13 @@ Currently exercises (long-context only — short-context panels are
future work):
- 4-cases prefill comparative study (gqa_helpers.long_ctx.gqa_prefill_long_ctx_4cases)
- 4-cases decode comparative study (gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases)
- Case-6 composite-command study (gqa_helpers.long_ctx.gqa_decode_long_ctx_composite)
Each sub-sweep writes its own ``sweep_{prefill,decode}.json`` into the
shared output dir ``benches/1H_milestone_output/gqa/gqa_long_ctx/``.
Each sub-sweep writes its own ``sweep_*.json`` into the shared output
dir ``benches/1H_milestone_output/gqa/gqa_long_ctx/``.
Selection via the env var ``GQA_1H_SWEEPS=prefill,decode`` (default
runs both). Toggle individual sweeps with ``GQA_1H_SWEEPS=prefill``
or ``GQA_1H_SWEEPS=decode``.
runs prefill+decode). The composite study is opt-in (it sweeps the
data-mode engine over several S_kv points): ``GQA_1H_SWEEPS=composite``.
Gated by ``GQA_1H_RUN=1`` to keep CI fast.
"""
@@ -24,6 +25,15 @@ import os
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_4cases import (
run_sweep as _run_decode_sweep,
)
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_composite import (
run_sweep as _run_composite_sweep,
)
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_streaming import (
run_sweep as _run_decode_streaming_sweep,
)
from kernbench.benches.gqa_helpers.long_ctx.gqa_prefill_compute_bound import (
run_sweep as _run_prefill_cb_sweep,
)
from kernbench.benches.gqa_helpers.long_ctx.gqa_prefill_long_ctx_4cases import (
run_sweep as _run_prefill_sweep,
)
@@ -57,6 +67,9 @@ def run(torch) -> None:
runners = {
"prefill": _run_prefill_sweep,
"decode": _run_decode_sweep,
"composite": _run_composite_sweep,
"prefill_cb": _run_prefill_cb_sweep,
"decode_streaming": _run_decode_streaming_sweep,
}
unknown = [s for s in sweeps if s not in runners]
if unknown:
+4
View File
@@ -264,6 +264,10 @@ class TLContext:
id=self._next_handle_id(),
addr=addr, shape=shape, dtype=dtype,
nbytes=nbytes, space="tcm",
# TCM-resident: a downstream composite operand reads it on-chip,
# not via a DMA_READ of its (bit-61) scratch address — same as
# the composite auto-output and recipe scratch handles below.
pinned=True,
)
# ── Reference (no DMA, metadata only) ────────────────────────
@@ -0,0 +1,191 @@
"""Phase 1 tests for the Case-6 composite-command decode variants.
Three command-form variants of the Case-6 (Cube-SP × PE-SP) long-context
decode kernel are compared in the "Use of Composite Commands" study:
async the current primitive kernel (``tl.dot`` + primitive
online-softmax merge) — the measured baseline.
composite per-tile GEMMs (Q·Kᵀ, P·V) issued as ``tl.composite``
GEMM commands; softmax merge stays primitive.
composite_extended per-tile attention as two composites — Q·Kᵀ GEMM +
a ``softmax_merge`` recipe composite folding the
online merge and P·V (ADR-0065).
This module currently holds the **refactor guard** only (Phase 1): it
freezes the byte-level command stream of the current baseline kernel so
the Phase-2 extraction of the shared (m,,O) reduce into a helper is
provably behavior-preserving. The variant kernels, the dispatch-ordering
test, and the e2e op_log test land in Phase 2 alongside the production
kernels.
"""
from __future__ import annotations
import hashlib
from pathlib import Path
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp import ( # noqa: E501
gqa_attention_decode_long_ctx_cube_sp_pe_sp_kernel as _baseline,
)
from kernbench.triton_emu.tl_context import TLContext, run_kernel
_TOPOLOGY = Path(__file__).resolve().parents[2] / "topology.yaml"
# C=8 CUBEs × P=8 PEs single-KV-head group; S_kv chosen so
# S_local = S_kv/(C·P) = 2048 > TILE_S_KV (1024) → the local-attention
# tile loop runs 2 tiles, exercising the per-tile merge path that the
# composite variants restructure.
_C, _P = 8, 8
_S_KV = 131_072
_D_HEAD, _H_Q, _H_KV, _T_Q = 128, 8, 1, 1
# Volatile per-handle identifiers (assigned by an internal counter); the
# refactor must not change *structure*, but raw ids are not part of the
# behavioral contract, so they are normalized out of the signature.
_VOLATILE_FIELDS = frozenset({"id", "completion", "cid"})
def _norm_value(v):
"""Normalize a command field to an id-free, hashable form."""
# TensorHandle / CompletionHandle-like: keep shape/addr/space/dir, drop id.
if hasattr(v, "shape") and hasattr(v, "addr"):
return ("H", getattr(v, "shape", None), getattr(v, "addr", None),
getattr(v, "space", None), getattr(v, "dtype", None))
if isinstance(v, (list, tuple)):
return tuple(_norm_value(x) for x in v)
if isinstance(v, dict):
return tuple(sorted((k, _norm_value(x)) for k, x in v.items()))
return v
def _cmd_signature(cmd) -> tuple:
fields = getattr(cmd, "__dict__", {})
return (
type(cmd).__name__,
tuple(sorted(
(name, _norm_value(val))
for name, val in fields.items()
if name not in _VOLATILE_FIELDS
)),
)
def _rank_stream(pe_id: int, cube_id: int) -> tuple:
tl = TLContext(
pe_id=pe_id, num_programs=_P, cube_id=cube_id, num_cubes=_C,
scratch_base=0x200000, scratch_size=1 << 20,
)
run_kernel(
_baseline, tl,
0x1000, 0x2000, 0x3000, 0x4000,
_T_Q, _S_KV, _H_Q, _H_KV, _D_HEAD, _C, _P,
)
return tuple(_cmd_signature(c) for c in tl.commands)
def _all_ranks_digest() -> str:
h = hashlib.sha256()
for cube_id in range(_C):
for pe_id in range(_P):
h.update(repr(_rank_stream(pe_id, cube_id)).encode())
return h.hexdigest()
# Golden digest captured from the current (pre-refactor) baseline kernel.
# Phase 2 extracts the intra-/inter-CUBE (m,,O) reduce into a shared
# helper; the refactored kernel MUST reproduce this exact stream.
_GOLDEN_DIGEST = "28069b4b1b19f427b33ee594fb71dc3987f3c92bf38b855058a2261e643047a9"
def test_case6_reduce_refactor_byte_equal():
"""Frozen command stream of the Case-6 baseline across all 64 ranks.
Guards the Phase-2 extraction of the shared reduce helper: the
refactored kernel must emit a byte-identical command stream."""
assert _all_ranks_digest() == _GOLDEN_DIGEST
# ── Variant command-form behaviour (ADR-0064 Rev2 / ADR-0065) ────────
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite import ( # noqa: E402,E501
gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_kernel as _composite,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext import ( # noqa: E402,E501
gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext_kernel as _composite_ext, # noqa: E501
)
from kernbench.common.pe_commands import CompositeCmd, PeCpuOverheadCmd # noqa: E402,E501
from kernbench.common.pe_cost_model import DEFAULT_PE_COST_MODEL # noqa: E402
def _dispatch(kernel, S_kv: int) -> tuple[int, float, int]:
"""(# PE_CPU dispatch commands, summed dispatch cycles, # composites)
emitted by ``kernel`` at the lrab center rank (cube 6, pe 0)."""
tl = TLContext(
pe_id=0, num_programs=_P, cost_model=DEFAULT_PE_COST_MODEL,
cube_id=6, num_cubes=_C, scratch_base=0x200000, scratch_size=1 << 20,
)
run_kernel(
kernel, tl, 0x1000, 0x2000, 0x3000, 0x4000,
_T_Q, S_kv, _H_Q, _H_KV, _D_HEAD, _C, _P,
)
n_disp = sum(1 for c in tl.commands if isinstance(c, PeCpuOverheadCmd))
cycles = sum(c.cycles for c in tl.commands if isinstance(c, PeCpuOverheadCmd))
n_comp = sum(1 for c in tl.commands if isinstance(c, CompositeCmd))
return n_disp, cycles, n_comp
# Two multi-tile S_kv points (S_local = 4096 / 8192 → 4 / 8 primitive tiles).
_S_KV_A, _S_KV_B = 262_144, 524_288
def test_variant_composite_counts():
"""The primitive kernel issues no composites; each composite variant
issues exactly two (Q·Kᵀ and P·V / the recipe head GEMM)."""
assert _dispatch(_baseline, _S_KV_A)[2] == 0
assert _dispatch(_composite, _S_KV_A)[2] == 2
assert _dispatch(_composite_ext, _S_KV_A)[2] == 2
def test_composite_dispatch_flat_in_skv():
"""Coarse composite issue is S_kv-independent: the kernel emits O(1)
commands and PE_SCHEDULER absorbs the per-tile fan-out, so dispatch
cost is identical at 4-tile and 8-tile S_kv (it *saturates*)."""
for kernel in (_composite, _composite_ext):
n_a, cyc_a, _ = _dispatch(kernel, _S_KV_A)
n_b, cyc_b, _ = _dispatch(kernel, _S_KV_B)
assert n_a == n_b, f"{kernel.__name__}: {n_a} != {n_b}"
assert cyc_a == cyc_b
def test_primitive_dispatch_grows_in_skv():
"""The hand-tiled primitive kernel issues O(n_tiles) commands, so its
dispatch cost grows with S_kv — the cost the composite forms remove."""
n_a, cyc_a, _ = _dispatch(_baseline, _S_KV_A)
n_b, cyc_b, _ = _dispatch(_baseline, _S_KV_B)
assert n_b > n_a
assert cyc_b > cyc_a
def test_composite_cheaper_than_primitive_at_long_ctx():
"""At a long-context (8-tile) S_kv, both composite forms issue far
fewer / cheaper PE_CPU commands than the primitive kernel — the
CPU-offload win (ADR-0064 Rev2). The shared (m,,O) reduce is a fixed
common term, so the margin is the local-attention command form alone."""
prim_cyc = _dispatch(_baseline, _S_KV_B)[1]
comp_cyc = _dispatch(_composite, _S_KV_B)[1]
ext_cyc = _dispatch(_composite_ext, _S_KV_B)[1]
assert prim_cyc > 1.5 * comp_cyc, (prim_cyc, comp_cyc)
assert prim_cyc > 1.5 * ext_cyc, (prim_cyc, ext_cyc)
def test_three_variants_complete_in_data_mode():
"""All three command forms run end-to-end through the 64-PE engine in
data mode and report a positive decode latency. Guards the composite
P·V path + the on-chip-operand engine fix (TCM operand not DMA-streamed)
at a multi-rank S_kv where S_local (=128) exceeds the recipe seed."""
from kernbench.benches.gqa_helpers.long_ctx.gqa_decode_long_ctx_composite import ( # noqa: E501
_engine_latency_ns,
)
for variant in ("primitive", "composite", "composite_extended"):
lat = _engine_latency_ns(variant, 8192, str(_TOPOLOGY))
assert lat > 0, f"{variant}: latency={lat}"
@@ -0,0 +1,61 @@
"""Tests for the compute-bound prefill composite-command study.
Unlike memory-bound decode (where command form is latency-neutral), in
compute-bound prefill the composite command keeps the MAC array fed by
pipelining DMA↔compute per HW tile, so it wins on wall-clock latency. The
rigorous claim here is therefore the *opposite* of the decode study:
``composite_latency < primitive_latency`` at a compute-bound context.
"""
from __future__ import annotations
from pathlib import Path
from kernbench.common.pe_commands import CompositeCmd
from kernbench.triton_emu.tl_context import TLContext, run_kernel
from kernbench.benches.gqa_helpers.long_ctx._gqa_prefill_compute_bound import (
gqa_prefill_composite_ext_kernel,
gqa_prefill_composite_kernel,
gqa_prefill_primitive_kernel,
)
_TOPOLOGY = Path(__file__).resolve().parents[2] / "topology.yaml"
_H_Q, _H_KV, _D_HEAD = 8, 1, 128
def _n_composites(kernel, T_q: int, S_kv: int) -> int:
tl = TLContext(pe_id=0, num_programs=1, scratch_base=1 << 61,
scratch_size=1 << 20)
run_kernel(kernel, tl, 0x1000, 0x2000, 0x3000, 0x4000,
T_q, S_kv, _H_Q, _H_KV, _D_HEAD, 1, 1)
return sum(1 for c in tl.commands if isinstance(c, CompositeCmd))
def test_command_form_composite_counts():
"""One Q-block (T_q=16 → M=128): the primitive issues no composites;
each composite form issues exactly two (Q·Kᵀ and P·V / recipe)."""
assert _n_composites(gqa_prefill_primitive_kernel, 16, 256) == 0
assert _n_composites(gqa_prefill_composite_kernel, 16, 256) == 2
assert _n_composites(gqa_prefill_composite_ext_kernel, 16, 256) == 2
def _latency_ns(variant: str, ctx_len: int) -> float:
from kernbench.benches.gqa_helpers.long_ctx.gqa_prefill_compute_bound import (
_run_panel,
)
return _run_panel(variant, ctx_len, str(_TOPOLOGY))["latency_ns"]
def test_three_variants_complete_in_data_mode():
"""All three prefill forms run end-to-end through the engine."""
for variant in ("primitive", "composite", "composite_extended"):
assert _latency_ns(variant, 256) > 0
def test_composite_faster_in_compute_bound_prefill():
"""At a compute-bound context (512), the composite command's DMA↔compute
pipelining beats the primitive's serial load→dot — the opposite of the
memory-bound decode study, where the two tie."""
prim = _latency_ns("primitive", 512)
comp = _latency_ns("composite", 512)
assert comp < prim, f"composite {comp} not < primitive {prim}"
+59
View File
@@ -0,0 +1,59 @@
"""Engine-fix spec test — composite GEMM with an on-chip (TCM) operand.
Root cause: ``_make_compute_out`` (tl_context) builds compute-result
handles with ``space="tcm"`` but omits ``pinned=True``. The PE_SCHEDULER
tiling gate (tiling.py:306) skips the operand DMA_READ only when
``pinned`` is set, so a compute result fed as a composite head-GEMM
operand ``a`` (e.g. the attention probabilities into P·V) is wrongly
scheduled as an HBM ``DMA_READ`` of its bit-61 scratch address →
``PhysAddrError`` in data mode.
Every other TCM-resident handle in tl_context is already marked
``pinned=True`` for exactly this reason (composite auto-output, recipe
scratch / primary-out). ``_make_compute_out`` is the one site that
forgot. ``.pinned`` is read ONLY by the tiling DMA gate, so setting it on
compute outputs is side-effect-free and byte-equal for existing benches
(which never feed an unpinned compute output into a composite operand).
Phase 1: this test FAILS until the fix (a DMA_READ for operand A is
emitted for the on-chip ``a``).
"""
from __future__ import annotations
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
from kernbench.triton_emu.tl_context import TLContext, run_kernel
def _pv_composite_kernel(*, tl) -> None:
"""P·V composite whose ``a`` is a real on-chip compute result."""
x = tl.load(0x1000, shape=(8, 1024), dtype="f16")
a = tl.exp(x) # compute output → TCM
V = tl.ref(0x3000, shape=(1024, 128), dtype="f16") # HBM, streamed
O = tl.zeros((8, 128), dtype="f16")
tl.composite(op="gemm", a=a, b=V, out=O)
def _pv_plan():
tl = TLContext(pe_id=0, num_programs=1,
scratch_base=0x200000, scratch_size=1 << 20)
run_kernel(_pv_composite_kernel, tl)
comp = next(c for c in tl.commands if isinstance(c, CompositeCmd))
return generate_plan_from_ops(
ops=comp.ops, tile_m=32, tile_k=64, tile_n=32,
bytes_per_element=2, pe_prefix="sip0.cube0.pe0",
)
def test_onchip_operand_a_not_dma_read():
"""The on-chip ``a`` (a compute result) must NOT be streamed via a
DMA_READ; only the HBM ref ``b`` (V) is."""
plan = _pv_plan()
read_operands = [
s.params.get("operand")
for t in plan.tiles for s in t.stages
if s.stage_type == StageType.DMA_READ
]
assert "A" not in read_operands, read_operands
assert "B" in read_operands, read_operands