|
Download code/alpamayo_tt/EXPERT_NOTES.md from changh95/Alpamayo2-Super-p300x2: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/EXPERT_NOTES.md
- Command line
-
hf download hf://changh95/Alpamayo2-Super-p300x2/code/alpamayo_tt/EXPERT_NOTES.md
-
curl -L -o EXPERT_NOTES.md https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/EXPERT_NOTES.md
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. | |