Files
kernbench2/docs/adr-proposed/ADR-0060-algo-gqa-fused-attention-ahbm-ko.md
T
ywkang 7a9d4ec47b gqa(adr): pivot ADR-0060 to composite hybrid + lazy tl.load; add ADR-0064 cost model
- ADR-0060: GEMMs (Q.Kt, P.V) via existing tl.composite (scheduler-managed
  tiling + K/V DMA streaming); softmax merge + IPCQ tree reduction stay in
  kernel. Front TL;DR pseudocode of the final composite kernel; new section
  B lists open design items (DDD sync, K pre-transpose, dma_read lever,
  kernel-vs-scheduler tiling, ring path).
- ADR-0062: redefined from a new load_async op to global lazy tl.load
  (non-blocking + auto-wait on first use; API unchanged; goldens regenerate).
- ADR-0064 (new): per-op-type CPU issue cost model (composite ~40ns >>
  primitive) so the hybrid's CPU-saturation win becomes measurable
  (currently dispatch_cycles=0 hides it). Cost-model impl deferred.
- KO mirrors for ADR-0060/0062/0064 (-ko suffix, adr-proposed).

Rationale: non-blocking CompositeCmd offloads tiling to PE_SCHEDULER,
decoupling CPU issue-rate from execution so the CPU can saturate the
engines; the prior 'composite = no latency benefit' claim was an artifact
of dispatch_cycles=0. Docs only; no production code changed.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-04 10:19:37 -07:00

42 KiB
Raw Blame History

ADR-0060: AHBM GQA Fused Attention 커널 (Llama3-70B)

Status

Proposed

Context model: Llama3-70B. Decision drivers: agentic workload → 낮은 batch, 긴 context; KV-load-bound decode; long-context prefill용 sequence-parallel (Ring KV).

Supersedes / extends: 기존 mesh-native 어텐션 커널 _attention_mesh_kv(prefill)와 _attention_mesh_mlo(decode), 그리고 milestone-gqa-llama70b eval bench. §A. 기존 kernbench 작업과의 관계 참조 — 그 코드가 본 ADR이 진짜 GQA·causal·long-context 커널로 업그레이드하는 baseline이다.

Supporting ADRs (efficiency / scale enabler — GQA blocker 아님; §8 정정 참조): ADR-0063 tl.scratch_scope(per-tile scratch 재활용 — 현실적 context 길이에 필요), ADR-0062 lazy tl.load(첫 사용 시점 auto-wait를 갖는 non-blocking load → load/compute 오버랩), ADR-0061 tl.broadcast(선택적 mask/범용 편의). 두 GEMM(Q·Kᵀ, P·V)은 scheduler가 관리하는 tl.composite 커맨드로 발행된다(기존 CompositeCmd; 새 커맨드 종류 없음); 진짜 GQA 자체는 커널 재구조화만 필요(§5.2).

Algorithm lineage. 이 커널은 FlashAttention(tiling + online/streaming softmax, P·V fused — full score matrix 미생성)이다. §4의 KV-parallel split-and-combine는 FlashDecoding(split-KV + log-sum-exp merge). §5.5의 Ring 경로는 Ring Attention(KV 블록을 mesh 주위로 회전, 같은 online softmax로 fold). 새 수학은 도입하지 않으며, 본 ADR은 이 알려진 알고리즘을 kernbench의 greenlet tl 프로그래밍 모델(ADR-0020, ADR-0046)과 IPCQ PE↔PE collective(ADR-0023/0025)에 매핑한다.


TL;DR — 최종 GQA 커널 (composite hybrid) pseudocode

결정(§1)을 한 곳에: GEMM → tl.composite; softmax 머지 + tree reduction → 커널; tl.load은 lazy. KV head별로, G를 matmul M 차원에 fold(진짜 GQA, broadcast 없음). 이것이 reference 형태이며, 산문 섹션이 각 부분을 부연한다.

def gqa_attention(q_ptr, k_ptr, v_ptr, o_ptr,
                  counter, start_pe, N, q_block, scale, *, tl):
    # ---- geometry (kernel arithmetic, §2) ----
    pe_id   = tl.program_id(axis=rank_axis)
    G       = H_q // H_kv                       # query heads per KV head (=8)
    causal  = q_block.is_prefill

    for kv in range(H_kv):                       # one KV head per iteration (§5.2)
        # Q group: G rows folded into M → [G·T_q, d]. Lazy load (ADR-0062):
        # issue now, auto-wait at first use inside the first composite.
        q_g     = tl.load(q_base(kv), (G * q_block.T_q, d))      # [G·T_q, d]
        my_len  = valid_len(counter, start_pe, pe_id, N) if not causal else S_kv
        n_tiles = ceil(my_len / TILE)

        # persistent arena (outside scratch_scope, ADR-0063): -inf, 0, zeros
        m, l, O = init_running(G * q_block.T_q, d)

        for j in range(n_tiles):
            if causal and tile_all_future(j, q_block):
                continue                                         # causal skip (kernel if)
            with tl.scratch_scope():                             # per-tile temporaries (ADR-0063)
                # --- GEMM #1 on the scheduler: Q·Kⱼᵀ; K streamed by composite ---
                Sj = tl.composite("gemm", a=q_g,
                                  b=tl.ref(k_tile(kv, j), (d, TILE))) * scale   # Kᵀ pre-stored [d,TILE]
                if causal and tile_partial(j, q_block):
                    Sj = Sj + causal_mask(j, q_block)            # additive boundary mask
                # --- online softmax merge in the kernel (MATH ops) ---
                m_new = tl.maximum(m, tl.max(Sj, axis=-1))
                P     = tl.exp(Sj - m_new)
                corr  = tl.exp(m - m_new)
                l     = l * corr + tl.sum(P, axis=-1)
                # --- GEMM #2 on the scheduler: P·Vⱼ; V streamed by composite ---
                Oj    = tl.composite("gemm", a=P,
                                     b=tl.ref(v_tile(kv, j), (TILE, d)))
                O     = O * corr + Oj                            # running merge
                m     = m_new

        # ---- cross-PE combine (§4): log-sum-exp tree to root, or store ----
        if N == 1:
            tl.store(o_base(kv), O / l)
        else:
            tree_reduce_and_store(m, l, O, pe_id, N, o_base(kv))  # tl.send/tl.recv

