a4a2683aad
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>
125 lines
4.4 KiB
Python
125 lines
4.4 KiB
Python
"""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()
|