"""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}"