"""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()