이 형태인 이유:tl.composite 호출이 tiling + K/V DMA 스트리밍 + cross-tile 파이프라이닝을 PE_SCHEDULER에 offload한다(CPU는 coarse descriptor를 내고 앞서 나가 → 엔진이 saturate 유지, §1); softmax 머지와 IPCQ tree reduction은 커널에 남는다(기존 CompositeCmd가 cross-tile 상태를 못 들기 때문). K는 reshape-not-transpose caveat를 피하려 transpose된 [d, TILE]로 미리 저장된다 (§3, §B). 전용 "flash-composite" 커맨드 종류는 도입하지 않는다(§8 항목 4).


A. 기존 kernbench 작업과의 관계 (먼저 읽을 것)

kernbench는 오늘날 이미 IPCQ 상에서 online-softmax (m, , O) 머지로 FlashAttention을 돌린다. 두 커널이 존재한다:

File Role Mechanism
src/kernbench/benches/_attention_mesh_kv.py prefill (Ring K/V) per-rank partial attention, bidirectional K/V fan-out, online-softmax fold
src/kernbench/benches/_attention_mesh_mlo.py decode (split-KV) per-rank one-shot partial attention, bidirectional (m,,O) fan-out, log-sum-exp merge

둘 다 milestone-gqa-llama70b(4 패널: {single,multi}_user × {prefill,decode}, src/kernbench/benches/milestone_gqa_llama70b.py)이 구동하며 tests/attention/test_milestone_gqa_llama70b.py에서 테스트된다.

