diff --git a/scripts/paper/paper_plot_gqa_decode_long_ctx_composite.py b/scripts/paper/paper_plot_gqa_decode_long_ctx_composite.py new file mode 100644 index 0000000..9c92733 --- /dev/null +++ b/scripts/paper/paper_plot_gqa_decode_long_ctx_composite.py @@ -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 — primitive rises O(n$_\\mathrm{tiles}$); " + "composite saturates O(1)" + ) + 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() diff --git a/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp.py b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp.py index 4a44f11..a1ca935 100644 --- a/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp.py +++ b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp.py @@ -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: diff --git a/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite.py b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite.py new file mode 100644 index 0000000..83c70e7 --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite.py @@ -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) diff --git a/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext.py b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext.py new file mode 100644 index 0000000..0dc8069 --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_composite_ext.py @@ -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) diff --git a/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_mlo_reduce.py b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_mlo_reduce.py new file mode 100644 index 0000000..e7db131 --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_mlo_reduce.py @@ -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) diff --git a/src/kernbench/benches/gqa_helpers/long_ctx/gqa_decode_long_ctx_composite.py b/src/kernbench/benches/gqa_helpers/long_ctx/gqa_decode_long_ctx_composite.py new file mode 100644 index 0000000..0ec83f1 --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/long_ctx/gqa_decode_long_ctx_composite.py @@ -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) diff --git a/src/kernbench/benches/milestone_1h_gqa.py b/src/kernbench/benches/milestone_1h_gqa.py index 451cade..8a7cb87 100644 --- a/src/kernbench/benches/milestone_1h_gqa.py +++ b/src/kernbench/benches/milestone_1h_gqa.py @@ -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,9 @@ 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_prefill_long_ctx_4cases import ( run_sweep as _run_prefill_sweep, ) @@ -57,6 +61,7 @@ def run(torch) -> None: runners = { "prefill": _run_prefill_sweep, "decode": _run_decode_sweep, + "composite": _run_composite_sweep, } unknown = [s for s in sweeps if s not in runners] if unknown: diff --git a/src/kernbench/triton_emu/tl_context.py b/src/kernbench/triton_emu/tl_context.py index e702c96..773f75f 100644 --- a/src/kernbench/triton_emu/tl_context.py +++ b/src/kernbench/triton_emu/tl_context.py @@ -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) ──────────────────────── diff --git a/tests/attention/test_gqa_decode_long_ctx_composite.py b/tests/attention/test_gqa_decode_long_ctx_composite.py new file mode 100644 index 0000000..129f155 --- /dev/null +++ b/tests/attention/test_gqa_decode_long_ctx_composite.py @@ -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}" diff --git a/tests/test_composite_tcm_operand_dma.py b/tests/test_composite_tcm_operand_dma.py new file mode 100644 index 0000000..3da30d6 --- /dev/null +++ b/tests/test_composite_tcm_operand_dma.py @@ -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