Files
kernbench2/docs/adr-proposed/ADR-0070-algo-gqa-short-context-attention.md
T

10 KiB
Raw Blame History

ADR-0070: GQA Short-Context Attention — Unified A1/A2/A4/B Mapping for Prefill and Decode

Status

Proposed — short-context (single tile to a handful of KV tiles) GQA attention benchmark covering both prefill and decode in four mapping modes (A1/A2/A4/B) and three composite tiers (without composite, GEMM-only composite, composite + softmax_merge fused). Compute-only attention — weight DMA, ITL batching, and end-to-end LLM driver scope are out of scope (separate ADRs).

Context

A GQA layer at the LLaMA-3.1-70B headline shape (h_q=64, h_kv=8, d_head=128, GQA group G = h_q/h_kv = 8) on a 4×4 cube SIP can be mapped to AHBM in several ways that trade KV-cache memory, HBM read volume per cube, and intra-cube IPCQ traffic against PE parallelism.

Bench-vs-headline shape. The kernels and tests use a d_head = 64 proxy (tests/attention/test_gqa_short_context.py D_HEAD = 64) so the per-tile working set fits comfortably in scratch across all four modes. The mapping decisions are d_head-independent — only the per-PE GEMM tile bytes scale.

The four candidate mappings (per GQA-mapping-on-AHBM-short-context note §1):

Mode  kv_per_cube  C        group_size  Reduce / broadcast topology
----  -----------  -------  ----------  ----------------------------------
A1            1    h_kv      P (=8)     row chain + col bridge (2×4 mesh)
A2            2    h_kv/2    P/2 (=4)   row chain only
A4            4    h_kv/4    P/4 (=2)   single intra_W / intra_E hop
B             8    1         1          no broadcast / no reduce (single PE)

Within each cube the 8 PEs split into kv_per_cube groups of group_size = P/kv_per_cube PEs; each group owns one KV head.

The kernel must be able to compare these mappings at multiple context lengths (8K..64K KV tokens) and across the three composite tiers so the team can size the per-cube HBM / IPCQ / GEMM tradeoff empirically.

Decision

1. Mapping (unified A1/A2/A4/B)

Two unified kernels select mode at launch via kv_per_cube ∈ {1,2,4,8}:

  • Prefill (gqa_attention_prefill_short_kernel):

    • Q-tile split — T_q rows floor-balanced across the group's group_size PEs.
    • FA2 head fusion — G Q heads fused into the M dimension of one batched GEMM per PE per tile.
    • IPCQ KV broadcast — group root (pe_in_group == 0) loads K/V from HBM and IPCQ-broadcasts each tile to the rest of its group; q-tiles are independent so no intra-group reduce.
  • Decode (gqa_attention_decode_short_kernel):

    • Sequence-shard — each PE owns S_local = S_kv/group_size tokens of the group's KV head.
    • FA2 head fusion — same as prefill.
    • IPCQ chain reduce — per-PE partial (m, , O) chain-reduce up to group root (PE 0), which normalizes and stores.

Geometry (2×4 mesh):

group_size = 8  : 2 row × 4 col   (A1)
group_size = 4  : 1 row × 4 col   (A2)
group_size = 2  : 1 row × 2 col   (A4)
group_size = 1  : single PE       (B)

2. Shard addressing (ADR-0011 D-VA1 contract)

Deploy is per-(sip, cube, pe) but tl.load receives a single global VA per tensor. The kernel computes its own shard base offset from tl.program_id(axis=0) (PE id) and tl.program_id(axis=1) (cube id):

# Decode shape (pe=row_wise): K is split across PEs by sequence shard.
cube_K_base       = cube_id        * kv_per_cube * K_HEAD_BYTES
head_K_base       = group_id_in_cube * K_HEAD_BYTES
pe_K_seq_offset   = pe_in_group     * n_tiles_per_pe * K_TILE_BYTES
k_shard_base      = k_ptr + cube_K_base + head_K_base + pe_K_seq_offset

For prefill the K/V dp is pe=replicate (group root reads and broadcasts) so the pe_in_group term drops out: k_head_shard_base = k_ptr + cube_K_base + head_K_base.

Skipping cube_id collapses all cubes onto cube 0's HBM region (observed pre-fix as an 11.5× per-cube DMA imbalance).

3. Three composite tiers

The same mapping is exercised against three GEMM/MATH dispatch styles so the contribution of the composite API vs the recipe-driven fusion can be isolated:

(1) without composite                  primitives only (tl.dot, tl.exp, …)
(2) with composite (GEMM-only)         tl.composite(op="gemm")
(3) with composite + softmax_merge     tl.composite(prologue=[softmax_merge],
                                                    op="gemm", out=O,
                                                    epilogue=[add])

File layout (src/kernbench/benches/gqa_helpers/short_ctx/):

_gqa_attention_{prefill,decode}_short.py                   (1)
_gqa_attention_{prefill,decode}_short_composite.py         (2)
_gqa_attention_{prefill,decode}_short_composite_fused.py   (3)

Tier (2) is GEMM-only by definition — no recipe-driven fusion, so both prefill's and decode's P·V stay tl.dot. Recipe-driven fusion is exactly what tier (3) adds.

Tier (3) on multi-cube modes (A1/A2/A4) of prefill relies on the D4 supplement: any composite operand that is an IPCQ recv'd slot (non-root PE's K_T in Q·Kᵀ, V in P·V) is pinned and read in place instead of being DMA-streamed from HBM. Without the supplement every recv slot fed to a composite — both Q·Kᵀ's b=K_T and P·V's b=V — PageFaults on PA decode (PE scratch addresses overflow the 51-bit PA range).

