2 Commits

Author SHA1 Message Date
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
15 changed files with 388 additions and 37 deletions
Binary file not shown.
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 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}). updates riding a 16\,B side-channel credit (\S\ref{sec:allreduce}).
\begin{figure}[t] \begin{figure*}[t]
\centering \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 \caption{Per-send data and control flow for the four PE-to-PE
signalling mechanisms (sender\,$\rightarrow$\,NoC\,$\rightarrow$\,receiver). signalling mechanisms (sender\,$\rightarrow$\,NoC\,$\rightarrow$\,receiver).
Doorbell and RDMA-CQ each issue two fabric transactions (payload then 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 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.} a \emph{design schematic}, not a measured comparison.}
\label{fig:ipcq-arch} \label{fig:ipcq-arch}
\end{figure} \end{figure*}
\begin{figure}[t] \begin{figure}[t]
\centering \centering
+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()
@@ -243,6 +243,12 @@ _FLOW_DATA = "#6699CC" # blue flit (data)
_FLOW_CTRL = "#ED7D65" # orange flit (control) _FLOW_CTRL = "#ED7D65" # orange flit (control)
_FLOW_CRED = "#7BB661" # green (credit / ack) _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): def _flow_swim(ax, x, y, w, h, label, sub=None):
ax.add_patch(mpatches.FancyBboxPatch( 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)) facecolor="white", edgecolor="#3a3a3a", linewidth=1.6))
ax.text(x + w / 2, y + h / 2 + (0.20 if sub else 0), ax.text(x + w / 2, y + h / 2 + (0.20 if sub else 0),
label, ha="center", va="center", label, ha="center", va="center",
fontsize=10.5, fontweight="bold", color="#111") fontsize=10.5 * _FS, fontweight="bold", color="#111")
if sub: if sub:
ax.text(x + w / 2, y + h / 2 - 0.35, sub, ax.text(x + w / 2, y + h / 2 - 0.35, sub,
ha="center", va="center", 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): 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)) facecolor=color, edgecolor="#222", linewidth=0.8))
if letter: if letter:
ax.text(x, y, letter, ha="center", va="center", 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"): 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): def _flow_label(ax, x, y, text, color="#222", fontsize=9):
ax.text(x, y, text, ha="center", va="center", 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"): 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], facecolor="white", edgecolor=_COLOR[kind],
linewidth=0.9, alpha=0.85, zorder=0)) linewidth=0.9, alpha=0.85, zorder=0))
ax.text(10, 0.65, _TITLE[kind], ax.text(10, 0.65, _TITLE[kind],
ha="center", va="center", fontsize=11.5, ha="center", va="center", fontsize=11.5 * _FS,
fontweight="bold", fontweight="bold",
color=("#1B5E20" if kind == "ipcq" else "white"), color=("#1B5E20" if kind == "ipcq" else "white"),
bbox=dict( bbox=dict(
@@ -449,7 +455,7 @@ def _flow_legend(fig, leg_y=0.02):
facecolor=color, edgecolor="#222", linewidth=0.6, facecolor=color, edgecolor="#222", linewidth=0.6,
transform=fig.transFigure)) transform=fig.transFigure))
fig.text(lx + 0.025, y + 0.011, text, 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") fontweight="bold", color="#222")
lx += 0.32 lx += 0.32
@@ -461,7 +467,11 @@ def _flow_legend(fig, leg_y=0.02):
def _plot_architecture_flow() -> Path: def _plot_architecture_flow() -> Path:
"""4 panels arranged 2×2 — each a horizontal Sender → NoC → Receiver """4 panels arranged 2×2 — each a horizontal Sender → NoC → Receiver
flow with small flit-style packets, dot-marker annotations and a flow with small flit-style packets, dot-marker annotations and a
bottom name plate. Aesthetic mirrors latency_model.png.""" 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)) fig, axes = plt.subplots(2, 2, figsize=(18.0, 11.5))
_draw_flow_case(axes[0, 0], "doorbell") _draw_flow_case(axes[0, 0], "doorbell")
_draw_flow_case(axes[0, 1], "hmq") _draw_flow_case(axes[0, 1], "hmq")
@@ -470,12 +480,14 @@ def _plot_architecture_flow() -> Path:
fig.suptitle( fig.suptitle(
"Per-send architecture = $\\Sigma$ control overhead + " "Per-send architecture = $\\Sigma$ control overhead + "
"data-path DMA + receiver wake", "data-path DMA + receiver wake",
fontsize=14.5, fontweight="bold", y=0.995) fontsize=14.5 * _FS, fontweight="bold", y=0.995)
_flow_legend(fig, leg_y=0.015) _flow_legend(fig, leg_y=0.015)
fig.tight_layout(rect=(0, 0.07, 1, 0.96)) fig.tight_layout(rect=(0, 0.07, 1, 0.96))
out = _OUT_DIR / "ipcq_alternatives_architecture_flow.png" out = _OUT_DIR / "ipcq_alternatives_architecture_flow.png"
fig.savefig(out, dpi=150, bbox_inches="tight") fig.savefig(out, dpi=150, bbox_inches="tight")
plt.close(fig) plt.close(fig)
finally:
_FS = 1.0
return out return out
@@ -483,19 +495,25 @@ def _plot_architecture_flow() -> Path:
def _plot_architecture_stacked() -> Path: def _plot_architecture_stacked() -> Path:
"""Same per-case flow content as the 2×2 figure, but stacked one """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 case per row (4 rows × 1 column). Wider panels give each case more
horizontal room.""" horizontal room. Rendered at full text width (two-column figure*), so
fig, axes = plt.subplots(4, 1, figsize=(16.0, 19.0)) 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"]): for ax, kind in zip(axes, ["doorbell", "hmq", "rdma", "ipcq"]):
_draw_flow_case(ax, kind) _draw_flow_case(ax, kind)
fig.suptitle( fig.suptitle(
"Per-send architecture = $\\Sigma$ control overhead + " "Per-send architecture = $\\Sigma$ control overhead + "
"data-path DMA + receiver wake", "data-path DMA + receiver wake",
fontsize=14.5, fontweight="bold", y=0.997) fontsize=14.5 * _FS, fontweight="bold", y=0.997)
_flow_legend(fig, leg_y=0.010) _flow_legend(fig, leg_y=0.010)
fig.tight_layout(rect=(0, 0.05, 1, 0.97)) fig.tight_layout(rect=(0, 0.06, 1, 0.97))
out = _OUT_DIR / "ipcq_alternatives_architecture_stacked.png" out = _OUT_DIR / "ipcq_alternatives_architecture_stacked.png"
fig.savefig(out, dpi=150, bbox_inches="tight") fig.savefig(out, dpi=150, bbox_inches="tight")
plt.close(fig) plt.close(fig)
finally:
_FS = 1.0
return out 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

@@ -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)