이들은 greenlet tl API로 작성되었다: tl.load, tl.dot, tl.softmax/tl.max/tl.sum/tl.exp, tl.send/tl.recv, 그리고 TensorHandle에 대한 Python -/*//(각각 MathCmd emit) — GEMM은 composite가 아니라 **blocking tl.dot**로. 실행 중 (m, , O)는 루프를 관통하는 Python TensorHandle일 뿐이다. 본 ADR은 실행 상태와 softmax 머지는 커널에 유지하되 두 GEMM은 scheduler가 관리하는 tl.composite 경로로 옮긴다(§1 참조) — 이것이 설계에 중요하다.

baseline의 의도된 세 한계 — 효율적 GQA 커널이 정확히 들어내야 하는 것:

  1. GQA 재사용 없음. h_q == h_kv == 1 (test_milestone_gqa_llama70b.py:137-142). 테스트는 이를 broadcast view에 대한 MemoryStore byte-conservation 실패로 돌리지만, 그 실패는 baseline의 head-packing 핵의 속성이다(_view(K, (h_q·d, S_kv))는 모든 head를 하나의 matmul 차원에 뭉치고 h_q == h_kv일 때만 byte를 보존). 올바른 수정은 broadcast op이 아니라 커널 재구조화다: 한 번에 한 KV head를 처리하고 G group 행을 matmul M 차원에 fold(§5.2). 이는 byte-보존 reshape만 쓰므로 진짜 GQA(h_q = G·h_kv)가 신규 primitive 없이 돈다 — §8 참조.
  2. O(N) reduction. baseline은 all-to-all bidirectional fan-out을 하여 모든 rank가 full 답을 갖는다(n_ranks 1 단계). 어텐션은 query owner에서만 O가 필요 → root로의 tree reduction⌈log₂ N⌉ 단계(§4).
  3. 검증 스케일만. S = 16인 이유는 1 MiB scratch bump allocator가 per-tile 임시값을 누수하고(test_milestone_gqa_llama70b.py:123-148) causal masking / tiling이 없기 때문 → ADR-0063(재활용) + §5(tiling, causal skip) + composite K/V 스트리밍(§3) + ADR-0062(lazy load 오버랩)으로 해결.

문서 부채(범위 밖이나 기록): baseline은 ADR-0055/0056/0057/0058/0059를 인용하나 파일로 존재하지 않는다 — ghost 참조다. 본 ADR은 이를 소급 작성하지 않으며; 권고는 Detailed Design Document의 Open Decisions 참조.


0. 참조 차원 (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

하드웨어 요약:

  • AHBM (chip) = CUBE(메모리 cube, 각각 PE를 담은 logic die) 집합 + IO die(ADR-0003).
  • IPCQ: PE↔PE 큐, PE당 4 mesh-방향 queue-pair (N/S/E/W, ADR-0023 D3; inter-SIP용 global_*, ADR-0032). 커널 API: tl.send(dir, src) / tl.recv(dir, shape, dtype) (tl_context.py:402-499).
  • Composite command (CompositeCmd, pe_commands.py:144-162): 단일 GEMM(또는 MATH) head + element-wise epilogue 단계 (bias/relu/scale/add/...), PE_SCHEDULER에 non-blocking으로 발행되며, scheduler가 tile plan을 생성하고 타일당 DMA→GEMM→write를 스트리밍한다 (ADR-0014 D6; pe_scheduler.py:104-143). 이는 일반적 multi-op DAG가 아니다: 두 GEMM을 chain할 수 없고, 인스턴스 간 register 상태를 못 들며, IPCQ를 pop/wait 못 한다. 따라서 본 ADR은 두 어텐션 GEMM(Q·Kᵀ, P·V)을 각각 자체 composite로 발행하고 cross-GEMM softmax 머지 + IPCQ reduction은 커널에 유지한다 — 새 "flash-composite" 커맨드 종류는 필요 없다(§1, §8 참조).
  • Allocation policy (SP): multi-user 패널에서 CUBE당 query head 하나; KV group의 G query head는 한 CUBE의 PE 안에 매핑.

직교하는 두 매핑 레이어 (round-robin placement):

  • Layer 1 (KV-parallel, intra-request): KV token i는 PE (start_pe + i) mod N에 안착 — round-robin이라 한 긴 request가 N PE에 걸쳐 분할, ≤1 token으로 균형.
  • Layer 2 (inter-request): start_pe = request_id mod N이 request별 "PE-0 역할"을 회전.
  • 결합: pe = (request_id + global_token_idx) mod N.

N = 하나의 (query head, request)에 대한 KV-parallel reduction group의 PE 수. Reduction은 intra-CUBE 유지(query head는 CUBE를 가로지르지 않음 — 명시적 non-goal).


0.5 커널 경계, 전제조건, I/O 계약

0.5.1 디코더 레이어 내 위치

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, K-cache, V-cache → O)
6. Output projection (GEMM)        out_proj kernel (SEPARATE, downstream)
7. residual add → FFN ...

전제조건 (upstream qkv_rope, 본 커널 아님):

  • P1. Q는 이미 RoPE-회전됨. 여기서 회전 없음.
  • P2. K-cache는 post-RoPE K 저장(어텐션 시점에 재회전 안 함 — post-RoPE 캐싱의 이유).
  • P3. decode의 경우, 새 step의 K 행은 RoPE-회전되어 그 소유 PE의 K-cache 슬롯에 본 커널 launch 전에 qkv_rope 추가한다. ⇒ 본 커널은 KV cache에 대해 pure read.
  • P4. V는 회전 안 함; V-cache는 raw projected V 보유.
  • P5. upstream RoPE 위치는 token의 절대 global 위치 global_idx = local_slot·N + ((pe_id start_pe) mod N)(§2.1). Round-robin placement ≠ RoPE 위치.

0.5.2 Shape 기호

T_q = 이번 launch의 query 길이(decode: 1; prefill/chunk: chunk 폭). S = 전체 context 길이. S_pe = 이 PE 소유 key ≈ ⌈S/N⌉. 저장 bf16(numpy proxy f16, memory_store.py:16); m,,O 누산기는 정밀도가 중요한 곳에서 f32.

0.5.3 INPUTS (커널 launch당)

Input Shape (per KV head) Location Notes
Q [G, T_q, d] per-PE HBM (loaded to TCM) post-RoPE (P1). group의 G query 행 batched.
K_cache [S_pe, d] per-PE HBM, base K_base[pe], contiguous post-RoPE (P2). Read-only. 로컬은 dense, global index는 strided (§2.1).
V_cache [S_pe, d] per-PE HBM, base V_base[pe] raw V (P4). Read-only.
global_token_counter scalar launch arg 커널이 S_pe, slot↔global, causal bound 도출.
start_pe (= request_id mod N) scalar launch arg Layer-2 회전.
pe_id, N scalars launch / tl.program_id reduction-group geometry.
q_block_meta {q_start, T_q} launch arg prefill/SP causal masking & skip.
O_base address launch arg 최종 O가 쓰이는 곳 (root PE만).
softmax_scale scalar launch arg 1/√d.

**Ring Attention (§5.5)**에는 ring step마다 추가: IPCQ로 들어오는 K_block, V_block을 ping-pong 버퍼에(post-RoPE), 그리고 causal step-skip용 step_kv_global_range.

0.5.4 OUTPUTS

Output Shape Location Notes
O (final) [G, T_q, d] O_base, root PE만 정규화 O_acc / _acc, 저장 dtype으로 cast.

중간값 (non-root PE, IPCQ 상 — 커널-가시 출력 아님): (m_i, _i, O_i) (m,: [G, T_q]; O_i: [G, T_q, d] 비정규화)이 tl.sendtree_parent에 push(§4). O_i가 무거운 부분.

No-SP (N=1): IPCQ partial 없음; 단일 PE의 실행 중 (m,,O)를 제자리에서 정규화하여 O_base에 기록.

명시적 비-출력: KV-cache write(upstream, P3), output projection(downstream), score S 및 prob P(미생성).


1. 결정 (메커니즘)

커널을 하이브리드로 구현한다: 두 GEMM(Q·Kᵀ, P·V)을 scheduler가 관리하는 tl.composite(op="gemm") 커맨드로 발행하고, online-softmax 머지와 cross-PE reduction은 커널 수준 tl op으로 유지한다. tl.loadlazy다(non-blocking; wait는 load된 데이터의 첫 사용 시점에 자동 삽입 — ADR-0062). 따라서 명시적 HBM load가 뒤따르는 compute와 오버랩된다. kernbench의 실행 + latency 모델에 근거한 근거:

  • GEMM tiling이 PE_SCHEDULER에 offload된다. CompositeCmd는 non-blocking (kernel_runner.py:182-191, pe_scheduler.py:104-121): 커널이 하나의 coarse descriptor(M = G·T_q, PE당 전체 tile sweep)를 push하면 scheduler가 tile plan을 생성하고 타일당 DMA→GEMM→write를 스트리밍한다(ADR-0014 D6). K/V는 scheduler가 HBM에서 스트리밍하는 tl.ref operand이므로, 타일당 K/V prefetch는 scheduler의 일이다 — 명시적 prefetch op 없음. CPU(greenlet)는 현재 composite가 도는 동안 다음 composite를 낼 수 있게 풀려나, scheduler가 타일을 가로질러 GEMM 엔진을 saturate 유지한다.
  • 이는 하드웨어를 반영하며 CPU issue-rate를 execution-rate로부터 decouple한다. 대조적으로 blocking per-op tl.dot 경로는 매 GEMM마다 CPU를 멈추고, 끼어드는 softmax MATH op 동안 GEMM 엔진 버블을 남긴다; CPU가 fine-grained per-tile 발행을 따라잡을 수 있을 때만 현실적이다.
  • 실행 중 (m, , O) flash 상태는 Python TensorHandle 루프를 관통한다 (baseline이 이미 그렇게 함); softmax 머지(max/exp/sum/rescale)는 두 GEMM composite 사이의 커널 수준 tl MATH다. 기존 CompositeCmd는 두 GEMM을 chain하거나 cross-tile register 상태를 못 들므로(§0), 머지는 필연적으로 커널에 산다 — 이것이 하이브리드 분할이지, 우회한 한계가 아니다.
  • Cross-PE 결합은 P·V 후 (m, , O)에 대한 log-sum-exp tree로, 커널 수준 tl.send/tl.recv(§4)를 통해 — 불변.

그래서 per-tile 내부 파이프라인은:

q_g = tl.load(Q group)                         # lazy; auto-wait at first use
per tile j:
  Sⱼ = tl.composite("gemm", a=q_g, b=tl.ref(Kⱼ))  → Sⱼ   # scheduler streams Kⱼ DMA + GEMM
  Sⱼ += maskⱼ                                          # kernel MATH, boundary tile only
  online-softmax: mⱼ, m_new, P, corr,                # kernel MATH
  Oⱼ = tl.composite("gemm", a=P, b=tl.ref(Vⱼ))   → Oⱼ   # scheduler streams Vⱼ DMA + GEMM
  O = O*corr + Oⱼ;  m = m_new                         # kernel MATH (running merge)

각 타일의 MATH 임시값을 tl.scratch_scope(ADR-0063)로 감싸 scratch를 O(1)로 유지하고, 다음 타일의 composite를 현재 타일 결과를 wait하기 전에 발행 (non-blocking 핸들)하여 scheduler가 타일을 가로질러 파이프라인되게 한다.

제어 흐름(tile skip, mask 생성, reduction 스케줄링, 주소 산술)은 커널에 산다 (greenlet 본문의 평범한 Python if/산술). 이는 kernbench의 greenlet 모델이 이미 허용하는 바 그대로다(kernel_runner.py, ADR-0020 D3).

이것이 무엇을 대체하나. 이전 iteration은 "composite는 latency 이점 없음"이라는 논거로 순수 greenlet primitive 경로(전부 tl.dot, composite 없음)를 제안했다. 그것은 오직 시뮬레이터가 현재 per-op CPU 발행 비용을 0으로 청구하기 때문에만(dispatch_cycles=0, pe_cpu.py) 성립한다 — descriptor offload가 숨기려 존재하는 바로 그 CPU issue-rate / DMA-program 비용을 모델이 지워버린 것이다. 하이브리드가 효율 커널의 충실한 표현이다. 이점의 측정 가능한 크기(CPU가 많은 타일에 대해 엔진을 saturate할 수 있는가?)는 op 종류별 발행 비용 모델링에 좌우되며, future work로 추적된다(cost model; §9). dispatch_cycles=0에서도 non-blocking composite 경로는 blocking tl.dot 경로가 남기는 GEMM 엔진 버블을 채운다.

이것이 효율적 선택인 이유. GEMM tiling + DMA 스트리밍 + cross-tile 파이프라이닝은 검증된 CompositeCmd scheduler 경로에 offload되고; softmax 머지와 검증된 IPCQ collective는 커널에 남는다. 진짜 새 기계장치는 작고 범용적인 두 primitive뿐이다(ADR-0062 lazy tl.load, ADR-0063 tl.scratch_scope); reduction은 tl.send/tl.recv를 재사용한다. 전용 "flash-composite" 커맨드(softmax 머지 + carried register 상태 + IPCQ-push epilogue를 내재화하는 한 종류)는 만들지 않는다 — 크고 특수목적이며, 하이브리드 대비 유일한 delta(완전 softmax offload)가 현재 모델링 충실도에서 정당화되지 않는다; §8 참조.


2. 메모리 레이아웃 & 드라이버 책임

2.1 KV cache 할당

  • per-PE KV 버퍼는 per-PE 최대 context ⌈max_context / N⌉ × d × dtype로 K와 V를 각각 sizing. kernbench에서는 sequence 차원에 대해 pe(또는 cube) 축이 row_wiseDPPolicy로 배치된다(policy/placement/dp.py; baseline은 K/V에 DPPolicy(pe="row_wise"), Q에 replicate).
  • PE 안에서 할당된 token은 연속 append된다(slot 0,1,2,…). global index는 strided(i, i+N, …)이나 per-PE 버퍼는 dense ⇒ DMA가 연속 유지.
  • Slot → global: global_idx = local_slot·N + ((pe_id start_pe) mod N).

2.2 드라이버 launch당 의무 (최소)

드라이버는 launch마다 base + counter + rotation을 공급하고; 커널이 나머지를 도출한다:

Launch arg Purpose
K_base[pe], V_base[pe] per-PE KV 버퍼 base (tensor VA에서)
O_base 어텐션 출력 목적지
global_token_counter 현재 sequence 위치
start_pe = request_id mod N Layer-2 회전
pe_id (tl.program_id), N geometry
q_block_meta prefill/SP query block 시작 + 길이

커널 도출(평범한 산술):

  • 이번 step에 내가 쓸 차례: (start_pe + counter) mod N == pe_id
  • 내 valid 길이: base = counter // N; rem = counter % N; my_len = base + (1 if ((pe_id start_pe) mod N) < rem else 0)
  • read 범위: tile 0 .. ⌈my_len / TILE⌉.
  • causal bound / per-tile skip: query block과 각 tile의 global 위치에서.

드라이버 = base + counter + rotation. 단일 공식 (request_id + token_idx) mod N이 placement policy 전부다.


3. Per-tile op 시퀀스 (greenlet tl)

한 iteration = 한 PE의 한 KV tile. 두 GEMM은 tl.composite(op="gemm") (scheduler 관리 tiling + K/V DMA 스트리밍); softmax 머지는 그 사이의 커널 tl MATH다. 실제 tl 이름(tl_context.py), lazy tl.load(ADR-0062), tl.scratch_scope(ADR-0063) 사용:

# running state (persistent arena — allocated once, outside the scope)
# m: [G, T_q]  l: [G, T_q]  O: [G, T_q, d]
q_g = tl.load(Q_group_ptr, (G*T_q, d))               # lazy; auto-wait at first use (ADR-0062)

with tl.scratch_scope():                              # per-tile MATH temporaries recycled
    Sj    = tl.composite("gemm", a=q_g,                # [G·T_q, TILE]; scheduler streams Kⱼ DMA
                         b=tl.ref(K_base + j*TILE*d, (TILE, d))) * softmax_scale
    if mask_j is not None:
        Sj = Sj + mask_j                              # additive causal mask (boundary tile)
    m_j   = tl.max(Sj, axis=-1)
    m_new = tl.maximum(m, m_j)
    P     = tl.exp(Sj - m_new)                        # no full-matrix softmax; streaming
    corr  = tl.exp(m - m_new)                         # rescale factor for old accumulators
    l     = l * corr + tl.sum(P, axis=-1)
    Oj    = tl.composite("gemm", a=P,                  # [G·T_q, d]; scheduler streams Vⱼ DMA
                         b=tl.ref(V_base + j*TILE*d, (TILE, d)))
    O     = O * corr + Oj                             # running merge (kernel MATH)
    m     = m_new

Notes:

  • q_g[G·T_q, d]로 reshape된 GQA-batched query(G group 행을 matmul M 차원에 fold; byte-보존). 한 K/V tile이 모든 G·T_q 행을 서비스 — GQA 재사용 레버 — broadcast 없이.
  • K/V는 tl.ref operand로 composite scheduler가 HBM에서 타일당 스트리밍한다 (pe_scheduler.py:104-143): 그것이 바로 prefetch/파이프라인이므로 명시적 prefetch op이 없다. 타일 j+1의 composite를 타일 j를 wait하기 전에 발행 (non-blocking 핸들)하면 scheduler가 타일을 가로질러 파이프라인되고 GEMM 엔진이 saturate 유지된다.
  • tl.trans는 kernbench에서 메타데이터 전용이고(tl_context.py:390) MemoryStore.read는 transpose가 아니라 reshape한다(memory_store.py:73). zero/structural 실행엔 무해하나; 비자명 수치 데이터에선 reshape-not-transpose가 되어 transpose된 K로의 Q·Kᵀ는 주의가 필요하다(§11) — K를 transpose해 [d, TILE]로 미리 저장하거나, 실제 tl.transpose를 추가(후보 primitive, 시뮬레이터의 성능-모델링 목적상 아마 불필요).
  • Masking: 커널이 query/KV global offset에서 boundary-tile mask를 만들어 더한다 (Sj + mask_j); full-past tile은 None 전달; full-future tile은 skip(if이 enqueue 안 함).
  • "flash-composite" 커맨드 종류(softmax 머지 + carried (m,,O)를 내재화하는 것)는 쓰지 않는다; 기존 CompositeCmd가 각 GEMM을 담당하고 머지는 커널에 남는다(§1, §8 항목 4).

4. Reduction (KV-parallel / SP combine) — root로의 tree

tile sweep 후 각 PE는 (m_i, _i, O_i)(비정규화)를 가진다. 결합/교환 가능한 log-sum-exp 머지로 결합(baseline의 fold와 동일 수학, _attention_mesh_mlo.py:117-122):

def merge(m_a, l_a, O_a, m_b, l_b, O_b):
    m = tl.maximum(m_a, m_b)
    sa, sb = tl.exp(m_a - m), tl.exp(m_b - m)
    return m, l_a*sa + l_b*sb, O_a*sa + O_b*sb
# final: O = O_root / l_root      # normalise once at the tree root
  • 타이밍: 교환은 P·V 후(각 PE가 최종 O_i를 가짐).
  • 토폴로지: mesh tree, depth ⌈log₂ N⌉, 물리 이웃 상. N=8의 경우: level-0 (0↔1,2↔3,4↔5,6↔7), level-1 (1↔3,5↔7), level-2 (3↔7), root = 7. 각 tree pair는 물리 N/S/E/W 이웃이어야 tl.send(dir)/tl.recv(dir)이 실제 link를 쓴다(튜닝 항목 §9; 이웃을 wiring하는 SFR install은 configure_sfr_intracube_pe_ring, baseline이 사용).
  • 왜 baseline fan-out이 아니라 tree인가: 어텐션은 한 곳(O_base를 쓰는 PE)에서 O가 필요하다. tree는 ⌈log₂ N⌉ 단계(N=8에 3) vs baseline all-to-all의 N1 (7). per-PE sweep이 짧은 decode에서 reduction이 지배적 비용이므로 실제 이득이다. (baseline fan-out은 답을 모든 rank에 복제 — 여기선 불필요.)
  • Payload: O_i(d-vector)가 무겁고; m,(g, T_q)당 scalar. 별도 tl.send로 전송(baseline은 triplet을 세 send로 — 같은 패턴).

커널 구조(greenlet):

# leaf / internal node: fold children that send to me, then forward up
for child_dir in tree_children_dirs(pe_id, N):
    m_c = tl.recv(child_dir, m.shape); l_c = tl.recv(child_dir, l.shape)
    O_c = tl.recv(child_dir, O.shape)
    m, l, O = merge(m, l, O, m_c, l_c, O_c)
if not is_root(pe_id):
    tl.send(parent_dir(pe_id), m); tl.send(parent_dir(pe_id), l); tl.send(parent_dir(pe_id), O)
else:
    tl.store(O_base, O / l)

tl.recv는 데이터 도착까지 blocking(tl_context.py:446-499); 머지 순서가 정적 tree로 고정되므로 data-dependent pop이 불필요 — 하드웨어/composite pop-as-dependency 변경이 필요 없는 이유(§6).


5. Case별 커널

네 case 모두 §3(tile sweep) + §4(reduction)를 공유하며; N, query-block 폭, 어느 tile_* predicate가 발동하는지에서만 다르다.

5.1 공통 skeleton

def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
                q_block, softmax_scale, *, tl):
    pe_id   = tl.program_id(axis=rank_axis)
    my_len  = valid_len(counter, start_pe, pe_id, N)
    n_tiles = ceil(my_len / TILE)
    q_g     = load_Q_group(q_ptr)            # [G·T_q, d]; lazy tl.load, G folded into M (no broadcast)
    m, l, O = init_running()                 # persistent arena: -inf, 0, zeros
    for j in range(n_tiles):                 # K/V streamed per tile by the composite scheduler (§3)
        if tile_all_future(j, q_block):      # causal skip (kernel if)
            continue
        run_tile(j, mask_or_null(j, q_block))   # §3: 2 composites + softmax MATH, in tl.scratch_scope
    if N == 1:
        tl.store(o_ptr, O / l)               # no reduction
    else:
        tree_reduce_and_store(pe_id, N, o_ptr)  # §4

5.2 DECODE, no SP (N=1, 한 PE가 head의 KV 소유)

  • T_q = 1, 모든 과거 KV에 attend ⇒ future tile 없음, 마지막 ragged tile만 mask.
  • GQA 재사용이 전부(decode는 KV-load-bound): KV head의 G=8 query 행을 matmul M 차원에 fold(q_g[G, T_q, d] → [G·T_q, d]로 reshape, byte-보존). 그러면 Q·Kᵀcomposite([G·T_q, d], Kᵀ[d, TILE]) → [G·T_q, TILE]이고 P·Vcomposite([G·T_q, TILE], V[TILE, d]) → [G·T_q, d]. KV tile([TILE, d])이 공유 K/V operand — 한 번 스트리밍되어 모든 G·T_q 행이 자동 재사용 — 그것들이 GEMM의 M 행이기 때문. K/V broadcast 불필요; composite의 tile plan 내 m = G·T_q가 timing이 모든 G 행의 작업을 올바르게 세게 한다(leading batch 축은 세어지지 않음 — §8 참조).
  • S_pe가 scratch에 맞으면(작은/중간 context) 이것은 one-shot partial attention으로 축약(Q·Kᵀ용 composite 하나, softmax 하나, P·V용 composite 하나) — 바로 baseline _attention_mesh_mlo_partial_attention, 단지 GQA-batched에 composite 경로. Tiling(§3)은 S_pe가 scratch scope의 tile 예산을 초과할 때만 발동.

5.3 DECODE, with SP / KV-parallel (N=8)

  • 한 request의 KV가 8 PE에 round-robin; 각각 ≈my_len token 소유. 각 PE가 자기 로컬 tile에 대해 G=8 query 행을 GQA-batch.
  • T_q=1 ⇒ 짧은 per-PE sweep ⇒ reduction 지배 ⇒ §4 tree(3단계)가 구조적으로 중요한 부분. Reduction latency는 batch될 때 다른 동시 decode token으로, 혹은 긴 single-stream context의 긴 per-PE tile sweep으로 숨겨진다.

5.4 PREFILL, no SP

  • 전체 prompt 상주; query는 T_q token 블록(chunk).
  • causality가 실재, query block [qs, qe) vs KV tile [ks, ke):
    • ke ≤ qstile_all_pastmask=None, full compute.
    • ks ≥ qetile_all_futureskip(커널 if).
    • overlap → tile_partial → 커널이 삼각 additive mask 생성, tile op 시퀀스가 더함(Sj + mask_j).
  • GQA는 group의 query head를 decode와 같은 축으로 batch.

5.5 PREFILL, with SP (Ring KV)

  • KV가 N PE(ring)에 sharded; 각 step이 peer의 KV 블록을 IPCQ로 전달. 실행 중 (m,,O)ring step을 가로질러 carry(Python 핸들, _attention_mesh_kv가 오늘 하는 그대로).
  • IPCQ가 다음 step의 KV 수신을 현재 step의 compute와 tl.recv_async/tl.wait로 오버랩(tl_context.py:543-560 — 이미 존재).
  • Causal ring skip: 들어오는 step의 KV 블록이 전부 로컬 query block 이후면 그 compute를 skip(recv-consume 안 함 / fold 안 함) — causal 어텐션에서 약 절반 step 제거.
  • 수신 버퍼는 ping-pong(persistent arena, 재활용 안 함).
  • ring 후, (m,,O)는 query-소유 PE별 최종; KV-parallel 분할이 남으면 §4가 tree-merge한 뒤 정규화 + store.

baseline _attention_mesh_kv가 이미 ring fold를 구현한다; 본 ADR은 GQA 재사용, causal step-skip, recv_async 오버랩을 추가한다.


6. 왜 하드웨어 / composite 변경이 불필요한가

  • HW/composite pop-as-dependency가 도움 될 유일한 곳은 poll을 숨길 다른 작업이 없는 외로운 reduction — 즉 batch=1 그리고 짧은 context. 타깃은 agentic = 낮은 batch, 긴 context ⇒ 각 PE가 많은 KV tile을 가짐; §4 tree의 tl.recv blocking은 동시 token / 긴 context의 sweep 작업으로 가려진다.
  • §4 reduction은 정적 tree를 쓰므로 recv 순서가 컴파일 타임에 고정 — data-dependent pop 없음. tl.recv(blocking)로 충분.
  • 결정: greenlet tl.send/tl.recv collective로 출시. 짧은-context, single-stream, latency-critical 타깃이 나타날 때만 HW pop 재검토.

7. 제어 vs 실행 분할 (load-bearing 원칙)

Concern Owner
tile skip (future), mask 생성, causal bound 커널 (greenlet 본문의 Python if + 산술)
주소 / offset / valid-length 산술 커널 (counter에서)
reduction 스케줄링, IPCQ send/recv 순서 커널 (정적 tree)
Q·Kᵀ, P·V (per-tile K/V DMA 스트리밍 + tiling 포함) tl.composite → PE_SCHEDULER
Q load, mask add, softmax math, 실행 중 (m,,O) 머지 tl op PE 엔진 상 (커널 발행)

커널이 결정하고; GEMM은 composite로 scheduler에 offload되며; 나머지 tl op은 이미 결정된 작업을 실행한다. 이는 kernbench가 이미 지원하는 greenlet + composite 모델 그대로 — 새 제어 추상화 없음.


8. 필요한 kernbench 변경

설계 iteration의 정정: 진짜 GQA(h_q > h_kv)는 신규 primitive 불필요 — §5.2의 커널 재구조화(KV head별, G를 M에 fold, byte-보존 reshape)만 필요. 보조 ADR들은 efficiency / scale enabler이지 GQA blocker가 아니다.

커널 내 알고리즘 작업 (신규 primitive 없음; 기존 tl API):

  • GQA Q축 batching(재사용 레버) — KV head별로 G·T_q를 matmul M 차원에 fold (§5.2); _view 식 byte-보존 reshape; GEMM은 M = G·T_qtl.composite(op="gemm"). 오늘 timing/data 모드 모두 동작.
  • composite를 통한 GEMM(§1/§3) — Q·Kᵀ와 P·V를 각각 non-blocking tl.composite(op="gemm")로 발행; PE_SCHEDULER가 tiling하고 tl.ref K/V operand의 DMA를 스트리밍(기존 CompositeCmd; 새 커맨드 종류 없음).
  • root로의 tree reduction(§4)이 baseline all-to-all fan-out 대체 — tl.send/ tl.recv 상의 순수 커널 제어 흐름.
  • causal tile skip + additive boundary mask(§3/§5.4) — 커널 if + +로 더한 mask 텐서.
  • round-robin KV placement / valid-length 산술(§2) — launch-arg 산술 + DPPolicy(pe="row_wise").

efficiency / scale을 위한 신규 primitive (각각 보조 ADR 있음):

  1. per-tile scratch 재활용ADR-0063(tl.scratch_scope). 스케일에 필요: S=16 상한(1 MiB bump allocator)을 제거해 현실적 context 길이가 돌게 함. 셋 중 최고 가치.
  2. Lazy tl.loadADR-0062(non-blocking load + 첫 사용 시점 auto-wait; API 표면 불변). 효율: 명시적 load(Q group, 비-composite 커널)를 뒤따르는 compute와 오버랩. 타일당 K/V prefetch는 composite scheduler가 처리(§1)하므로, 이것은 나머지 명시적 load를 담당. 전역 시맨틱 변경 → 기존 골든 재생성(ADR-0062 D3).
  3. GQA head / mask broadcastADR-0061(tl.broadcast). 선택적 편의, GQA blocker 아님(위 정정 참조). G·T_q 행에 걸친 additive-mask 구성과 범용 커널에 유용; np.matmul이 data 모드에서 이미 broadcast하므로 정확성엔 불필요. 최저 우선순위.

명시적 REJECTED (효율적 대안 선택):

  1. 내부 루프 전체를 내재화하는 전용 "flash-composite" 커맨드 종류 — DMA→MM→VEC→DMA→MM→VEC + carried (m,,O) register 상태 + 꼬리 IPCQ push. 두 GEMM은 기존 CompositeCmd쓴다(§1/§3) — 그것이 scheduler 관리 tiling, K/V DMA 스트리밍, cross-tile 파이프라이닝을 준다. 기각되는 것은 softmax 머지 + cross-tile register 수명 + IPCQ-push epilogue까지 흡수하는 커맨드 종류다: 크고 특수목적이며, 하이브리드 대비 유일한 delta(완전 softmax offload)가 현재 모델링 충실도(per-op CPU 발행 비용 = 0; §1 참조)에서 정당화되지 않는다. cost model(§9)이 완전 offload를 측정 가능하게 가치 있게 만들면 재검토.
  2. 하드웨어 pop-as-dependency. 범위 밖(§6).
  3. 본 커널 내 RoPE / QKV projection / KV-cache write. Upstream qkv_rope (P1P5). RoPE를 fold하면 매 decode step마다 과거 tile을 재회전해야 하고 Ring Attention의 post-RoPE pass-through를 깬다.

9. Open 튜닝 항목 (kernbench에서 측정, blocking 아님)

  1. Composite tile-pipeline depth — KV-load-bound에 지배적; 커널이 wait 전에 non-blocking composite를 얼마나 앞서 발행하는지, 그리고 scheduler의 타일당 스트리밍 depth.
  2. reduction tree의 PE↔mesh-이웃 매핑 — 각 depth-⌈log₂ N⌉ pair가 물리 N/S/E/W 이웃인지 보장; 나쁜 매핑은 hop 추가. SFR install 대조 검증.
  3. TILE 크기 — scratch 상주(S/P tile + O_acc + G-way GQA)와 DMA 효율 균형; ADR-0063 및 scheduler의 TILE_M/K/N(pe_scheduler.py)과 상호작용.
  4. §5.5의 Ring 버퍼 ping-pong vs recv_async depth.
  5. decode의 one-shot vs tiled 교차점(§5.2) — tiling이 단일 composite를 이기는 S_pe 임계값.
  6. per-op CPU 발행 비용 (cost model) — 현재 dispatch_cycles=0(pe_cpu.py)이라 composite-vs-primitive 발행 오버헤드가 보이지 않음. op 종류별 차등 발행 비용 (tl.composite descriptor push ≫ primitive op)이 하이브리드의 CPU-saturation 이점을 측정 가능하게 한다(§1). ADR-0064에 명세; 별도 future work로 추적.

10. Coverage 요약

Case KV placement Inner loop Reduction Masking
Decode, no SP 1 PE, all KV one-shot or tiled, GQA-batched none last tile only
Decode, SP round-robin N PEs tiled, GQA-batched §4 tree (⌈log₂ N⌉) last tile only
Prefill, no SP resident tiled per q-block none skip future / triangular boundary
Prefill, SP (Ring) ring-sharded per ring step (recv_async overlap) running state + §4 tree causal step-skip + boundary

Case별 I/O (전체 계약 §0.5):

Case Inputs Output IPCQ partials
Decode, no SP Q[G,1,d], full K/V[S,d] on 1 PE O[G,1,d] at O_base none
Decode, SP Q[G,1,d], per-PE K/V[S_pe,d] O[G,1,d] (root) (m,,O_i) per PE → tree
Prefill, no SP Q[G,T_q,d], K/V[≤end,d] O[G,T_q,d] at O_base none
Prefill, SP (Ring) Q[G,T_q,d], ring-delivered K/V blocks O[G,T_q,d] (root) running (m,,O) + tree

모든 case에서: KV-cache write와 RoPE는 upstream에서; output projection은 downstream에서; score S와 prob P는 미생성.


11. 검증 계획 (Phase 1 테스트 개요)

SPEC/ADR coverage: R5 (PE↔PE IPCQ, PE↔HBM), R2 (traversal에 의한 latency), ADR-0023/0025 (IPCQ), ADR-0046 (tl 계약), ADR-0054 (eval bench).

시뮬레이터의 계약은 bit-exact 수치가 아니라 traversal에 의한 latency + 결정성 + 구조적 정확성(SPEC §0, §0.1)이다 — Phase 2 데이터는 주로 data path를 exercise하려 존재하고, tl.trans는 reshape-not-transpose이며, bf16f16으로 모델된다. 따라서 검증은 구조/타이밍 우선, 수치 parity는 제한된 2차 체크.

  1. data 모드 실행(enable_data=True): GQA 커널(h_q = G·h_kv)이 baseline head-packing이 부딪히는 byte-conservation 에러 없이 — 네 case 모두 — 완료. (이는 §5.2 재구조화가 필요하지, 신규 primitive가 아님.) 수치 parity(2차): reshape-as-transpose가 정확한 symmetric/identity 입력에 대해, 커널 O가 numpy FlashAttention reference와 fp tolerance 내 일치. 완전 asymmetric parity는 실제 tl.transpose에 좌우됨(범위 밖; 플래그됨).
  2. GQA 재사용: h_q = G·h_kv에서 K/V dma_read_countG와 무관(타일당 한 load, group에 걸쳐 재사용), 반면 GEMM 작업은 G로 스케일. 레버가 실제 발동함을 assert.
  3. tree reduction 단계 수: decode-SP가 ⌈log₂ N⌉ reduction 라운드(N1이 아님)를 발행; op_log ipcq_send/recv 수가 tree와 일치.
  4. Causal skip: prefill이 모든 tile_all_future tile을 skip — GEMM 수가 full grid가 아니라 하삼각 tile 수와 일치.
  5. Long context (ADR-0063): scope 없이 1 MiB를 넘기는 S sweep이 완료하고 reference와 일치.
  6. Load/compute 오버랩: tiled sweep의 end-to-end latency가 직렬 Σ(load+compute)보다 낮음 — composite scheduler의 타일당 K/V 스트리밍(§1/§3)과 Q load를 오버랩하는 lazy tl.load(ADR-0062)에서. (오버랩은 실제 모델 동시성, 빼기가 아님.)
  7. Composite GEMM offload (구조적): 각 tile의 Q·Kᵀ와 P·V가 blocking tl.dot이 아니라 CompositeCmd(non-blocking)를 PE_SCHEDULER에 emit; op_log가 composite tile plan을 보이고 커널이 다음 tile의 composite를 wait 전에 발행(cross-tile 파이프라이닝).
  8. 결정성: 동일 입력 → 동일 op_log + latency (SPEC §0.1).

B. 하이브리드 전환에서 나온 Open 설계 항목 (추후 검토)

이들은 결정이 순수 greenlet primitive 경로에서 composite hybrid + lazy tl.load(이번 개정)로 옮겨갈 때 생겼다. 설계를 막는 것은 없으며; 각각 구현 중 검증 패스가 필요하다. (묻지 않고) 작업 합의에 따라 여기 기록한다 — 권고는 내가 예측한 기본값이며; 검토 시 수정.

  1. DDD-0060이 아직 미동기화. Detailed Design Document는 여전히 옛 tl.load_async 더블버퍼 경로와 primitive tl.dot 내부 루프를 기술한다(그 §4.3/§5/§10). 하이브리드(composite GEMM, lazy load, K 사전 transpose)로 갱신해야 한다. DDD가 파생 how-to이고 큰 재작성이라 검토용으로 남김; ADR이 이제 권위 있는 기록이다. 권고: 구현 시작 전 후속으로 DDD 동기화.

  2. composite GEMM의 K operand 방향. Q·Kᵀ는 b = [d, TILE]가 필요하나 KV cache는 K를 [S_pe, d]로 저장한다. tl.trans는 메타데이터 전용이고 MemoryStore.read는 transpose가 아니라 reshape(memory_store.py:73) — 비자명 데이터엔 런타임 transpose가 틀리다. 권고: cache에 K를 사전 transpose [d, S_pe]로 저장(pseudocode와 §3이 이를 가정)하여 tl.ref(k_tile, (d, TILE))이 연속 slice가 되게. upstream qkv_rope write 레이아웃이 이를 지원하는지 검증, 아니면 실제 tl.transpose 추가(더 무거움; 연기).

  3. composite 출력 버퍼 vs tl.scratch_scope. 각 Q·Kᵀ composite는 Sjout_addr에 쓰고; 커널이 softmax MATH용으로 그것을 읽는다. 그 출력 버퍼와 in-flight composite의 타깃은 per-tile scratch_scope(ADR-0063)가 소비 전에 재활용하지 않는 곳에 살아야 한다 — in-flight lazy load와 같은 규율(ADR-0062 D-Negative). 권고: 현재 tile의 composite 출력은 scoped arena에(같은 iteration에 소비); persistent (m,,O)는 바깥에. 다음 tile의 composite를 일찍 발행할 때(cross-tile 파이프라이닝) use-after-recycle이 없음을 검증.

  4. composite 스트리밍 하의 GQA dma_read_count 레버. 레버(§11.2: K/V dma_read_countG와 무관)는 composite가 모든 G·T_q M-행에 재사용되는 하나의 K/V tile DMA를 emit한다고 가정한다. scheduler의 generate_gemm_planTILE_M/K/N(pe_scheduler.py:35-37, 32/64/32)으로 tiling한다 — G·T_q에 대한 M-tiling이 M-tile마다 공유 K/V tile DMA를 재발행하지 않음을 확인(즉 operand DMA가 M-tile에 걸쳐 공유되거나, 아니면 레버가 약해짐). 권고: levers 테스트에서 assert; 위반 시 GQA 이득은 DMA가 아니라 compute에만 — 여전히 정확하나 헤드라인이 바뀜.

  5. 커널 TILE vs scheduler TILE_M/K/N. 커널은 논리적 KV TILE을 다루고; scheduler는 고정 TILE_M/K/N으로 내부 재tiling한다. 두 tiling 레이어가 상호작용(scratch 상주, 파이프라인 depth). 권고: 커널 TILE을 K/V 스트리밍 granularity로 보고 scheduler가 GEMM을 sub-tile하게; DDD에 관계를 문서화하고 둘 다 sweep(§9 항목 1, 3).

  6. cost model은 별도 ADR. 하이브리드의 CPU-saturation 이점은 dispatch_cycles=0인 동안 보이지 않는다. op 종류별 issue-cost 모델은 ADR-0064에 명세; 본 ADR의 §1/§9는 측정 가능한(구조적뿐 아닌) 이점을 위해 그것에 의존. 권고: eval에서 하이브리드 latency 이점을 주장하기 전에 ADR-0064 모델을 도입.

  7. Ring 경로(§5.5) GEMM. §5.5는 여전히 primitive op + recv_async로 ring fold를 기술한다. 일관성을 위해 ring의 per-step Q·Kᵀ / P·V도 composite여야 하고; IPCQ recv_async 오버랩은 직교하며 유지. 권고: 구현 중 ring step에 같은 하이브리드 형태 적용; 낮은 리스크, §3 미러.