changh95's picture
Add files using upload-large-folder tool
0190e6b verified
|
Raw History Blame Contribute Delete
10.5 kB

Action expert on one Blackhole chip — implementation notes

Code: alpamayo_tt/expert.py (device graph, trace, Euler loop), alpamayo_tt/action_proj.py (host action_in_proj / action_out_proj / rope tables), alpamayo_tt/expert_ref.py (torch oracle of one layer with HF semantics), tests/test_expert.py.

What runs where

Piece Where Notes
action_in_proj (Fourier 60 -> MLP 512 -> 1536 -> LayerNorm) host, torch fp32 64x1536 output uploaded as bf16 (196 KB) per step
64 decoder layers + final RMSNorm device, one trace inputs: x_dev (emb), cos_dev/sin_dev, mask_dev, the 64 VLM KV caches
action_out_proj (1536 -> 2) host, torch fp32 reads the final-normed hidden (196 KB)
Euler loop (10 steps) host x <- x + dt * v, t = linspace(0, 1, 11)[:-1]

The reference model runs its action_in_proj in bf16 with bf16-rounded Fourier frequencies (freqs is an nn.Module buffer, so from_pretrained(dtype=bf16) rounds 59.948 -> 60.0 etc.). Using exact frequencies gives a systematic 0.07 max-abs error vs the golden emb; with the rounded frequencies (default) the error is 0.009 and PCC 0.999995 (the remainder is the reference's bf16 MLP arithmetic).

Layouts (batch 1, 64 expert tokens)

Tensor Shape / dtype Notes
residual stream x [1, 1, 64, 1536] bf16 TILE DRAM 2 tile rows
wqkv [1536, 4096] bfp8 (or bf16) HF cat(q, k, v).T; q heads 0..15, then k heads 0..7, then v heads 0..7
wo [2048, 1536] w1/w3: [1536, 6144], w2: [6144, 1536], all HF .T
RMSNorm gammas [1, 1, 1, dim] bf16 ttnn.rms_norm requires gamma width == input width in this tt-metal
q / k / v heads [1, 16, 64, 128], [1, 8, 64, 128] nlp_create_qkv_heads(transpose_k_heads=False)
cos / sin [1, 1, 64, 128] bf16 HF rotate_half tables for positions i + rope_delta + S, theta 5e6
VLM KV cache (input) [1, 8, Smax, 128] bf16 or bfp8 TILE DRAM per layer HF head order, rows [0, S) valid
mask [1, 1, 128, Smax] bfp4 (decode mode) / [1, 1, 64 or 128, Smax] bf16 (standard) additive, 0 / -1e9

Per layer (19 ttnn ops): rms_norm -> linear(wqkv) -> nlp_create_qkv_heads -> rms_norm(q), rms_norm(k) (per-head, eps 1e-6) -> rotary_embedding_hf(q), (k) -> fill_cache(K), fill_cache(V) -> attention -> nlp_concat_heads -> linear(wo) -> add -> rms_norm -> linear(w1, silu), linear(w3) -> mul -> linear(w2) -> add. TTExpert.layer_fn is the swap point for a fused kernel (signature (layer_idx, x, K_cache, V_cache) -> x).

Attention over the VLM cache: masking approach

The expert's layer i attends non-causally to all S cached VLM positions of layer i plus its own 64 tokens. Instead of slicing / concatenating (every copy of the ~19 MB cache costs ~40-75 us), the own roped K and V are written into the cache itself, rows [Smax-64, Smax), with ttnn.fill_cache(cache, k, batch_idx=0, update_idx=Smax-64) (tile-aligned offset write, ~4 us). One attention op then reads the whole cache and an additive mask hides rows [S, Smax-64) (garbage / unused). Attention is permutation invariant over keys, so placing the own tokens at the tail is exact. Consequences:

  • the traced graph does not depend on S; per trajectory only cos/sin (positions) and the mask are rewritten from host (set_context, ~1 ms for the bfp4 mask of 128 x 8192);
  • the last 64 rows of every VLM KV cache are expert scratch: the text model must never store tokens there (S <= Smax - 64, asserted). With max_seq_len = 8192 and prompts of ~4.6K + <= 256 CoC tokens this is far away;
  • rows [S, Smax-64) may contain anything (tests fill them with 7.0).

Two attention kernels are implemented (attn_mode):

  • decode (default) — ttnn.transformer.scaled_dot_product_attention_decode, the flash-decode kernel that splits the K/V sequence across cores and reduces. The cache [1, 8, Smax, 128] is viewed (zero copy) as [8, 1, Smax, 128] = 8 users with one kv head each; the queries [1, 16, 64, 128] are viewed as [1, 8, 128, 128] = per user 128 "heads" (the 2 GQA q heads x 64 tokens; tile order is identical, so both views are free). Mask [1, 1, 128, Smax] bfp4, broadcast over users. k_chunk_size=128, 16 cores per user. 1024 "heads" in a single user would overflow L1 (tested), the 8-user view does not.
  • standard — ttnn.transformer.scaled_dot_product_attention(is_causal=False) with GQA q [1, 16, 64, 128] (or the 2 q heads of a group folded onto rows, fold_heads=True). Every (head, q-chunk) work item re-reads the entire K/V, so it is either per-core latency bound (few chunks) or DRAM bound (many chunks): ~520 us at Smax = 8192 regardless of chunking.

Measured attention time per layer (bf16 cache, S = 4700, trace-timed):

Kernel Smax = 8192 Smax = 5120
decode, mask bfp4, k_chunk 128 172 us 149 us
decode, bfp8 cache 143 us 126 us
standard, GQA q, q_chunk 32, mask bfp4 516 us 322 us
standard, folded heads, q_chunk 128 708 us 460 us

Matmul program configs (M = 64 rows, bfp8 weights, trace-timed)

1D multicast configs (MatmulMultiCoreReuseMultiCast1DProgramConfig, in0 broadcast, N split over the 11x10 grid, per_core_M=2) vs the ttnn default heuristics:

Matmul shape default tuned tuned config
wqkv 1536 x 4096 38 us 25 us in0_block_w 8, per_core_N 2
wo 2048 x 1536 43 us 19 us in0_block_w 8, per_core_N 1
w1 (silu fused), w3 1536 x 6144 45 us each 35-38 us each in0_block_w 8, per_core_N 2
w2 6144 x 1536 113 us 41 us in0_block_w 48, per_core_N 1

Total weight traffic per layer 40.5 MB bfp8 -> ~158 us = ~255 GB/s average (DRAM peak ~512 GB/s).

Trace safety (fixed a second-run corruption seen in the integrated pipeline)

A trace bakes in the addresses of the buffers its ops use; its intermediates are placed in memory that is free at capture time. A device buffer allocated after some trace was captured can therefore alias that trace's intermediates and is overwritten whenever the trace replays (verified directly on this tt-metal, see test_sample_repeat_with_foreign_trace). In the pipeline the text decode trace is captured before the expert first runs, so every expert input buffer (x_dev, cos_dev, sin_dev, mask_dev sized cfg.max_seq_len) is allocated in TTExpert.__init__, and sample() re-uploads cos/sin/mask at its start (~1 ms) instead of trusting the cached context across trajectories. The expert trace output is read back immediately after each replay. The expert trace itself is keyed on the KV cache and mask buffer addresses and re-captured if they change. Rule for integrators: allocate all persistent buffers (weights, caches, model inputs) before the first trace capture in the process.

Accuracy (tests/test_expert.py, golden = ref_out/030c760c_5100000_ref.pt, S = 4592)

Test Result
torch layer oracle vs transformers Qwen3VLTextModel (random weights) PCC 1.0000000, max abs 1.4e-6
action_in_proj vs golden emb, steps 0..9 max abs 9.0e-3, PCC 0.999995 (exact freqs: 6.9e-2 / 0.99953)
layer 0, device bfp8 / bf16 vs torch oracle (golden exp.layer0.in, golden cache rows + 13 synthetic) PCC 0.99993 / 0.99994
layer 31 PCC 0.999997 / 0.999997
layer 63 PCC 0.999998 / 0.999998
same, vs golden exp.layer{0,31,63}.out (13 of 4592 cache rows synthetic) PCC 0.99993 / 1.00000 / 0.99994
traced 64-layer step vs eager bitwise identical
sample() repeated with a foreign trace replayed in between bitwise identical

bf16 weights are not measurably better than bfp8 here (bf16 activations / bf16 cache dominate), so bfp8 is the default (2.55 GB for the 64 layers).

Timing (chip 2, S = 4700, Smax = 8192, bf16 cache, bfp8 weights, trace-timed)

traced 64-layer step (device, incl. 196 KB emb upload + hidden download) 32.5 ms (min 32.3)
full step() incl. host action_in_proj / action_out_proj 33.5 ms
10-step sample() ~335 ms
one layer, back-to-back trace 497 us

Per-op device time of one layer (us):

op us op us
rms_norm 17.3 concat_heads 27.7
wqkv (1536x4096 bfp8) 25.6 wo (2048x1536) 18.6
nlp_create_qkv_heads 52.8 add 6.5
q_norm / k_norm 6.5 / 5.8 rms_norm2 17.4
rope q / k 11.0 / 9.1 w1+silu / w3 (1536x6144) 40.2 / 37.3
fill_cache K / V 3.9 / 3.8 mul 10.9
attention (decode kernel) 175.3 w2 (6144x1536) 44.6
add2 6.7
sum 521

Where the time goes: attention 34%, weight matmuls 32% (40.5 MB/layer at ~245 GB/s), head split + concat 15%, the other 9 small ops 19%. The goal of < 30 ms/step as an op graph is missed by ~8%; what was tried and rejected: batched per-head matmuls to skip nlp_create_qkv_heads (382 us), batched o_proj + sum (62 us vs 46 us), the DRAM-sharded decode matmul (requires 32 rows), SDPA LoFi / exp-approx / other k-chunks (within 3%), 1024-"head" decode SDPA (L1 overflow). Remaining op-graph options: a bfp8 KV cache (attention 143 us, -2 ms/step), a 5120-row cache (-1.5 ms/step), fusing rms_norm+residual. The fused-layer kernel (megakernel stream) replaces the whole layer via TTExpert.layer_fn.

Known gaps

  • The golden dump only holds the VLM cache for the prefill (4579 rows); the 13 generated CoC positions are missing, so layer tests compare against the torch oracle with synthetic rows for the generated span and only report (not assert) the PCC vs the golden layer output (0.99993-1.0 anyway). End-to-end parity (per-step v, final action) is covered by the integrated pipeline (teacher-forced trajectories 0.02-0.31 m from the reference).
  • S must satisfy S <= Smax - 64 (tail rows are scratch). If the integrator uses a smaller max_seq_len (5120 fits the trajectory-profile prompts) the attention gets ~15% faster.
  • The trace bakes in the buffer addresses of the 64 KV caches and of the mask; it is re-captured automatically when they change (_trace_key). See "Trace safety" above.
  • nlp_create_qkv_heads (53 us) and nlp_concat_heads (28 us) are slow for 64-row inputs (few cores busy); no cheaper ttnn formulation was found — a fused kernel is the way out.
  • Only batch 1.