sweep+plot: add primitive hand-tiled (16×16×16) variant to Case-6 decode composite study

Wires a fourth variant into the Cube-SP × PE-SP long-context decode
composite command-form study to surface the worst-case PE_CPU dispatch
inflation that the coarse composite forms delegate to PE_SCHEDULER
(ADR-0065): per-block DMA of Q/K/V slices, tl.dot per 16³ block,
deferred K-inner sum outside the K loop.

Renamed _tiled.py → _hand_tiled_16x16x16.py; coarse-primitive kernel
retained for the 4-cases bench, paper scripts, and the golden byte-equal
regression guard.

New breakdown bar chart at S_kv=128K shows engine (~460 μs) dominates
all three variants; hand-tiled adds ~68 μs PE_CPU dispatch on top —
composite forms sit essentially on the memory-bound floor.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-13 09:28:57 -07:00
parent f8caf2e16c
commit eec61838b5
11 changed files with 540 additions and 135 deletions
@@ -0,0 +1,131 @@
"""Latency breakdown bar chart for the Case-6 composite-command decode study.
Reads sweep_decode_composite.json and writes a single-figure stacked bar
chart comparing three variants — primitive hand-tiled (16×16×16),
composite GEMM, composite + softmax_merge — at the S_kv = 131 072 point:
bottom stack: PE_CPU dispatch time (from pe_cpu_dispatch_cycles, ns at
1 GHz — ADR-0064 Rev2 D3)
top stack: engine time (latency_ns dispatch cycles) — DMA / GEMM /
MATH / IPCQ work the engine flushes on the critical path
The dispatch/engine split is a first-order breakdown; in reality the
two paths overlap partially. The note on the plot calls this out.
Run (after the composite bench sweep):
python scripts/paper/paper_plot_gqa_decode_composite_breakdown.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"
)
# The three variants to compare (coarse `primitive` intentionally dropped).
_ORDER = ("primitive_tiled", "composite", "composite_extended")
_LABELS = {
"primitive_tiled": "primitive hand-tiled\n(16×16×16)",
"composite": "composite\nGEMM",
"composite_extended": "composite +\nsoftmax_merge",
}
_S_KV_TARGET = 131_072
def main() -> None:
sweep = json.loads(_SWEEP_JSON.read_text())
rows = {(r["variant"], r["S_kv"]): r for r in sweep["rows"]}
dispatch_us = []
engine_us = []
totals_us = []
for v in _ORDER:
r = rows[(v, _S_KV_TARGET)]
disp_ns = r["pe_cpu_dispatch_cycles"]
total_ns = r["latency_ns"]
eng_ns = max(0.0, total_ns - disp_ns)
dispatch_us.append(disp_ns / 1e3)
engine_us.append(eng_ns / 1e3)
totals_us.append(total_ns / 1e3)
xs = list(range(len(_ORDER)))
labels = [_LABELS[v] for v in _ORDER]
fig, ax = plt.subplots(figsize=(7.5, 5.0))
bars_eng = ax.bar(
xs, engine_us,
color="#4f8a4f", edgecolor="#2a4a2a", label="engine (DMA + GEMM + MATH + IPCQ)",
)
bars_disp = ax.bar(
xs, dispatch_us, bottom=engine_us,
color="#c0504d", edgecolor="#5a2624", label="PE_CPU dispatch",
)
# Segment value labels — engine at mid, dispatch at mid of its stack.
for i, (eng, disp, tot) in enumerate(zip(engine_us, dispatch_us, totals_us)):
ax.text(i, eng / 2, f"{eng:.1f} µs",
ha="center", va="center", fontsize=9, color="white")
if disp / max(totals_us) > 0.04: # only label if visible
ax.text(i, eng + disp / 2, f"{disp:.1f} µs",
ha="center", va="center", fontsize=9, color="white")
else:
ax.annotate(f"{disp:.2f} µs",
xy=(i, tot), xytext=(0, 6),
textcoords="offset points",
ha="center", va="bottom", fontsize=8,
color="#5a2624")
offset_pts = 14 if disp / max(totals_us) > 0.04 else 20
ax.annotate(f"total {tot:.1f}",
xy=(i, tot), xytext=(0, offset_pts),
textcoords="offset points",
ha="center", va="bottom", fontsize=9,
color="#333", fontweight="bold")
ax.set_xticks(xs)
ax.set_xticklabels(labels, fontsize=10)
ax.set_ylabel("time (µs)")
ax.set_title(
f"Case-6 decode latency breakdown at $S_{{kv}}=${_S_KV_TARGET // 1024}K\n"
"(engine dominates all three — primitive hand-tiled adds "
f"{dispatch_us[0]:.0f} µs PE_CPU dispatch overhead)",
fontsize=11,
)
ax.legend(loc="upper right", fontsize=9)
ax.grid(True, axis="y", ls=":", alpha=0.5)
ax.set_ylim(0, max(totals_us) * 1.15)
fig.text(
0.5, 0.02,
"First-order breakdown: dispatch and engine paths overlap partially on the real "
"critical path;\ntreat the split as an upper bound on dispatch's contribution.",
ha="center", fontsize=8, color="#666",
)
fig.tight_layout(rect=(0, 0.06, 1, 1))
out = _FIG_DIR / "gqa_decode_long_ctx_composite_breakdown.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()
@@ -45,11 +45,12 @@ _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", "^"),
"primitive_tiled": ("primitive hand-tiled (16×16×16)", "#8b5a2b", "D"),
"primitive": ("primitive (tl.dot, hand-tiled)", "#c0504d", "o"),
"composite": ("composite GEMM", "#3b6ea5", "s"),
"composite_extended": ("composite + softmax_merge", "#4f8a4f", "^"),
}
_ORDER = ("primitive", "composite", "composite_extended")
_ORDER = ("primitive_tiled", "composite", "composite_extended")
def _load() -> dict: