diff --git a/scratch_dot_tiling_experiment.py b/scratch_dot_tiling_experiment.py new file mode 100644 index 0000000..924f00c --- /dev/null +++ b/scratch_dot_tiling_experiment.py @@ -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() diff --git a/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_tiled.py b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_tiled.py new file mode 100644 index 0000000..dad3cc1 --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/long_ctx/_gqa_attention_decode_long_ctx_cube_sp_pe_sp_tiled.py @@ -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)