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.