diff --git a/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short.py b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short.py index 8dd5526..a5d34f6 100644 --- a/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short.py +++ b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short.py @@ -1,4 +1,4 @@ -"""GQA decode kernel: short context, attention only, multi-tile per PE. +"""GQA decode kernel: short context, attention only, multi-tile per PE (1). Unified A1/A2/A4/B decode mapping per ADR-0070 (phase mirror of ``_gqa_attention_prefill_short.py``). Mode selected at launch via diff --git a/src/kernbench/benches/_gqa_attention_decode_short_composite.py b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short_composite.py similarity index 78% rename from src/kernbench/benches/_gqa_attention_decode_short_composite.py rename to src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short_composite.py index a9d8875..1c4fbf8 100644 --- a/src/kernbench/benches/_gqa_attention_decode_short_composite.py +++ b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short_composite.py @@ -1,19 +1,22 @@ -"""GQA decode kernel — composite (second-level) variant. +"""GQA decode kernel: composite GEMM-only variant (2). Identical mapping to ``_gqa_attention_decode_short.py`` (sequence-shard -+ chain reduce + FA2 head fusion, unified A1/A2/A4/B). Per-tile fusion -follows the ADR-0065 two-composite pattern (cf. decode_opt2): ++ chain reduce + FA2 head fusion, unified A1/A2/A4/B). Difference vs +first-level: the Q·Kᵀ GEMM is issued via ``tl.composite(op="gemm", ...)`` +instead of ``tl.dot``. - tile 0 primitives establish (m, ℓ, O) - tile k>0 #1 Q·Kᵀ composite → scores (pinned primary-out) - #2 softmax_merge prologue → P (pinned, auto-binds to GEMM a) - + P·V composite, ``out=O_local`` - + add epilogue folding the result into O_local +This is the GEMM-only tier: Q·Kᵀ is a composite (operands ``a=Q`` / +``b=K_T`` are pinned ``tl.load`` results), while P·V stays a plain +``tl.dot`` and softmax stays a primitive MATH chain — no +``softmax_merge`` fusion. Folding the per-tile softmax into a P·V +composite (the ``softmax_merge`` prologue making ``P`` a pinned +primary-out bound to the head GEMM) is variant (3) in +``_gqa_attention_decode_short_composite_fused.py``. -The softmax_merge recipe (tl_recipes.py) folds max/exp/sum into the -composite, eliminating the GEMM/MATH engine bubble that a flat -primitives chain incurs. The intra-group chain reduce after the tile -loop is unchanged. +Three-variant comparison: + (1) without composite : ``_gqa_attention_decode_short.py`` + (2) with composite (GEMM-only, no fuse) : this file + (3) with composite + softmax_merge fuse : ``…_composite_fused.py`` Shard addressing, layouts, and caller contract are identical to the first-level decode kernel (ADR-0011 D-VA1). @@ -112,23 +115,17 @@ def gqa_attention_decode_short_composite_kernel( Q = tl.load(q_base, shape=(G, d_head), dtype="f16") - # ── Tile 0: establish running (m, ℓ, O) with primitives ────────── - # (Reference: decode_opt2 — running state is set up with tl.dot/MATH - # primitives, not composite. Recipe-driven composite enters in tile 1+.) + # ── Tile 0: establish (m, ℓ, O) — Q·Kᵀ composite, softmax + P·V primitives ── K_T = tl.load(k_shard_base, shape=(d_head, TILE_S_KV), dtype="f16") V = tl.load(v_shard_base, shape=(TILE_S_KV, d_head), dtype="f16") - scores = tl.dot(Q, K_T) + scores = tl.composite(op="gemm", a=Q, b=K_T) 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 = tl.dot(exp_scores, V) + O_local = tl.dot(exp_scores, V) # P·V stays primitive in the GEMM-only tier - # ── Tiles 1..n_tiles_per_pe-1: two composites per tile ────────── - # #1 Q·Kᵀ composite → scores (pinned primary-out, fed into #2). - # #2 softmax_merge prologue + P·V GEMM + add epilogue, all in one - # composite: updates (m, ℓ) in place, computes P, runs P·V with - # pinned auto-bind, and folds the result into O_local. + # ── Tiles 1..n_tiles_per_pe-1: Q·Kᵀ composite + softmax + P·V tl.dot ── for tile_idx in range(1, n_tiles_per_pe): with tl.scratch_scope(): K_T_t = tl.load(k_shard_base + tile_idx * K_TILE_BYTES, @@ -136,12 +133,17 @@ def gqa_attention_decode_short_composite_kernel( V_t = tl.load(v_shard_base + tile_idx * TILE_S_KV * KV_ROW_BYTES, shape=(TILE_S_KV, d_head), dtype="f16") scores_t = tl.composite(op="gemm", a=Q, b=K_T_t) - tl.composite( - prologue=[{"op": "softmax_merge", "s": scores_t, - "m": m_local, "l": l_local, "O": O_local}], - op="gemm", b=V_t, out=O_local, - epilogue=[{"op": "add", "other": O_local}], + 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 = tl.dot(exp_scores_t, V_t) + 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) # ── Chain reduce ── group_cols = min(4, group_size) diff --git a/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short_composite_fused.py b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short_composite_fused.py new file mode 100644 index 0000000..5ff1865 --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_decode_short_composite_fused.py @@ -0,0 +1,195 @@ +"""GQA decode kernel: composite + softmax_merge fused variant (3). + +Identical mapping to ``_gqa_attention_decode_short.py``. Difference vs +the GEMM-only composite: per-tile softmax is folded into +the P·V composite via the ``softmax_merge`` prologue recipe, eliminating +the GEMM/MATH engine bubble that a flat primitives chain incurs. + +Per-tile structure (ADR-0065 two-composite pattern, cf. decode_opt2): + tile 0 primitives establish (m, ℓ, O) + tile k>0 #1 Q·Kᵀ composite → scores (pinned primary-out) + #2 softmax_merge prologue → P (pinned, auto-binds to GEMM a) + + P·V composite, ``out=O_local`` + + add epilogue folding the result into O_local + +The intra-group chain reduce after the tile loop is unchanged. + +Three-variant comparison: + (1) without composite : ``_gqa_attention_decode_short.py`` + (2) with composite (GEMM-only, no fuse) : ``…_composite.py`` + (3) with composite + softmax_merge fuse : this file + +Shard addressing, layouts, and caller contract are identical to the +first-level decode kernel (ADR-0011 D-VA1). +""" +from __future__ import annotations + + +TILE_S_KV = 1024 + + +def _validate_config(*, kv_per_cube: int, T_q: int, P: int, C: int, + h_q: int, h_kv: int, S_kv: int) -> None: + """Validate composite-decode config — caller-side, sim-cost 0. + + Mirrors first-level decode ``_validate_config``. + """ + if kv_per_cube not in (1, 2, 4, 8): + raise ValueError(f"kv_per_cube must be in {{1,2,4,8}}; got {kv_per_cube}") + if T_q != 1: + raise ValueError(f"decode requires T_q == 1; got {T_q}") + if P != 8: + raise ValueError(f"2x4 mesh chain-reduce geometry requires P == 8; got {P}") + if h_q % h_kv != 0: + raise ValueError(f"GQA group G = h_q/h_kv must be integer; got h_q={h_q}, h_kv={h_kv}") + if C != h_kv // kv_per_cube: + raise ValueError( + f"C must equal h_kv/kv_per_cube = {h_kv // kv_per_cube}; " + f"got C={C} (h_kv={h_kv}, kv_per_cube={kv_per_cube})" + ) + group_size = P // kv_per_cube + if S_kv % (group_size * TILE_S_KV) != 0: + raise ValueError( + f"S_kv ({S_kv}) must be a multiple of group_size·TILE_S_KV " + f"({group_size * TILE_S_KV}); each PE's sequence shard must " + f"be a whole number of tiles" + ) + + +def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl): + """Used only by the intra-group chain reduce below (recipe handles + per-tile fold internally).""" + 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_short_composite_fused_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, + kv_per_cube: int, + *, + tl, +) -> None: + """Composite-GEMM decode — same mapping as first-level + tl.composite. + + Caller must invoke ``_validate_config(...)`` first. + """ + group_size = P // kv_per_cube + S_local = S_kv // group_size + n_tiles_per_pe = S_local // TILE_S_KV + + G = h_q // h_kv + pe_id = tl.program_id(axis=0) + cube_id = tl.program_id(axis=1) + pe_in_group = pe_id % group_size + group_id_in_cube = pe_id // group_size + + KV_ROW_BYTES = d_head * 2 + K_TILE_BYTES = d_head * TILE_S_KV * 2 + Q_ROW_BYTES = G * d_head * 2 + K_HEAD_BYTES = S_kv * d_head * 2 + V_HEAD_BYTES = S_kv * KV_ROW_BYTES + + # Global VA per ADR-0011 D-VA1: kernel computes its own shard base. + q_base = q_ptr + cube_id * kv_per_cube * Q_ROW_BYTES + group_id_in_cube * Q_ROW_BYTES + k_shard_base = (k_ptr + + cube_id * kv_per_cube * K_HEAD_BYTES + + group_id_in_cube * K_HEAD_BYTES + + pe_in_group * n_tiles_per_pe * K_TILE_BYTES) + v_shard_base = (v_ptr + + cube_id * kv_per_cube * V_HEAD_BYTES + + group_id_in_cube * V_HEAD_BYTES + + pe_in_group * n_tiles_per_pe * TILE_S_KV * KV_ROW_BYTES) + + Q = tl.load(q_base, shape=(G, d_head), dtype="f16") + + # ── Tile 0: establish running (m, ℓ, O) with primitives ────────── + # (Reference: decode_opt2 — running state is set up with tl.dot/MATH + # primitives, not composite. Recipe-driven composite enters in tile 1+.) + K_T = tl.load(k_shard_base, shape=(d_head, TILE_S_KV), dtype="f16") + V = tl.load(v_shard_base, shape=(TILE_S_KV, d_head), dtype="f16") + scores = tl.dot(Q, K_T) + 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 = tl.dot(exp_scores, V) + + # ── Tiles 1..n_tiles_per_pe-1: two composites per tile ────────── + # #1 Q·Kᵀ composite → scores (pinned primary-out, fed into #2). + # #2 softmax_merge prologue + P·V GEMM + add epilogue, all in one + # composite: updates (m, ℓ) in place, computes P, runs P·V with + # pinned auto-bind, and folds the result into O_local. + for tile_idx in range(1, n_tiles_per_pe): + with tl.scratch_scope(): + K_T_t = tl.load(k_shard_base + tile_idx * K_TILE_BYTES, + shape=(d_head, TILE_S_KV), dtype="f16") + V_t = tl.load(v_shard_base + tile_idx * TILE_S_KV * KV_ROW_BYTES, + shape=(TILE_S_KV, d_head), dtype="f16") + scores_t = tl.composite(op="gemm", a=Q, b=K_T_t) + tl.composite( + prologue=[{"op": "softmax_merge", "s": scores_t, + "m": m_local, "l": l_local, "O": O_local}], + op="gemm", b=V_t, out=O_local, + epilogue=[{"op": "add", "other": O_local}], + ) + + # ── Chain reduce ── + group_cols = min(4, group_size) + group_rows = (group_size + group_cols - 1) // group_cols + pe_col_in_group = pe_in_group % group_cols + pe_row_in_group = pe_in_group // group_cols + + if group_cols > 1: + if pe_col_in_group < group_cols - 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_in_group > 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_in_group == 0 and group_rows > 1: + if pe_row_in_group < group_rows - 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_in_group > 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) + + if pe_in_group == 0: + O_final = O_local / l_local + o_base = (o_ptr + + cube_id * kv_per_cube * Q_ROW_BYTES + + group_id_in_cube * Q_ROW_BYTES) + tl.store(o_base, O_final) diff --git a/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short.py b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short.py index 59fdbcc..8c1aa29 100644 --- a/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short.py +++ b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short.py @@ -1,4 +1,4 @@ -"""GQA prefill kernel: short context, attention only, multi-tile. +"""GQA prefill kernel: short context, attention only, multi-tile (1). Unified A1/A2/A4/B prefill mapping per ADR-0070 (supersedes ADR-0060 §B.split.2 prefill-short clause). Mode selected at launch diff --git a/src/kernbench/benches/_gqa_attention_prefill_short_composite.py b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short_composite.py similarity index 82% rename from src/kernbench/benches/_gqa_attention_prefill_short_composite.py rename to src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short_composite.py index 385b1a3..1a8426d 100644 --- a/src/kernbench/benches/_gqa_attention_prefill_short_composite.py +++ b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short_composite.py @@ -1,24 +1,22 @@ -"""GQA prefill kernel — composite (second-level) variant. +"""GQA prefill kernel: composite GEMM-only variant (2). Identical mapping to ``_gqa_attention_prefill_short.py`` (Q-tile split + -FA2 head fusion + IPCQ KV broadcast, unified A1/A2/A4/B). +FA2 head fusion + IPCQ KV broadcast, unified A1/A2/A4/B). Difference vs +first-level: Q·Kᵀ uses ``tl.composite(op="gemm")`` instead of ``tl.dot``. +``P·V`` stays a plain ``tl.dot``: this is the GEMM-only tier, where +softmax remains a primitive MATH chain with no ``softmax_merge`` fusion. -Composite hybrid (option c1): Q·Kᵀ uses ``tl.composite(op="gemm")`` so -GEMM and any future prologue/epilogue can overlap; P·V stays a plain -``tl.dot``. The reason: prefill broadcasts V over IPCQ, so on non-root -PEs V is a recv'd TCM slot (unpinned per ADR-0065 D4). Feeding it to a -composite would trigger the scheduler's "stream from HBM" branch and -PhysAddr-decode the slot address (PE scratch, bit-61 set) → PageFault. -``tl.dot`` reads TCM in place, preserving the broadcast model. +Full second-level fusion (with the ``softmax_merge`` prologue making +P a pinned primary-out bound to a P·V composite) is variant (3) in +``_gqa_attention_prefill_short_composite_fused.py``. -This means prefill cannot host the ``softmax_merge`` prologue (which -attaches to the P·V composite). Full second-level fusion would require -either: (i) pinning recv handles (sim/API change, ADR-0065 supplement), -or (ii) skipping the broadcast model — both out of scope here. +Three-variant comparison: + (1) without composite : ``_gqa_attention_prefill_short.py`` + (2) with composite (GEMM-only, no fuse) : this file + (3) with composite + softmax_merge fuse : ``…_composite_fused.py`` -Shard addressing, layouts, and the caller contract are identical to -the first-level kernel (see its docstring for the full Q/K/V/O layout -and ADR-0011 D-VA1 offset computation). +Shard addressing, layouts, and caller contract are identical to the +first-level kernel (ADR-0011 D-VA1). """ from __future__ import annotations @@ -151,13 +149,8 @@ def gqa_attention_prefill_short_composite_kernel( tl.send(dir="intra_E", src=K_T) tl.send(dir="intra_E", src=V) - # Q·Kᵀ is a composite GEMM (pinned load operands). P·V stays a - # primitive tl.dot: in the broadcast topology, V on non-root PEs is - # a recv'd TCM slot (unpinned), which the composite scheduler would - # try to stream from HBM. tl.dot reads TCM in-place, preserving - # the broadcast model. softmax_merge prologue requires a composite - # P·V to host it — so prefill keeps Q·Kᵀ-only composite (no fused - # softmax epilogue); see kernel docstring (c1). + # Q·Kᵀ is a composite GEMM; P·V stays a primitive tl.dot in this + # GEMM-only tier (no softmax_merge fusion — that is variant 3). scores = tl.composite(op="gemm", a=Q, b=K_T) m_local = tl.max(scores, axis=-1) centered = scores - m_local diff --git a/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short_composite_fused.py b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short_composite_fused.py new file mode 100644 index 0000000..3d89ede --- /dev/null +++ b/src/kernbench/benches/gqa_helpers/short_ctx/_gqa_attention_prefill_short_composite_fused.py @@ -0,0 +1,204 @@ +"""GQA prefill kernel: composite + softmax_merge fused variant (3). + +Identical mapping to ``_gqa_attention_prefill_short.py``. Difference vs +the GEMM-only composite: per-tile softmax is folded into +the P·V composite via the ``softmax_merge`` prologue recipe, eliminating +the GEMM/MATH engine bubble. + +Multi-cube (A1/A2/A4): non-root PEs receive K/V via IPCQ ``tl.recv``. +Recv'd slots in on-chip memory are pinned (read in place as a composite +operand, not DMA-streamed from HBM). + +Three-variant comparison: + (1) without composite : ``_gqa_attention_prefill_short.py`` + (2) with composite (GEMM-only, no fuse) : ``…_composite.py`` + (3) with composite + softmax_merge fuse : this file + +Shard addressing, layouts, and caller contract are identical to the +first-level kernel (ADR-0011 D-VA1). +""" +from __future__ import annotations + + +TILE_S_KV = 1024 + + +def _validate_config(*, kv_per_cube: int, T_q: int, P: int, C: int, + h_q: int, h_kv: int, S_kv: int) -> None: + """Validate composite-prefill config — caller-side, sim-cost 0. + + Mirrors first-level prefill ``_validate_config``. + """ + if kv_per_cube not in (1, 2, 4, 8): + raise ValueError(f"kv_per_cube must be in {{1,2,4,8}}; got {kv_per_cube}") + if P != 8: + raise ValueError(f"2x4 mesh broadcast geometry requires P == 8; got {P}") + if h_q % h_kv != 0: + raise ValueError(f"GQA group G = h_q/h_kv must be integer; got h_q={h_q}, h_kv={h_kv}") + if C != h_kv // kv_per_cube: + raise ValueError( + f"C must equal h_kv/kv_per_cube = {h_kv // kv_per_cube}; " + f"got C={C} (h_kv={h_kv}, kv_per_cube={kv_per_cube})" + ) + group_size = P // kv_per_cube + if T_q < group_size: + raise ValueError( + f"T_q ({T_q}) must be >= group_size ({group_size}) so every PE " + f"has q-tile work for the broadcast chain" + ) + if S_kv % TILE_S_KV != 0: + raise ValueError( + f"S_kv ({S_kv}) must be a multiple of TILE_S_KV ({TILE_S_KV})" + ) + + +def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl): + 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_prefill_short_composite_fused_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, + kv_per_cube: int, + *, + tl, +) -> None: + """Composite-GEMM prefill — same mapping as first-level + tl.composite. + + Caller must invoke ``_validate_config(...)`` first. + """ + group_size = P // kv_per_cube + + G = h_q // h_kv + pe_id = tl.program_id(axis=0) + cube_id = tl.program_id(axis=1) + pe_in_group = pe_id % group_size + group_id_in_cube = pe_id // group_size + + q_start = (T_q * pe_in_group) // group_size + q_end = (T_q * (pe_in_group + 1)) // group_size + T_q_pe = q_end - q_start + + n_tiles = S_kv // TILE_S_KV + + Q_ROW_BYTES = G * d_head * 2 + KV_ROW_BYTES = d_head * 2 + K_TILE_BYTES = d_head * TILE_S_KV * 2 + K_HEAD_BYTES = n_tiles * K_TILE_BYTES + V_HEAD_BYTES = S_kv * KV_ROW_BYTES + + # Global VA per ADR-0011 D-VA1: kernel computes its own shard base. + q_base = (q_ptr + + cube_id * kv_per_cube * T_q * Q_ROW_BYTES + + (group_id_in_cube * T_q + q_start) * Q_ROW_BYTES) + k_head_shard_base = (k_ptr + + cube_id * kv_per_cube * K_HEAD_BYTES + + group_id_in_cube * K_HEAD_BYTES) + v_head_shard_base = (v_ptr + + cube_id * kv_per_cube * V_HEAD_BYTES + + group_id_in_cube * V_HEAD_BYTES) + + Q = tl.load(q_base, shape=(T_q_pe * G, d_head), dtype="f16") + + group_cols = min(4, group_size) + group_rows = (group_size + group_cols - 1) // group_cols + pe_col_in_group = pe_in_group % group_cols + pe_row_in_group = pe_in_group // group_cols + + # ── Tile 0 — broadcast + persistent (m, ℓ, O) ── + if pe_in_group == 0: + K_T = tl.load(k_head_shard_base, + shape=(d_head, TILE_S_KV), dtype="f16") + V = tl.load(v_head_shard_base, + shape=(TILE_S_KV, d_head), dtype="f16") + if group_cols > 1: + tl.send(dir="intra_E", src=K_T) + tl.send(dir="intra_E", src=V) + if group_rows > 1: + tl.send(dir="intra_S", src=K_T) + tl.send(dir="intra_S", src=V) + elif pe_col_in_group == 0 and pe_row_in_group > 0: + K_T = tl.recv(dir="intra_N", + shape=(d_head, TILE_S_KV), dtype="f16") + V = tl.recv(dir="intra_N", + shape=(TILE_S_KV, d_head), dtype="f16") + if group_cols > 1: + tl.send(dir="intra_E", src=K_T) + tl.send(dir="intra_E", src=V) + else: + K_T = tl.recv(dir="intra_W", + shape=(d_head, TILE_S_KV), dtype="f16") + V = tl.recv(dir="intra_W", + shape=(TILE_S_KV, d_head), dtype="f16") + if pe_col_in_group < group_cols - 1: + tl.send(dir="intra_E", src=K_T) + tl.send(dir="intra_E", src=V) + + # Tile 0 — primitives establish (m, ℓ, O); recipe fusion enters tile 1+. + scores = tl.dot(Q, K_T) + 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 = tl.dot(exp_scores, V) + + # ── Tiles 1..n_tiles-1 ── + for tile_idx in range(1, n_tiles): + with tl.scratch_scope(): + if pe_in_group == 0: + K_T_t = tl.load(k_head_shard_base + tile_idx * K_TILE_BYTES, + shape=(d_head, TILE_S_KV), dtype="f16") + V_t = tl.load(v_head_shard_base + tile_idx * TILE_S_KV * KV_ROW_BYTES, + shape=(TILE_S_KV, d_head), dtype="f16") + if group_cols > 1: + tl.send(dir="intra_E", src=K_T_t) + tl.send(dir="intra_E", src=V_t) + if group_rows > 1: + tl.send(dir="intra_S", src=K_T_t) + tl.send(dir="intra_S", src=V_t) + elif pe_col_in_group == 0 and pe_row_in_group > 0: + K_T_t = tl.recv(dir="intra_N", + shape=(d_head, TILE_S_KV), dtype="f16") + V_t = tl.recv(dir="intra_N", + shape=(TILE_S_KV, d_head), dtype="f16") + if group_cols > 1: + tl.send(dir="intra_E", src=K_T_t) + tl.send(dir="intra_E", src=V_t) + else: + K_T_t = tl.recv(dir="intra_W", + shape=(d_head, TILE_S_KV), dtype="f16") + V_t = tl.recv(dir="intra_W", + shape=(TILE_S_KV, d_head), dtype="f16") + if pe_col_in_group < group_cols - 1: + tl.send(dir="intra_E", src=K_T_t) + tl.send(dir="intra_E", src=V_t) + # Two-composite fusion: Q·Kᵀ composite → softmax_merge prologue + # binds P (pinned primary-out) to the P·V composite, folding + # the new tile's contribution into O_local. + scores_t = tl.composite(op="gemm", a=Q, b=K_T_t) + tl.composite( + prologue=[{"op": "softmax_merge", "s": scores_t, + "m": m_local, "l": l_local, "O": O_local}], + op="gemm", b=V_t, out=O_local, + epilogue=[{"op": "add", "other": O_local}], + ) + + O_final = O_local / l_local + o_base = (o_ptr + + cube_id * kv_per_cube * T_q * Q_ROW_BYTES + + (group_id_in_cube * T_q + q_start) * Q_ROW_BYTES) + tl.store(o_base, O_final) diff --git a/tests/attention/test_gqa_short_context.py b/tests/attention/test_gqa_short_context.py index 30c1094..5712ef8 100644 --- a/tests/attention/test_gqa_short_context.py +++ b/tests/attention/test_gqa_short_context.py @@ -64,7 +64,7 @@ def _count(op_log, name: str) -> int: def _run_prefill(*, kv_per_cube: int, C: int, T_q: int, S_kv: int, h_q: int = 8): """Run the unified prefill kernel in the given mode.""" - from kernbench.benches._gqa_attention_prefill_short import ( + from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_prefill_short import ( _validate_config as _validate_prefill_config, gqa_attention_prefill_short_kernel, ) @@ -182,7 +182,7 @@ def test_prefill_multitile_scaling(kv_per_cube, C): def _run_decode(*, kv_per_cube: int, C: int, S_kv: int, h_q: int = 8): """Run the unified decode kernel in the given mode.""" - from kernbench.benches._gqa_attention_decode_short import ( + from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_decode_short import ( _validate_config as _validate_decode_config, gqa_attention_decode_short_kernel, ) @@ -387,7 +387,7 @@ def test_prefill_per_cube_dma_balanced(kv_per_cube, C): def _run_prefill_composite(*, kv_per_cube: int, C: int, T_q: int, S_kv: int, h_q: int = 8): """Same helper as _run_prefill but for the composite-GEMM kernel.""" - from kernbench.benches._gqa_attention_prefill_short_composite import ( + from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_prefill_short_composite import ( _validate_config as _validate_prefill_cmp_config, gqa_attention_prefill_short_composite_kernel, ) @@ -428,7 +428,7 @@ def _run_prefill_composite(*, kv_per_cube: int, C: int, T_q: int, def _run_decode_composite(*, kv_per_cube: int, C: int, S_kv: int, h_q: int = 8): """Same helper as _run_decode but for the composite-GEMM kernel.""" - from kernbench.benches._gqa_attention_decode_short_composite import ( + from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_decode_short_composite import ( _validate_config as _validate_decode_cmp_config, gqa_attention_decode_short_composite_kernel, ) @@ -467,9 +467,92 @@ def _run_decode_composite(*, kv_per_cube: int, C: int, S_kv: int, ) +def _run_prefill_composite_fused(*, kv_per_cube: int, C: int, T_q: int, + S_kv: int, h_q: int = 8): + """Variant (3): GEMM-only composite + softmax_merge prologue fusion.""" + from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_prefill_short_composite_fused import ( + _validate_config as _validate_prefill_fuse_config, + gqa_attention_prefill_short_composite_fused_kernel, + ) + _validate_prefill_fuse_config(kv_per_cube=kv_per_cube, T_q=T_q, P=P, + C=C, h_q=h_q, h_kv=H_KV, S_kv=S_kv) + n_tiles = S_kv // TILE_S_KV + Q_ROWS = kv_per_cube * T_q + Q_COLS = (h_q * D_HEAD) // kv_per_cube + topo = resolve_topology(str(TOPOLOGY_DEFAULT)) + + def _bench_fn(ctx): + configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg()) + dp_q = DPPolicy(cube="column_wise", pe="replicate", num_cubes=C, num_pes=P) + dp_kv = DPPolicy(cube="row_wise", pe="replicate", num_cubes=C, num_pes=P) + dp_o = DPPolicy(cube="column_wise", pe="replicate", num_cubes=C, num_pes=P) + q = ctx.zeros((Q_ROWS, Q_COLS), dtype=DTYPE, dp=dp_q, + name=f"q_pre_fuse_kv{kv_per_cube}") + k = ctx.zeros((H_KV * n_tiles * D_HEAD, TILE_S_KV), + dtype=DTYPE, dp=dp_kv, name=f"k_pre_fuse_kv{kv_per_cube}") + v = ctx.zeros((H_KV * S_kv, D_HEAD), + dtype=DTYPE, dp=dp_kv, name=f"v_pre_fuse_kv{kv_per_cube}") + o = ctx.empty((Q_ROWS, Q_COLS), dtype=DTYPE, dp=dp_o, + name=f"o_pre_fuse_kv{kv_per_cube}") + ctx.launch( + f"gqa_prefill_fuse_kv{kv_per_cube}", + gqa_attention_prefill_short_composite_fused_kernel, + q, k, v, o, + T_q, S_kv, h_q, H_KV, D_HEAD, C, P, kv_per_cube, + _auto_dim_remap=False, + ) + + return run_bench( + topology=topo, bench_fn=_bench_fn, + device=resolve_device(None), engine_factory=_engine_factory, + ) + + +def _run_decode_composite_fused(*, kv_per_cube: int, C: int, S_kv: int, + h_q: int = 8): + """Variant (3): GEMM-only composite + softmax_merge prologue fusion.""" + from kernbench.benches.gqa_helpers.short_ctx._gqa_attention_decode_short_composite_fused import ( + _validate_config as _validate_decode_fuse_config, + gqa_attention_decode_short_composite_fused_kernel, + ) + T_q = 1 + _validate_decode_fuse_config(kv_per_cube=kv_per_cube, T_q=T_q, P=P, + C=C, h_q=h_q, h_kv=H_KV, S_kv=S_kv) + Q_ROWS = kv_per_cube * T_q + Q_COLS = (h_q * D_HEAD) // kv_per_cube + k_rows = (H_KV * S_kv * D_HEAD) // TILE_S_KV + topo = resolve_topology(str(TOPOLOGY_DEFAULT)) + + def _bench_fn(ctx): + configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg()) + dp_q = DPPolicy(cube="column_wise", pe="replicate", num_cubes=C, num_pes=P) + dp_kv = DPPolicy(cube="row_wise", pe="row_wise", num_cubes=C, num_pes=P) + dp_o = DPPolicy(cube="column_wise", pe="replicate", num_cubes=C, num_pes=P) + q = ctx.zeros((Q_ROWS, Q_COLS), dtype=DTYPE, dp=dp_q, + name=f"q_dec_fuse_kv{kv_per_cube}") + k = ctx.zeros((k_rows, TILE_S_KV), dtype=DTYPE, dp=dp_kv, + name=f"k_dec_fuse_kv{kv_per_cube}") + v = ctx.zeros((H_KV * S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv, + name=f"v_dec_fuse_kv{kv_per_cube}") + o = ctx.empty((Q_ROWS, Q_COLS), dtype=DTYPE, dp=dp_o, + name=f"o_dec_fuse_kv{kv_per_cube}") + ctx.launch( + f"gqa_decode_fuse_kv{kv_per_cube}", + gqa_attention_decode_short_composite_fused_kernel, + q, k, v, o, + T_q, S_kv, h_q, H_KV, D_HEAD, C, P, kv_per_cube, + _auto_dim_remap=False, + ) + + return run_bench( + topology=topo, bench_fn=_bench_fn, + device=resolve_device(None), engine_factory=_engine_factory, + ) + + @pytest.mark.parametrize("kv_per_cube,C", MODES) def test_prefill_composite_smoke(kv_per_cube, C): - """Composite-GEMM prefill completes in all 4 modes.""" + """Variant (2) GEMM-only composite prefill — all 4 modes.""" r = _run_prefill_composite(kv_per_cube=kv_per_cube, C=C, T_q=8, S_kv=1024) assert r.completion.ok, ( f"prefill-composite kv_per_cube={kv_per_cube}: {r.completion}" @@ -478,8 +561,48 @@ def test_prefill_composite_smoke(kv_per_cube, C): @pytest.mark.parametrize("kv_per_cube,C", MODES) def test_decode_composite_smoke(kv_per_cube, C): - """Composite-GEMM decode completes in all 4 modes.""" + """Variant (2) GEMM-only composite decode — all 4 modes.""" r = _run_decode_composite(kv_per_cube=kv_per_cube, C=C, S_kv=8192) assert r.completion.ok, ( f"decode-composite kv_per_cube={kv_per_cube}: {r.completion}" ) + + +@pytest.mark.parametrize("kv_per_cube,C", MODES) +def test_decode_composite_fused_smoke(kv_per_cube, C): + """Variant (3) composite + softmax_merge fused decode — all 4 modes.""" + r = _run_decode_composite_fused(kv_per_cube=kv_per_cube, C=C, S_kv=8192) + assert r.completion.ok, ( + f"decode-composite-fused kv_per_cube={kv_per_cube}: {r.completion}" + ) + + +@pytest.mark.parametrize("kv_per_cube,C", MODES) +def test_decode_composite_fused_multitile(kv_per_cube, C): + """Fused decode with n_tiles_per_pe == 2 in EVERY mode, so the tile-1+ + fused composite loop actually runs. + + The smoke test at S_kv=8192 gives A1 (group_size=8) n_tiles_per_pe=1, + which never enters the fused path. S_kv = group_size·2·TILE_S_KV forces + 2 tiles per PE for all modes (mirrors test_decode_multitile_per_pe).""" + group_size = P // kv_per_cube + S_kv = group_size * 2 * TILE_S_KV + r = _run_decode_composite_fused(kv_per_cube=kv_per_cube, C=C, S_kv=S_kv) + assert r.completion.ok, ( + f"decode-composite-fused-multitile kv={kv_per_cube} S_kv={S_kv}: " + f"{r.completion}" + ) + + +@pytest.mark.parametrize("kv_per_cube,C", MODES) +def test_prefill_composite_fused_smoke(kv_per_cube, C): + """Variant (3) composite + softmax_merge fused prefill — all 4 modes. + + S_kv=2048 (n_tiles=2) so the tile-1+ fused composite loop actually + runs; S_kv=1024 gives n_tiles=1 and never enters the fused path. + Multi-cube (A1/A2/A4) works now that recv'd K/V slots are pinned.""" + r = _run_prefill_composite_fused(kv_per_cube=kv_per_cube, C=C, + T_q=8, S_kv=2048) + assert r.completion.ok, ( + f"prefill-composite-fused kv_per_cube={kv_per_cube}: {r.completion}" + )