23 KiB
ADR-001: AHBM GQA Fused Attention Kernel (Llama3-70B)
Status: Proposed Context model: Llama3-70B Decision drivers: agentic workload → low batch, long context; KV-load-bound decode; SP (Ring KV) for long-context prefill.
Algorithm lineage. This kernel is FlashAttention (tiling + online/streaming softmax with fused P·V — no full score matrix materialized). The KV-parallel split-and-combine in §4 is FlashDecoding (split-KV with log-sum-exp merge). The SP path in §5.5 is Ring Attention (block-wise KV rotated around the ring, accumulated by the same online softmax). No new math is introduced; this ADR maps those known algorithms onto AHBM composite commands + IPCQ.
0. Reference dimensions (Llama3-70B)
| Symbol | Meaning | Value |
|---|---|---|
H_q |
query heads | 64 |
H_kv |
KV heads | 8 |
G |
GQA group size = H_q / H_kv |
8 |
d |
head dim | 128 |
L |
layers | 80 |
D |
model dim | 8192 |
Hardware recap:
- AHBM (chip) = set of CUBEs (memory cubes, each with a logic die containing PEs) + IO die.
- IPCQ: PE↔PE queues. Each PE holds 4 neighbor queue-pairs (one per mesh direction).
- Composite command: a fused multi-stage command (load → matmul → vector op → ...) executed as one unit. Cannot express a data-dependent
popas an internal dependency (see §7). It can take a precomputed mask tile as an input and apply it as a stage. - Allocation policy (SP): one query head per CUBE. For GQA-8, one KV head's 8 query heads map to 8 CUBEs.
Two orthogonal mapping layers (decided in design discussion):
- Layer 1 (KV-parallel, intra-request): within a request, KV token
ilands on PE(start_pe + i) mod N— round-robin so a single long request is split acrossNPEs and all PEs stay balanced (chunk length differs by ≤1 token). - Layer 2 (inter-request):
start_pe = request_id mod Nrotates the "PE-0 role" per request to spread write/start load. - Combined destination PE for a token:
pe = (request_id + global_token_idx) mod N.
Here N = number of PEs participating in the KV-parallel reduction group for one (query head, request). Within a CUBE the reduction group is the set of PEs splitting that head's sequence; reduction stays intra-CUBE (query head never spans CUBEs — explicit non-goal).
0.5 Kernel boundary, preconditions, and I/O contract
0.5.1 Where this kernel sits in the decoder layer
1. RMSNorm
2. QKV projection (GEMM) ─┐ qkv_rope kernel (SEPARATE, upstream)
3. RoPE on Q and K ─┤
4. write new K,V → KV cache ─┘
5. ===== THIS KERNEL: FlashAttention ===== (post-RoPE Q, post-RoPE K-cache, V-cache → O)
6. Output projection (GEMM) out_proj kernel (SEPARATE, downstream)
7. residual add → FFN ...
Preconditions (responsibilities of the upstream qkv_rope kernel, NOT this kernel):
- P1. Q is already RoPE-rotated. This kernel does no rotation on Q.
- P2. K-cache stores post-RoPE K. Past tiles are never re-rotated at attention time (this is the reason for post-RoPE caching).
- P3. For decode, the new step's K row is RoPE-rotated and appended to its owning PE's K-cache slot by
qkv_ropebefore this kernel launches. For prefill, the whole query block's K is post-RoPE in cache. ⇒ this kernel is pure read on the KV cache; it never writes KV. - P4. V is not rotated (RoPE applies to Q,K only). V-cache holds raw projected V.
- P5. RoPE position used upstream is the token's absolute global position, i.e.
global_idx = local_slot*N + ((pe_id - start_pe) mod N)(§2.1). Round-robin placement must not be confused with RoPE position — the angle depends on global index, never on the local slot.
(If a fused megakernel is ever desired, stages 2–4 would prepend to §5.1; the chosen design keeps them separate so KV cache is the clean interface and Ring Attention only ever passes post-RoPE tensors.)
0.5.2 Symbols for shapes
T_q = query length this launch (decode: 1; prefill/chunk: chunk width). S = total context length (keys seen). S_pe = keys owned by this PE = valid_len (§2.2), ≈ ceil(S/N). bf16 storage, fp32 for m,l,O accumulators.
0.5.3 INPUTS (per kernel launch)
| Input | Shape (per head) | Layout / location | Notes |
|---|---|---|---|
Q |
[G, T_q, d] |
on-chip regs/SRAM of the PE | post-RoPE (P1). G=8 query rows of the GQA group batched together. |
K_cache |
[S_pe, d] |
per-PE HBM buffer, contiguous, base = K_base[pe] |
post-RoPE (P2). Read-only. Strided in global index, dense locally (§2.1). |
V_cache |
[S_pe, d] |
per-PE HBM buffer, base = V_base[pe] |
raw V (P4). Read-only. |
global_token_counter |
scalar | launch arg | kernel derives S_pe, slot↔global, causal bounds. |
start_pe |
scalar = request_id mod N |
launch arg | Layer-2 rotation. |
pe_id, N |
scalars | known to PE / launch | reduction-group geometry. |
q_block_meta |
{q_start, T_q} |
launch arg | prefill/SP causal masking & skip. |
O_base |
address | launch arg | where final O is written (root PE only). |
softmax_scale |
scalar | launch arg | 1/sqrt(d) (=1/√128). |
For Ring Attention (§5.5) add: incoming K_block, V_block per ring step arrive via IPCQ/comms into ping-pong receive buffers (post-RoPE), plus step_kv_global_range so the kernel can apply the causal step-skip.
0.5.4 OUTPUTS
| Output | Shape (per head) | Location | Notes |
|---|---|---|---|
O (final) |
[G, T_q, d] |
O_base, written by tree-root PE only |
normalized: O_acc / l_acc. fp32→bf16 cast on store. |
Intermediate (per non-root PE, on IPCQ, not a kernel-visible output):
| Partial | Shape | Channel | Notes |
|---|---|---|---|
(m_i, l_i, O_i) |
m,l: [G, T_q] scalars; O_i: [G, T_q, d] unnormalized |
pushed to tree_parent via neighbor IPCQ |
log-sum-exp merge payload (§4). O_i is the heavy part. |
No-SP cases (N=1): there are no IPCQ partials; the single PE's running (m,l,O) is normalized in place and written to O_base. The kernel's only output is O.
Explicitly NOT outputs of this kernel: KV cache writes (done upstream, P3), output projection (downstream stage 6), the score matrix S and probabilities P (never materialized — FlashAttention).
1. Decision
Fuse the attention inner pipeline into a single composite per KV tile:
K_j load → Q·Kjᵀ → (+ mask stage) → online-softmax update → V_j load → P·V → running (m,l,O) update → downstream IPCQ push
Cross-PE combination (KV-parallel / SP) is a log-sum-exp tree reduction over (m, l, O) performed after P·V, flowed through IPCQ with a kernel-driven async-submit + bounded-polling loop (no hardware pop-as-dependency change required for the target workload).
Control flow (tile skip, mask generation, reduction scheduling, address arithmetic) lives in the kernel; the composite only executes already-decided work.
2. Memory layout & driver responsibilities
2.1 KV cache allocation
- Pre-allocate per-PE KV buffers sized to per-PE max context =
ceil(max_context / N)tokens ×d× dtype, for K and V separately. - Within a PE, tokens assigned to it are appended contiguously (slot 0,1,2,...). Even though global indices are strided (
i, i+N, i+2N, ...), the per-PE buffer is dense ⇒ DMA stays contiguous, bandwidth not fragmented. - Slot → global index is implicit:
global_idx = local_slot * N + ((pe_id - start_pe) mod N).
2.2 Driver per-step duties (kept minimal)
The driver does not compute write slots or read ranges. It supplies, per kernel launch:
| Launch arg | Purpose |
|---|---|
K_base[pe], V_base[pe] |
base address of each PE's KV buffer |
O_base |
attention output destination (separate from KV write) |
global_token_counter |
current sequence position; kernel derives everything from it |
start_pe (= request_id mod N) |
Layer-2 rotation |
pe_id |
known to PE; used with counter for mod/div |
q_block_meta |
for prefill/SP: query block start + length |
From global_token_counter, start_pe, pe_id the kernel derives:
- "is it my turn to write this step":
(start_pe + counter) mod N == pe_id - my valid length:
base = counter / N; rem = counter % N; my_len = base + ( ((pe_id - start_pe) mod N) < rem ? 1 : 0 ) - write offset:
my_len(next free slot) — kernel computesK_base[pe] + my_len*d. - read range: tiles
0 .. ceil(my_len / TILE). - causal bounds & per-tile skip: from global positions of query block vs each tile.
Driver = address bases + counter + rotation. One formula ((request_id + token_idx) mod N) is the entire placement policy.
3. Composite structure
One composite instance = one KV tile's worth of work on one PE.
COMPOSITE attn_tile(q_reg, K_base, V_base, tile_idx, mask_tile_or_null,
m_reg, l_reg, O_reg, push_target):
S1: DMA K_tile ← K_base + tile_idx*TILE*d # contiguous load
S2: MM S = q_reg · K_tileᵀ (scale 1/sqrt(d))
S3: VEC if mask_tile != null: S += mask_tile # masked boundary tile
S4: VEC online softmax:
m_new = max(m_reg, rowmax(S))
P = exp(S - m_new)
l_reg = l_reg*exp(m_reg - m_new) + rowsum(P)
O_reg = O_reg*exp(m_reg - m_new) # rescale accumulator
m_reg = m_new
S5: DMA V_tile ← V_base + tile_idx*TILE*d # issued early, hidden behind S2-S4
S6: MM O_reg += P · V_tile
S7: (tail) PUSH (m_reg,l_reg,O_reg) → push_target # only on reduction-producing composite
Notes:
m_reg, l_reg, O_regare carried across composites (running state), so consecutive tile composites form an in-order chain on one PE.- S5 (V load) is scheduled so its latency overlaps S2–S4; with prefetch depth ≥ 2 the next tile's K/V are already in flight.
- S7 push is only on the last composite of a PE's tile sweep (when its
O_localis final). Push is a tail action, not a dependency — composites can express push, not pop. - Mask handling: composite only applies a mask tile it is given. The kernel generates the boundary-tile mask from query/KV offsets and passes it in. Full-past tiles get
null; full-future tiles are never enqueued (skipped in kernel).
4. Reduction (KV-parallel / SP combine)
After every PE finishes its tile sweep it holds (m_i, l_i, O_i) (unnormalized). Combine via associative/commutative log-sum-exp merge:
merge((m_a,l_a,O_a),(m_b,l_b,O_b)):
m = max(m_a,m_b)
l = l_a*exp(m_a-m) + l_b*exp(m_b-m)
O = O_a*exp(m_a-m) + O_b*exp(m_b-m)
return (m,l,O)
final O = O_root / l_root # normalize once at tree root
- Timing: exchange is AFTER P·V (each PE must finish its
O_local). - Topology: mesh tree, depth
ceil(log2(N)). ForN=8: level-0 pairs (0↔1,2↔3,4↔5,6↔7), level-1 (1↔3,5↔7), level-2 (3↔7), root = PE matching last-in-group. PE↔mesh-neighbor numbering must be chosen so each tree pair is a physical 4-direction neighbor (tuning item, §9). - Payload:
Ois the heavy part (d-vector = 128 elems);m,lare scalars appended.
5. Kernel pseudocode
5.1 Common loop skeleton (all four cases share this)
def attn_kernel(pe_id, start_pe, counter, K_base, V_base, O_base, q_block):
# ---- derive geometry ----
my_len = valid_len(counter, start_pe, pe_id, N)
n_tiles = ceil(my_len / TILE)
q_reg = load_Q(q_block) # decode: 1 row; prefill: q_block rows
m,l,O = -inf, 0, zeros(d)
PREFETCH = 2
issued = 0
# prime prefetch
for j in range(min(PREFETCH, n_tiles)):
if not tile_all_future(j, q_block): # kernel-level if (§7)
enqueue(build_composite(j, q_block)); issued += 1
# main sweep
for j in range(n_tiles):
wait_complete(j) # in-order; running m,l,O updated
nxt = j + PREFETCH
if nxt < n_tiles and not tile_all_future(nxt, q_block):
enqueue(build_composite(nxt, q_block))
# opportunistic reduction progress (see 5.5) — only when this PE
# already produced O_local AND has spare queue depth
try_drain_reduction()
# ---- reduction phase ----
push_local((m,l,O), to=tree_parent(pe_id))
if is_tree_internal_or_root(pe_id):
reduce_tree_node(pe_id) # async-submit + bounded poll
if is_root(pe_id):
O_final = O_acc / l_acc
store(O_base, O_final)
def build_composite(j, q_block):
if tile_partial(j, q_block): # boundary / diagonal tile
mask = make_mask_tile(j, q_block) # kernel computes from offsets
else: # tile_all_past
mask = null
return composite.attn_tile(q_reg, K_base, V_base, j, mask, m,l,O, push_target)
tile_all_future / tile_partial / tile_all_past are pure position comparisons (query global pos vs tile's KV global-pos range). Decode degenerates them (see below).
5.2 DECODE, no SP (single PE owns the whole head's KV)
N = 1for the reduction group → no cross-PE merge, no tree.- Query length = 1; it attends to all past KV ⇒
tile_all_futurenever true,tile_partialonly on the final ragged tile, no real masking. - GQA reuse: the
G=8query heads sharing one KV head are processed in the same PE/composite stream so each K/V tile is loaded once and reused by 8 query rows (Q·Kᵀ becomes an 8×TILE GEMV-batch, P·V an 8×d). This is the main lever; decode is KV-load-bound, so amortizing K/V load over the group is where the win is. - Bottleneck = K/V tile DMA. Tune
PREFETCH(≥2) so the queue never starves.
load q_reg = [8 query rows for this KV head] # GQA group batched
for j in 0..n_tiles:
composite.attn_tile(...) # 8×TILE Q·Kᵀ, mask=null (except last)
store O[8 rows] # no reduction
5.3 DECODE, with SP / KV-parallel (N=8, head split across 8 PEs)
- One request's KV is round-robin across 8 PEs; each PE owns ~
my_lentokens (≤1 apart). - Each PE still batches its
G=8query rows (GQA) over its local tiles. - Query length 1 ⇒ no future tiles, minimal mask ⇒ tile sweep is short per PE; reduction is the structurally interesting part, done once per token at sweep end.
- Reduction latency for a single token is hidden by the remaining tiles of other concurrent decode tokens if batched; for strict batch=1 long-context, hidden by the fact that each PE still has many tiles (long context) — sweep work >> reduction handshake.
per PE:
my tiles sweep (short, GQA-batched) → (m_i,l_i,O_i)
push to tree_parent
tree merge (depth 3) → root normalizes → O[8 rows] for this KV head
5.4 PREFILL, no SP
- Whole prompt resident on the PE(s) for the head; query is a block of
Ttokens (chunked). - Now causality is real: for query block
[qs, qe)and KV tile covering[ks, ke):ke <= qs→tile_all_past→ mask=null, full compute.ks >= qe→tile_all_future→ skip enqueue (kernelif).- overlap →
tile_partial→ kernel builds triangular mask tile, composite applies it.
- Chunked prefill walks query blocks; for each query block runs the tile sweep with the skip/mask decisions above.
- GQA: 8 query heads of a group share K/V tiles within the block ⇒ batched along the same axis as decode.
for q_block in chunks(prompt):
m,l,O = reset
for j in tiles_up_to(q_block.end):
if tile_all_future(j,q_block): continue # skip
composite.attn_tile(..., mask=mask_or_null(j,q_block))
store O[q_block]
5.5 PREFILL, with SP (Ring KV)
- KV is sharded across
NPEs (the SP ring); each ring step delivers a different KV block from a peer. - Composite is instantiated per ring step's KV block:
K load → Q·Kᵀ → mask → online-softmax update → V → P·V → running (m,l,O). Running(m,l,O)is carried across ring steps. - IPCQ overlaps next step's KV receive (comms) with current step's compute. KV and V of a step arrive together (it's comms, not local DMA) ⇒ V is already present when P·V runs.
- Causal ring optimization via kernel
if: if the incoming step's KV block is entirely after the local query block, the whole step's compute is skipped (don't enqueue the composite) — can eliminate ~half the ring steps for causal attention. - Receive buffers ping-pong; buffer swap aligned to composite boundary.
- After the ring completes,
(m,l,O)is final per PE → same log-sum-exp tree reduction to produce O. (If the ring already linearly accumulated the full sequence on each query-owning PE, the "reduction" is just the running state; if KV-parallel splits remain, tree-merge them.)
init m,l,O
recv KV_block[0]
for step in 0..N-1:
issue_recv(KV_block[step+1]) # overlap comms
if step_all_future(step, q_block): # causal ring skip
continue
mask = step_mask(step, q_block) # null if fully past
composite.attn_tile(q_reg, Kbuf[step%2], Vbuf[step%2], 0, mask, m,l,O, _)
swap buffers each step
# m,l,O now final for this PE's query block
reduce_tree → normalize → store O
5.6 Reduction drain loop (the no-hardware-change mechanism)
Because composites cannot pop as a dependency, the kernel does it explicitly. The poll never sits on the critical path because there is always other tile work to enqueue (long context):
def reduce_tree_node(pe_id):
children = tree_children(pe_id)
acc = local_state
pending = set(children)
while pending:
for c in list(pending):
if ipcq_nonempty(dir_to(c)): # poll ONLY the child direction
part = ipcq_pop(dir_to(c))
acc = merge(acc, part)
pending.remove(c)
if pending:
# miss → do useful work instead of busy-wait
if has_unissued_tiles(): enqueue(next_tile_composite())
else: spin_short() # rare: only batch=1 + short ctx
if not is_root(pe_id):
ipcq_push(dir_to(parent), acc)
else:
global_acc = acc
6. Why no hardware change for the target workload
- The only place a HW
pop-as-dependency helps is a lone reduction with no other work to hide the poll — i.e. batch=1 and short context. - Target is agentic = low batch but long context ⇒ each PE has many KV tiles; a reduction poll-miss is always covered by issuing the next tile composite. The handshake cost is hidden, not on the critical path.
- HW
popwould also remove the kernel's ability to choose "do another tile instead of waiting" and complicate buffer-lifetime/deadlock reasoning (a blocked composite holds PE resources). For long-context that flexibility is worth more than the saved handshake. - Decision: ship with software async-submit + bounded poll. Revisit HW
poponly if a short-context, single-stream, latency-critical target emerges.
7. Control vs execution split (the load-bearing principle)
| Concern | Owner |
|---|---|
| tile skip (future), mask generation, causal bounds | kernel (if, arithmetic) |
| address/offset/valid-length arithmetic | kernel (from counter) |
| reduction scheduling, IPCQ poll/push timing | kernel |
| K/V load, Q·Kᵀ, mask-apply, softmax, P·V, push | composite (no internal branch) |
Composites never branch on data and never pop. They receive fully-decided inputs (which tile, which mask, where to push). This is what lets the design avoid HW changes.
8. Required changes to the composite command implementation
- Multi-stage fusion across DMA + 2×MM + VEC in one command. Current composites must support the chain
DMA(K) → MM → VEC(mask+softmax) → DMA(V) → MM(P·V) → VEC(accumulate)as a single fused unit, with the two DMAs schedulable at different dependency points (K early, V mid) rather than both up front. - Carried register state across composite instances.
(m, l, O)accumulators must persist between consecutive composites on a PE (running flash state) instead of being command-local. Needs a stable register/SRAM binding the kernel passes in and out. - Mask tile as an input operand. Composite must accept an optional mask tile and apply it as a VEC stage (
S += mask). No mask generation inside the composite. - Tail push action to a named IPCQ direction. Composite must be able to, as a terminal action, write
(m,l,O)to a neighbor queue-pair identified by direction. Push only — pop remains a kernel operation (no internal pop-dependency). - Early/decoupled V DMA scheduling. The V load stage must be issuable so its latency overlaps the preceding MM/VEC stages (prefetch depth ≥ 2 across instances), not serialized after softmax.
- GQA batching on the Q axis. Q·Kᵀ and P·V stages must accept a batched Q of
G=8rows so one K/V tile load serves the whole query group (single KV-head reuse). - (NOT required) internal data-dependent branch / pop-as-dependency. Explicitly out of scope; kernel handles all branching. Listed so reviewers don't add it speculatively.
- (NOT required here) RoPE / QKV projection. Handled by the upstream
qkv_ropekernel (§0.5). This kernel assumes post-RoPE Q and post-RoPE K-cache and performs no rotation and no KV-cache write. Listed so reviewers don't fold RoPE into the attention composite (doing so would force re-rotating past tiles every decode step and break Ring Attention's post-RoPE pass-through).
9. Open tuning items (not blocking; measured in Kernbench)
- KV tile prefetch depth — dominant for KV-load-bound decode/prefill; sweep 2→4.
- PE↔mesh-neighbor mapping for the reduction tree — ensure each depth-3 tree pair is a physical 4-direction neighbor; bad mapping adds hops.
- TILE size — balance SRAM residency (K_tile + S/P + V_tile + O_acc + 8-way GQA) against DMA efficiency.
- Ring buffer ping-pong vs prefetch depth interaction in §5.5.
10. Coverage summary
| Case | KV placement | Composite unit | Reduction | Masking |
|---|---|---|---|---|
| Decode, no SP | 1 PE, all KV | per tile, GQA-batched | none | last tile only |
| Decode, SP | round-robin N PEs | per tile, GQA-batched | tree (depth 3) | last tile only |
| Prefill, no SP | resident | per tile per q-block | none | skip future / triangular boundary |
| Prefill, SP (Ring) | ring-sharded | per ring step | running state + tree | causal step-skip + boundary |
All four share the §5.1 skeleton and the §3 composite; they differ only in N, query-block width, and which tile_* predicates fire.
I/O per case (see §0.5 for full contract):
| Case | Inputs | Output | Partials on IPCQ |
|---|---|---|---|
| Decode, no SP | Q[G,1,d] post-RoPE, full K_cache[S,d], V_cache[S,d] on 1 PE |
O[G,1,d] at O_base |
none |
| Decode, SP | Q[G,1,d], per-PE K_cache[S_pe,d], V_cache[S_pe,d] |
O[G,1,d] (root PE) |
(m,l,O_i) per PE → tree |
| Prefill, no SP | Q[G,T_q,d] post-RoPE, K_cache[≤end,d], V_cache |
O[G,T_q,d] at O_base |
none |
| Prefill, SP (Ring) | Q[G,T_q,d], ring-delivered K_block,V_block (post-RoPE) per step |
O[G,T_q,d] (root PE) |
running (m,l,O) + tree partials |
In all cases: KV-cache writes and RoPE happen upstream; output projection happens downstream; score S and probs P are never materialized.