4. Caller contract — _validate_config

Each kernel module exposes a single _validate_config(...) helper called by the bench wrapper before launch. The kernel itself is lean (no inline if guards): the contract block is enforced caller-side, sim cost zero, but catches every silent-shard-corruption foot-gun (kv_per_cube=3 → group_size=2 via integer division, non-integer GQA group, mismatched cube count, partial-tile S_kv, …) before any address arithmetic runs.

5. Tensor layouts (host-side, mode-invariant byte totals)

  • Q: (kv_per_cube·T_q, h_q·d_head/kv_per_cube) dp=(cube=column_wise, pe=replicate) over C = h_kv/kv_per_cube.
  • K: (h_kv·S_kv·d_head/TILE_S_KV, TILE_S_KV) tile-major. Decode: dp=(cube=row_wise, pe=row_wise). Prefill: dp=(cube=row_wise, pe=replicate) (broadcast model).
  • V: (h_kv·S_kv, d_head) native. Same dp as K.
  • O: same dp as Q.

Caller pre-scales Q by 1/√d_head.

Verification

Smoke + regression (tests/attention/test_gqa_short_context.py)

  • Prefill / decode smoke (all 4 modes, single-tile S_kv).
  • Multi-tile coverage (all 4 modes, S_kv chosen so each PE owns ≥2 tiles).
  • Op-count invariants — prefill dma_writes = group_size·kv_per_cube·C, dma_reads = P·C + 2·kv_per_cube·C·n_tiles; decode dma_writes = kv_per_cube·C.
  • IPCQ topology — short-context kernel never emits inter-CUBE E/W IPCQ.
  • ADR-0011 D-VA1 regressionper_cube_disjoint_src_addrs, per_cube_dma_balanced (max/min DMA busy ratio < 1.1×) per mode per phase.

Composite smoke

  • test_{prefill,decode}_composite_smoke[A1/A2/A4/B] — tier (2), 4 modes × 2 phases.
  • test_{prefill,decode}_composite_fused_smoke[A1/A2/A4/B] — tier (3), 4 modes × 2 phases.
  • test_decode_composite_fused_multitile[A1/A2/A4/B]n_tiles_per_pe == 2 per mode so the recipe-fused tile loop actually executes.

Mode × context sweep (tests/attention/test_gqa_short_context_sweep_*.py)

Six sweep files (3 tiers × {prefill, decode}) measure each (mode, S_kv) cell and dump the same metric set to CSV:

wall_us, n_pe, gemm_util, math_count, math_pipeline_us, hbm_bw_util,
hbm_read_mb, hbm_write_kb, ipcq_kb, kv_cache_per_cube_mb

S_kv ∈ {8K, 16K, 32K, 64K}; outputs under docs/sweeps/.

Composite/fused sweeps count both gemm_f16 (non-pipeline) and TileToken/GEMM (pipeline composite path); fused additionally counts TileToken/MATH. Baseline sweeps count gemm_f16 only.

Known limitations

  1. Causal mask — kernel is non-causal. Adding causal masking is orthogonal to mapping and is scheduled as a separate ADR.
  2. f16 online-softmax accumulator(m, , O) are f16 throughout. Long-context numerical drift will need an f32 accumulator before correctness-grade use.
  3. End-to-end numeric validation — current tests are op-count / topology / per-cube-DMA invariants. No full-kernel np.allclose against a torch/numpy reference yet.
  4. Skinny-M GEMM underfill — decode (T_q=1) gives M = G = 8, well below PE_SCHEDULER supertile TILE_M = 32. The two GEMM paths handle this differently:
    • Primitive tl.dot dispatches the GEMM at the actual m = 8 (no padding); GEMM time reflects the real 8-row work.
    • tl.composite(op="gemm") tiles at TILE_M = 32, padding M 4× with zeros; GEMM time reflects the padded 32-row work. Consequence: when sweep CSVs show tier (2)/(3) gemm_util higher than tier (1), most of that gap is supertile padding overhead, not extra useful work or fusion savings — read it as "composite tile shape doesn't match decode-skinny shape" rather than "composite/fusion is more compute-intensive." This is the intended decode shape for this benchmark; lifting it requires batched-M (multi-request inference, B_sys > 1), tracked as a separate batched variant — see "Future work".
  5. Multi-cube prefill composite without the ADR-0065 D4 supplement PageFaults. The supplement — IPCQ recv'd slots are pinned and read in place as composite operands — ships alongside this ADR.

Future work

  • Batched-M variant (B_sys ∈ {1,2,4,8,16}) so composite/fused pipeline overlap shows up in wall_clock rather than only in per-engine utilization.
  • Long-context (S_kv ≥ 128K) sweep extension for the headline LLaMA decode target.
  • Causal-mask + f32 accumulator promotion.
  • 3-variant comparison plots from the sweep CSVs (wall, gemm util, hbm bw util, ipcq, kv cache) — generated under docs/diagrams/.

References

  • ADR-0011 — PhysAddr + VA contract (D-VA1: kernel-computed shard offset).
  • ADR-0060 — GQA fused attention on AHBM (precursor of the unified mapping).
  • ADR-0064 — CPU issue-cost model (composite supertile motivation).
  • ADR-0065 — Flat ops, composite, softmax_merge recipe (with the D4 supplement applied here).
  • ADR-0070 supersedes ADR-0060 §B.split.2 short-prefill clause.
  • GQA-mapping-on-AHBM-short-context note (research design rationale).