changh95 commited on
Commit
00e0473
·
verified ·
1 Parent(s): afc42a7

Optimized build (2026-10-03): model call 87 -> 28 ms, pose path 38 ms

Browse files

code/ updated to the independently verified optimized build (12x10 compute grid). README Demo & Performances, GPU_COMPARISON.md update, OPT_REPORT/OPT_BASELINE, VERIFICATION_2026-10-03.md and patches/tt-metal-eth-dispatch.patch added. tt-model.yaml and SERVING.md are unchanged; they still describe the container image.

This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. GPU_COMPARISON.md +43 -0
  2. OPT_BASELINE.md +228 -0
  3. OPT_REPORT.md +0 -0
  4. README.md +42 -20
  5. VERIFICATION_2026-10-03.md +85 -0
  6. code/bench_breakdown.py +147 -0
  7. code/models/demos/mast3r/postprocess.py +10 -3
  8. code/models/demos/mast3r/tt/fattn.py +130 -0
  9. code/models/demos/mast3r/tt/fused.py +215 -0
  10. code/models/demos/mast3r/tt/heads_rope.py +1138 -0
  11. code/models/demos/mast3r/tt/kernels/add2_compute.cpp +34 -0
  12. code/models/demos/mast3r/tt/kernels/add3_compute.cpp +62 -0
  13. code/models/demos/mast3r/tt/kernels/add3_reader.cpp +27 -0
  14. code/models/demos/mast3r/tt/kernels/add3_writer.cpp +23 -0
  15. code/models/demos/mast3r/tt/kernels/add3s_reader.cpp +40 -0
  16. code/models/demos/mast3r/tt/kernels/add3s_writer.cpp +37 -0
  17. code/models/demos/mast3r/tt/kernels/fattn_reader.cpp +108 -0
  18. code/models/demos/mast3r/tt/kernels/fsdpa/compute_common.hpp +2484 -0
  19. code/models/demos/mast3r/tt/kernels/fsdpa/compute_streaming.hpp +0 -0
  20. code/models/demos/mast3r/tt/kernels/fsdpa/sdpa.cpp +277 -0
  21. code/models/demos/mast3r/tt/kernels/heads_concat_reader.cpp +38 -0
  22. code/models/demos/mast3r/tt/kernels/heads_concat_writer.cpp +43 -0
  23. code/models/demos/mast3r/tt/kernels/heads_rope_compute.cpp +180 -0
  24. code/models/demos/mast3r/tt/kernels/heads_rope_reader.cpp +110 -0
  25. code/models/demos/mast3r/tt/kernels/heads_rope_writer.cpp +57 -0
  26. code/models/demos/mast3r/tt/kernels/ln_add_compute.cpp +436 -0
  27. code/models/demos/mast3r/tt/kernels/ln_add_writer.cpp +43 -0
  28. code/models/demos/mast3r/tt/kernels/ln_add_writer_rb.cpp +64 -0
  29. code/models/demos/mast3r/tt/kernels/ln_fast_compute.cpp +202 -0
  30. code/models/demos/mast3r/tt/kernels/ln_reader_split.cpp +209 -0
  31. code/models/demos/mast3r/tt/kernels/mast3r_gelu_poly.h +116 -0
  32. code/models/demos/mast3r/tt/kernels/mm_gelu_activation.hpp +123 -0
  33. code/models/demos/mast3r/tt/kernels/mm_gelu_compute.cpp +657 -0
  34. code/models/demos/mast3r/tt/kernels/mm_in0_heads_reader.cpp +472 -0
  35. code/models/demos/mast3r/tt/kernels/phase_il_rm_reader.cpp +45 -0
  36. code/models/demos/mast3r/tt/kernels/phase_il_rm_writer.cpp +25 -0
  37. code/models/demos/mast3r/tt/kernels/phase_il_rm_writer_hs.cpp +29 -0
  38. code/models/demos/mast3r/tt/kernels/strip_gather.cpp +56 -0
  39. code/models/demos/mast3r/tt/kernels/tail_il_reader.cpp +30 -0
  40. code/models/demos/mast3r/tt/kernels/tail_il_writer.cpp +42 -0
  41. code/models/demos/mast3r/tt/kernels/tailf_compute.cpp +57 -0
  42. code/models/demos/mast3r/tt/kernels/tailf_reader.cpp +41 -0
  43. code/models/demos/mast3r/tt/kernels/tailf_writer.cpp +67 -0
  44. code/models/demos/mast3r/tt/kernels/ups2_compute.cpp +66 -0
  45. code/models/demos/mast3r/tt/kernels/ups2_reader.cpp +56 -0
  46. code/models/demos/mast3r/tt/kernels/ups2_writer.cpp +26 -0
  47. code/models/demos/mast3r/tt/kernels/ups2h_compute.cpp +53 -0
  48. code/models/demos/mast3r/tt/kernels/ups2h_reader.cpp +48 -0
  49. code/models/demos/mast3r/tt/kernels/ups2h_writer.cpp +20 -0
  50. code/models/demos/mast3r/tt/ttnn_dust3r.py +0 -0
GPU_COMPARISON.md CHANGED
@@ -184,3 +184,46 @@ HF_HUB_OFFLINE=1 /home/deepgadget/experiments/tt-models/.venv-gpu/main/bin/pytho
184
  ```
185
 
186
  GPU released after the run: `nvidia-smi --query-compute-apps=pid --format=csv,noheader` -> empty.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
184
  ```
185
 
186
  GPU released after the run: `nvidia-smi --query-compute-apps=pid --format=csv,noheader` -> empty.
187
+
188
+ ## Update 2026-10-03: optimized build and RTX 5090
189
+
190
+ The GPU numbers above are unchanged. They come from the 2026-09-14 run. This section compares them with the optimized build at commit `ad39150` (see `OPT_REPORT.md`).
191
+
192
+ ### New Blackhole numbers (independent verification, 2026-10-03)
193
+
194
+ The verifier measured these numbers with `bench_breakdown.py` (warm, batch 1, one 512x512 pair, 30 iterations per run, median / min). The hardware is a Blackhole chip with a 12x10 compute grid. Details: `VERIFICATION_2026-10-03.md`.
195
+
196
+ | Blackhole metric | Configuration | Value |
197
+ |---|---|---:|
198
+ | Model call `model(img1, img2)` (H2D + trace + D2H of both maps) | served path config: worker dispatch, 2 CQ, uint8 input | 28.23 / 27.55 and 27.99 / 27.68 ms (2 runs) |
199
+ | Device trace (host wall, synchronized) | served path config | 25.39 / 24.82 and 25.23 / 24.52 ms |
200
+ | Model call | ETH dispatch, 1 CQ, float input | 29.99 / 29.84 and 29.54 / 29.27 ms |
201
+ | Device trace | ETH dispatch, 1 CQ, float input | 24.72 / 24.47 and 25.03 / 24.41 ms |
202
+ | Device span (profiler cycles at nominal 1.35 GHz, 3 replays) | ETH dispatch, 1 CQ | 23.484 / 23.473 / 23.485 ms |
203
+ | `return_pose` forward: `sym` graph / two-pass | served path config, real pair | 38.00 / 55.82 ms |
204
+
205
+ For the ratios, this section uses 28.1 ms for the model call and 25.3 ms for the device trace. Each value is the mean of the two run medians (served path config). The model call matches the GPU `incl_h2d` definition. The device trace matches the GPU `excl_h2d` definition. The previous Blackhole value for the model call was 73.9 ms.
206
+
207
+ ### Recomputed ratios (Blackhole ms / GPU ms; > 1 means the GPU is faster)
208
+
209
+ | GPU variant / precision | GPU incl_h2d | ratio vs 28.1 ms | GPU excl_h2d | ratio vs 25.3 ms |
210
+ |---|---:|---:|---:|---:|
211
+ | ref, fp32 strict | 103.409 | 0.27 (Blackhole 3.68x faster) | 102.278 | 0.25 (Blackhole 4.04x faster) |
212
+ | ref, tf32 | 70.753 | 0.40 (Blackhole 2.52x faster) | 69.983 | 0.36 (Blackhole 2.77x faster) |
213
+ | ref, bf16 autocast | 63.753 | 0.44 (Blackhole 2.27x faster) | 62.613 | 0.40 (Blackhole 2.47x faster) |
214
+ | ref, fp16 autocast | 56.552 | 0.50 (Blackhole 2.01x faster) | 55.505 | 0.46 (Blackhole 2.19x faster) |
215
+ | bf16_weights, eager | 42.329 | 0.66 (Blackhole 1.51x faster) | 41.277 | 0.61 (Blackhole 1.63x faster) |
216
+ | fp16_weights, eager | 39.459 | 0.71 (Blackhole 1.40x faster) | 38.308 | 0.66 (Blackhole 1.51x faster) |
217
+ | SDPA, bf16_weights, eager | 35.327 | 0.80 (Blackhole 1.26x faster) | 34.285 | 0.74 (Blackhole 1.36x faster) |
218
+ | fp16 autocast + compile reduce-overhead | 23.214 | 1.21 (GPU 1.21x faster) | 22.149 | 1.14 (GPU 1.14x faster) |
219
+ | bf16_weights + compile reduce-overhead | 21.106 | 1.33 (GPU 1.33x faster) | 20.109 | 1.26 (GPU 1.26x faster) |
220
+
221
+ ### Reading
222
+
223
+ - The optimized build is faster than every eager GPU variant, including bf16 weights with SDPA attention (1.26x on the model call).
224
+ - `torch.compile` with CUDA graphs is still faster on the GPU: 1.21-1.33x on the model call and 1.14-1.26x on the device forward.
225
+ - On the measurement host, the Blackhole chip has an x1 PCIe link. The overlapped D2H readback stalls one device kernel and costs about 1.2 ms per pair (`OPT_REPORT.md`, round 9).
226
+ - Under load, the chip clock drops to about 1.26-1.30 GHz. Thus, the trace wall time is about 1.3-1.7 ms longer than the device span.
227
+ - With `return_pose`, the `sym` graph takes 38.0 ms. Two GPU forwards with bf16 weights and SDPA take 2 x 35.3 = 70.6 ms. This comparison is specific to the pose mode.
228
+ - The served `/predict` total stays host-bound on both sides. The npz encode takes about 160-172 ms, and the build did not change the encode code. The verifier did not measure a served `/predict` total for this build.
229
+ - The verifier did not measure p150a power, so this section makes no efficiency comparison.
OPT_BASELINE.md ADDED
@@ -0,0 +1,228 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # mast3r-p150 — optimization baseline (Galaxy BH chip 9)
2
+
3
+ Date: 2026-10-02. tt-metal `8b98410e730` + ETH-dispatch patch (shared tree, unmodified by me).
4
+ Model code is the baseline commit `20bdff0`, unmodified. Weights: `naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt` @ `61c57447d7b0` (shared HF cache).
5
+ Default path = `TT_FUSED=1` fused graph (llama RoPE, `dit_minimal_matmul_addcmul_fused` residual linears, SDPA 128/256 chunks, cached conv weights) captured into one metal trace.
6
+
7
+ ## Chip 9 vs p150
8
+
9
+ | | p150a (authors) | Galaxy BH chip 9 |
10
+ |---|---|---|
11
+ | Tensix grid, default dispatch | 11x10 | 12x10 (13 Tensix columns, 1 used for dispatch) |
12
+ | Tensix grid, ETH dispatch | 12x10 | 13x10 (policy: capped at 12x10, chipenv sets `TT_METAL_CORE_GRID_OVERRIDE_TODEPRECATE=11,9`) |
13
+ | PCIe | x16 | x1 Gen5. H2D and D2H are much slower, and readback is dominated by host untilize. |
14
+ | CQs with ETH dispatch | — | 1 only. This model uses 1 CQ, so ETH dispatch is usable. |
15
+
16
+ I reproduced the p150 11x10 grid with `TT_METAL_CORE_GRID_OVERRIDE_TODEPRECATE=10,9` and Tensix dispatch (`bench_breakdown.py --mode worker11`).
17
+
18
+ ## How to run
19
+
20
+ ```bash
21
+ ROOT=/home/ttuser/experiments/tt-models; M=$ROOT/models/mast3r-p150
22
+ # one time: $ROOT/tools/mkvenv.sh $M (done; numpy/cv2/hf_hub come from the shared env, nothing extra installed)
23
+ cd $M/code
24
+ source $ROOT/tools/chipenv.sh 9 $M/.venv # forces the 12x10 cap
25
+ export TT_WEIGHTS_REVISION=61c57447d7b0adc8a1a30b2b0adec7a8935aa2a3
26
+ python test_mast3r.py --layer end_to_end --runs 25 # authors' gate: synthetic pair, PCC >= 0.99, best-of-N latency (traced)
27
+ python test_mast3r.py --layer full_encoder|full_decoder|dpt_head_device|dpt_head_device_2 --runs 5 # eager per-stage gates
28
+ python tools_prof/real_pair_acc.py # real pair media/source_{1,2}.png vs fp32 torch (pts3d/conf PCC)
29
+ python bench_breakdown.py --mode worker|worker11|eth12 # per-phase timing: host prep / H2D / trace / D2H / e2e
30
+ MAST3R_TRACE=0 python -m tracy -r -p -v -o $TT_METAL_PROFILER_DIR/mast3r_eager --op-support-count 6000 tools_prof/prof_eager.py # device op profile
31
+ python tools_prof/d2h_probe.py # readback-format micro-benchmark
32
+ ```
33
+ The first run on a new grid JIT-compiles about 460 kernels, which takes about 80 s. After that the kernel cache in `chipstate/chip9/cache` is warm and the first call takes about 4.5 s.
34
+
35
+ ## Accuracy (baseline, all PASS)
36
+
37
+ | Test | Result |
38
+ |---|---|
39
+ | `end_to_end` synthetic pair, 11x10 Tensix | PCC 0.9987 (head1 0.9983 / head2 0.9985). README: 0.9987 |
40
+ | `end_to_end` synthetic pair, 12x10 Tensix | PCC 0.9984 (0.9983 / 0.9982) |
41
+ | `bench_breakdown` 12x10 ETH | head1 0.99836 / head2 0.99895 |
42
+ | `full_encoder` / `full_decoder` / `dpt_head_device` / `dpt_head_device_2` (eager) | 0.9997 / 1.0001 (b1 0.9999, b2 1.0000) / 1.0000 / 0.9999 |
43
+ | Real pair (kitchen source_1/2), 12x10 | pts3d PCC head1 0.98985 / head2 0.99078; conf PCC 0.9924 / 0.9944; raw PCC 0.9919 / 0.9765; median rel depth err 1.3%. README range: pts3d 0.987-0.9995, conf ~0.99 |
44
+
45
+ Proposed accuracy gates for optimization work, so we stay at or above the baseline: synthetic e2e PCC ≥ 0.998; real pair pts3d PCC ≥ 0.989 for each head; conf PCC ≥ 0.99.
46
+
47
+ ## Performance (warm, synthetic 512x512 pair, chip 9)
48
+
49
+ `bench_breakdown.py`, 30 iterations, median (min):
50
+
51
+ | Phase | 11x10 Tensix (p150-equiv) | 12x10 Tensix | 12x10 ETH |
52
+ |---|---|---|---|
53
+ | host prep (cat, bf16, host tensor) | 0.85 | 0.86 | 0.85 |
54
+ | H2D input `[2,3,512,512]` bf16 RM (3 MB) | 1.07 | 1.06 | 1.05 |
55
+ | **trace (device forward, execute + sync)** | **62.62** (62.49) | **61.82** (61.73) | **62.19** (62.12) |
56
+ | D2H, 2 heads (`to_torch` + reshape/permute/float) | 24.89 | 24.01 | 23.19 |
57
+ | **e2e model call** | 89.16 (84.39) | 87.20 (84.00) | 86.39 (83.94) |
58
+
59
+ `test_mast3r.py --layer end_to_end --runs 25` best-of-25: **84.91 ms** on 11x10 and **83.64 ms** on 12x10. The README reports 73.1 ms on p150; the 11 ms gap is the x1 PCIe readback.
60
+
61
+ - Moving from 11x10 to 12x10 gains only 0.8 ms (1.3%) on the device forward, because most ops do not scale to the full grid (see grid usage below).
62
+ - ETH dispatch vs Tensix dispatch at the same 12x10 grid makes no difference (62.2 vs 61.8 ms). The trace is kernel-bound: the sum of device kernel durations is 60.97 ms against 62.2 ms of trace, so dispatch gaps total only about 1.2 ms over 1252 ops (about 1 µs per op).
63
+
64
+ ## Host vs device; every host↔device transfer per forward
65
+
66
+ Everything from the patch-embed im2col to the final 1x1 head conv runs on device inside one trace. No host round-trips happen mid-graph on the fused path: RoPE LUTs, `trans_mat`, the `ones` vectors and all weights (including prepared conv weights) are uploaded during the eager warm-up, before capture.
67
+
68
+ On the host, per call:
69
+ 1. `torch.cat` + bf16 cast + `ttnn.from_torch` RM: about 0.85 ms.
70
+ 2. Post-processing of the outputs: reshape `(1,512,512,4)`, permute to NCHW, `.float()`. About 0.4 ms per head. This is included in D2H.
71
+
72
+ Transfers per forward (traced path):
73
+
74
+ | # | Direction | Tensor | Bytes on the wire | Time |
75
+ |---|---|---|---|---|
76
+ | 1 | H2D | input `[2,3,512,512]` bf16 ROW_MAJOR (`copy_host_to_device_tensor`) | 3.1 MB | 1.0 ms |
77
+ | 2 | D2H | head1 output `[1,1,262144,4]` bf16 **TILE, padded to 32 channels** | **16.8 MB** (8x the useful 2 MB) | ~12 ms |
78
+ | 3 | D2H | head2 output, same | 16.8 MB | ~12 ms |
79
+
80
+ The readback probe (`tools_prof/d2h_probe.py`, one head, so these are per-head times):
81
+ - `to_torch` of the padded TILE output takes 28-30 ms in isolation (12 ms in the model loop). Of that, raw `from_device` DMA is only 5.4 ms; the rest is host untilize of 16 MB.
82
+ - A 2 MB ROW_MAJOR tensor with 8-byte pages (`[262144,4]`) also costs 5.7 ms, because the cost is per page.
83
+ - A 2 MB ROW_MAJOR tensor with wide pages (`[1024,1024]`) costs **0.93 ms**, and a 2 MB TILE `[1024,1024]` takes 0.73 ms raw or 1.43 ms with `to_torch`.
84
+
85
+ **Conclusion:** writing the output in a compact, wide-page layout on device inside the trace would cut D2H from about 24 ms to about 2 ms. Examples: a channel-4 ROW_MAJOR tensor reshaped to `[512, 2048]`, or both heads packed into one `[2,512,2048]` RM tensor.
86
+
87
+ ## Grid / core usage
88
+
89
+ - **There is no hard-coded 11x10 or core grid in the port.** Every explicit program config takes `device.compute_with_storage_grid_size()` (`ttnn_dust3r.py:_fused_consts`, lines 108-135): `MinimalMatmulConfig` (8,4,4 / 4,4,4, 2x2 subblocks) and `SDPAProgramConfig(q=128, k=256)`. Everything else uses ttnn auto configs: `ttnn.linear` with no program config, interleaved `layer_norm`, NLP create/concat heads, conv2d auto sharding, and upsample. So the code follows whatever grid the device exposes. Its own docstrings note that the block sizes were tuned on the 11x10 p150.
90
+ - Actual core use per op, from the device profiler on 12x10 (120 cores), for one forward:
91
+
92
+ | Cores used | Ops | Kernel time |
93
+ |---|---|---|
94
+ | 120 (full) | 483 | 29.2 ms (48%) |
95
+ | 96-119 | 288 | 18.0 ms (30%). Auto `ttnn.linear` (MatmulMultiCoreReuseMultiCast on 96 cores, `in0_block_w=1`), conv2d 108/114, tilize 119 |
96
+ | 64-95 | 263 | 9.5 ms (16%). Encoder LayerNorm and NlpCreateHeads (64 cores = tile rows), conv2d at 80 cores |
97
+ | < 64 | 218 | 4.3 ms (7%). Decoder LayerNorm, create/concat heads (32 cores = 1024/32 tile rows) |
98
+
99
+ About half the device time runs on ops that do not use the full grid, which is why the 11→12 column change gave so little. Getting real value from 12x10 means sharded or explicit program configs sized to a 12-column grid (for example M=2048 → 64 tile rows, and N splits that divide 12).
100
+
101
+ ## What is fused / traced already; op count
102
+
103
+ - Whole graph traced: one `execute_trace` per pair. Persistent input buffer; outputs live in trace-owned buffers.
104
+ - **1252 device programs per forward**, by phase: patch-embed 5 (0.78 ms), encoder 289 (25.9 ms), decoder 534 (16.4 ms), DPT x2 424 (17.9 ms). Python-level ttnn calls are about 946; conv2d expands into halo, move, I2S/S2I and DRAM-slice ops.
105
+ - Already fused:
106
+ - RoPE: 1 `rotary_embedding_llama` per q/k, with the channel permutation folded into the weights.
107
+ - The 5 residual linears per block run as `dit_minimal_matmul_addcmul_fused` (bias + residual add fused).
108
+ - SDPA flash attention.
109
+ - Fused QKV / fused cross-attention KV plus `split_query_key_value_and_split_heads`.
110
+ - `concatenate_heads`.
111
+ - Relu fused into the conv1 / head2 convs.
112
+ - TILE reshape of the DPT taps.
113
+ - The encoder runs both views at B=2.
114
+ - **Not fused:** GELU. `ttnn.linear(..., activation="gelu")` with bias and an auto program config produces `fused_activation=nullopt` plus a **separate GELU UnaryDeviceOperation**: 48 ops, 4.97 ms (encoder 24 x 149 µs, decoder 24 x 57 µs). Also unfused:
115
+ - LayerNorm, separate from the following matmul.
116
+ - Create-heads separate from QKV.
117
+ - Pre-activation relu in ResConvUnit: 14 ops, 0.25 ms.
118
+ - Residual adds in the DPT (BinaryNg: 20 ops, 0.5 ms).
119
+ - Layout churn around the upsamples (tilize, untilize, S2I, I2S: about 4 ms in total).
120
+ - Weights: all bf16 (bf8 was rejected by the authors for accuracy), matmul HiFi2, conv HiFi4 + fp32 accumulation, LN HiFi4.
121
+
122
+ ## Biggest device-time sinks (kernel time per forward, 12x10)
123
+
124
+ | Op | Count | ms | % |
125
+ |---|---|---|---|
126
+ | `ttnn.linear` (MatmulDeviceOperation: qkv, fc1, cq, ckv, decoder_embed, patch, DPT 1x1) | 165 | 13.44 | 22.0 |
127
+ | MinimalMatmul (dit fused residual: proj, fc2, cproj) | 120 | 7.76 | 12.7 |
128
+ | Conv2d (DPT) | 60 | 6.96 | 11.4 |
129
+ | SDPA | 72 | 5.75 | 9.4 |
130
+ | Unary (48 GELU + 14 relu) | 62 | 5.21 | 8.5 |
131
+ | RotaryEmbeddingLlama | 144 | 4.07 | 6.7 |
132
+ | LayerNorm | 147 | 3.83 | 6.3 |
133
+ | NlpCreateHeads | 72 | 3.08 | 5.1 |
134
+ | Upsample (bilinear) | 10 | 1.87 | 3.1 |
135
+ | Tilize | 11 | 1.80 | 3.0 |
136
+ | Sharded→Interleaved / Interleaved→Sharded / Halo / Move / slices (conv plumbing) | ~270 | ~4.6 | 7.5 |
137
+
138
+ Top individual items:
139
+ - encoder fc1 matmul: 24 x 172 µs = 4.1 ms, 96 cores
140
+ - encoder fc2 dit: 24 x 153 µs = 3.7 ms
141
+ - encoder GELU: 3.6 ms
142
+ - encoder qkv: 24 x 148 µs = 3.5 ms
143
+ - encoder SDPA: 24 x 139 µs = 3.3 ms
144
+ - DPT refinenet-1 resconv 3x3 256→256 at 128x128: 8 x 355 µs = 2.8 ms, 80 cores
145
+ - decoder SDPA: 48 x 51 µs = 2.4 ms
146
+ - RoPE: 4.1 ms. The encoder RoPE costs 44 µs per call, about a third of an SDPA.
147
+ - head tilize `[262144,128]`: 2 x 628 µs
148
+ - head upsample 256→512: 2 x 588 µs
149
+ - final 128→4 1x1 matmul: 2 x 288 µs, a wasted N=32 pad
150
+
151
+ For scale: the encoder qkv matmul does 12.9 GFLOP in 148 µs, about 87 TFLOP/s. That is well below Blackhole bf16 peak, so matmul tuning has headroom.
152
+
153
+ ## Ranked optimization opportunities
154
+
155
+ Estimates are for the 12x10 grid. "e2e" is the model call including transfers.
156
+
157
+ 1. **Compact output readback** (−20 ms e2e, about 23%; device +≤0.5 ms). Inside the trace:
158
+ - Drop the 32-channel TILE padding (untilize to RM, or a wide-page reshape such as `[512, 2048]`).
159
+ - Optionally pack both heads into one buffer and read them with a single `to_torch`.
160
+ - Also stop running the 128→4 head conv as an N=32-padded TILE matmul.
161
+
162
+ This is the single biggest win on Galaxy's x1 PCIe, and it helps p150 as well. It does not change numerics.
163
+ 2. **Fuse GELU into fc1** (−4.5 to −5 ms device). Give an explicit `MatmulMultiCoreReuseMultiCastProgramConfig` with `fused_activation=(GELU, approx=False)`, or a `minimal_matmul` that fuses GELU, and check accuracy. Done together with item 3.
164
+ 3. **Explicit 12x10 program configs for the qkv / fc1 / ckv / cq / embed linears** (−4 to −6 ms). These currently run the auto MultiCast config on 96 cores with `in0_block_w=1`. Options: `minimal_matmul` / explicit block sizes for M=64 (encoder) and M=32 (decoder) tile rows on 12 columns, or 2D-sharded activations. A larger `in0_block_w` alone should help a lot. The authors' `MAST3R_FUSED_MM=minimal` was slower on p150 at (8,4,4); retune it for 12x10.
165
+ 4. **Sharded LayerNorm + create-heads / concat-heads on the full grid** (−3 to −4 ms). LN currently uses 64/32 cores at 42/18 µs, create-heads 59/35 µs, concat 25/13 µs. Width/block-sharded LN fused with the following matmul input (or `ttnn.rms/layer_norm` with a sharded program config) and `nlp_create_qkv_heads` from sharded input should be about 3x faster. Longer term: an L1-resident block pipeline (activations sharded in L1 across the whole block instead of DRAM-interleaved round trips), which is the "megakernel-lite" step.
166
+ 5. **DPT head restructuring** (−5 to −8 ms of the 17.9 ms DPT):
167
+ - Refinenet-1 3x3 convs at 128x128x256 run at 80 cores (2.8 ms). Force a 120-core height-sharded config and enable `enable_act_double_buffer` / `enable_weights_double_buffer`.
168
+ - Remove the tilize→upsample→S2I/I2S churn: about 4 ms of tilize, untilize, sharded↔interleaved, halo and slice ops.
169
+ - The final 512x512 head2 conv is DRAM-sliced into several conv passes. Doing the upsample+conv on an L1-sharded layout, or fusing bilinear ×2 and a 3x3 conv as transposed conv/sub-pixel, avoids the slicing.
170
+ - Run the two heads' independent work interleaved or batched where the weights allow; they cannot share weights, but layout plumbing can overlap.
171
+ 6. **RoPE** (−2 to −3 ms). 144 calls take 4.1 ms. Options:
172
+ - Fold RoPE into create-heads, using the `nlp_create_qkv_heads` + rotary fused variant (`rotary_embedding_llama_fused_qk` takes q and k in one call; check prefill support).
173
+ - At minimum, apply q and k in one op: 72 instead of 144 programs.
174
+ 7. **SDPA tuning** (−1 to −2 ms). Encoder 139 µs, decoder 51 µs. Re-sweep q/k chunk sizes for 12x10 (for example 256/256 or 128/512) and try `exp_approx_mode`, with an accuracy check.
175
+ 8. **bf8 weights / LoFi for the MLP only** (−3 to −5 ms, accuracy risk). The authors saw head1 pts3d PCC drop to 0.98 when all linears were bf8. Trying bfp8 only for fc1/fc2 of the encoder (or the decoder), with a real-pair gate, may be acceptable. This is lower priority because of the accuracy constraints.
176
+ 9. **ETH dispatch** (about 0 now). Dispatch overhead is only about 1.2 ms per trace, so moving dispatch to ETH matters only for the extra column. On Galaxy, Tensix dispatch already gives 12x10, and the 13x10 grid is forbidden by policy. ETH dispatch costs about 0.9 µs per op (dispatch_s off), so after the op count drops, compare both again. Keep ETH + 12x10 as the p150-representative configuration.
177
+ 10. **Host side** (−1 ms). Cast/cat into a pinned preallocated host tensor (0.85 ms). Upload the input as bf16 in a single `[2,3,512,512]` page-friendly layout; H2D is already 1 ms. Trace region: 512 MiB is reserved; measure the real size and shrink it to free DRAM (no speed gain).
178
+
179
+ Device-only potential: about 62 → 40-45 ms with items 2-7. E2E: about 84 → 45-50 ms on Galaxy, dominated by item 1.
180
+
181
+ ## Artifacts
182
+
183
+ - `code/bench_breakdown.py`: per-phase timing (host prep / H2D / trace / D2H / e2e) plus PCC, on any dispatch mode.
184
+ - `code/tools_prof/prof_eager.py`: device-profiler driver (eager fused graph, 2nd iteration = steady state).
185
+ - `code/tools_prof/real_pair_acc.py`: real-image accuracy gate.
186
+ - `code/tools_prof/d2h_probe.py`: readback-format micro-benchmark.
187
+ - Profiler CSV: `/home/ttuser/experiments/tt-models/chipstate/chip9/profiler/mast3r_eager/reports/2026_10_02_10_14_18/ops_perf_results_2026_10_02_10_14_18.csv` (2 iterations, 1252 ops each).
188
+
189
+ ## Post-crash re-verification (2026-10-02 11:18 UTC)
190
+
191
+ After the unexpected host shutdown, chip 9 opened cleanly with no reset needed and the kernel cache was 100% warm (456/456 hits). The baseline was re-run with the later optimization pass disabled (`MAST3R_OPT=none python test_mast3r.py --layer end_to_end --runs 25`, 12x10):
192
+ - PCC: 0.9984, PASS.
193
+ - Best-of-25 latency: 85.43 ms. Before the crash it was 83.64 ms; the difference is host-load noise.
194
+
195
+ The baseline numbers above still hold. Later commits (`0e84074` compact output, WIP `8f82803`) are optimization work beyond this baseline and are gated by `MAST3R_OPT`.
196
+
197
+ ## Post-crash #2 re-verification (2026-10-02 16:23 UTC)
198
+
199
+ After the second host crash, chip 9 opened cleanly with no reset, and the kernel cache was 456/456 warm. I re-ran the baseline (`MAST3R_OPT=none`, 12x10) and it reproduces:
200
+ - `test_mast3r.py --layer end_to_end --runs 25`: PCC 0.9984 (head1 0.9983 / head2 0.9982), PASS, best-of-25 85.53 ms.
201
+ - `bench_breakdown.py --mode eth12`, median:
202
+ - host_prep 0.81 ms, H2D 1.06 ms, trace 62.20 ms, D2H 23.52 ms, e2e 84.85 ms.
203
+ - PCC head1 0.99836 / head2 0.99895.
204
+ - `real_pair_acc.py`: pts3d PCC 0.98988 / 0.99076, conf 0.99237 / 0.99439. Identical to the first baseline.
205
+
206
+ I also checked the current default path (`MAST3R_OPT` unset, giving out, mm, dpt, l1, dptf). It includes the WIP `dptf` change from snapshot `13363fc`, which was unverified at the time of the crash. It now passes every proposed gate:
207
+ - e2e PCC 0.9988, PASS, best-of-25 53.66 ms.
208
+ - bench eth12, median: trace 49.30 ms, D2H 2.47 ms, e2e 53.78 ms.
209
+ - Real pair: pts3d 0.99032 / 0.99081, conf 0.99232 / 0.99466. Against the half-pixel-upsample reference, pts3d is 0.99988 / 0.99997.
210
+
211
+ So `dptf` is verified: the trace drops from 50.3 to 49.3 ms with no loss of accuracy. Compared with the baseline, the optimization knobs so far bring the device trace from 62.2 to 49.3 ms (-21%) and e2e from 84.9 to 53.8 ms (-37%), on the same chip, the same 12x10 grid and the same timing method.
212
+
213
+ ## Post-crash #3 re-verification (2026-10-03 00:46 UTC)
214
+
215
+ After the reboot, chip 9 (BDF 0000:42:00.0, /dev/tenstorrent/25) opened cleanly with no reset, and the JIT cache was 456/456 warm. Results:
216
+
217
+ - Baseline (`MAST3R_OPT=none`, 12x10):
218
+ - `test_mast3r.py --layer end_to_end --runs 25`: PCC 0.9984, PASS, best-of-25 **84.49 ms**.
219
+ - `bench_breakdown.py --mode eth12`, median (min):
220
+ - host_prep 0.82 ms, H2D 1.05 ms
221
+ - trace **62.21 ms** (62.14)
222
+ - D2H 23.57 ms (20.43)
223
+ - e2e 92.37 ms (84.76). The e2e median is noisier than before because of D2H spikes up to 37 ms; the device trace is unchanged.
224
+ - PCC head1 0.99836 / head2 0.99895.
225
+ - `real_pair_acc.py`: pts3d 0.98988 / 0.99076, conf 0.99237 / 0.99439. This is bit-identical to the earlier runs.
226
+ - Current default path (all OPT knobs): bench eth12 gives PCC 0.99852 / 0.99896, trace 49.28 ms, D2H 3.30 ms, e2e 54.44 ms (53.86).
227
+
228
+ The baseline numbers and the ranked opportunity list above still hold. The WIP snapshot `8db75c8` (`tools_prof/fid_probe.py`, `gelu_probe.py`) contains only probe scripts and does not affect the model path.
OPT_REPORT.md ADDED
The diff for this file is too large to render. See raw diff
 
README.md CHANGED
@@ -23,7 +23,7 @@ base_model:
23
 
24
  # mast3r-p150
25
 
26
- DUSt3R two-view 3D reconstruction (naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt: ViT-L/16 encoder, dual-branch decoder, two DPT heads; the backbone of MASt3R) running entirely on one Tenstorrent Blackhole p150a via tt-nn: an image pair in, two dense 512×512 pointmaps with confidence (plus optional PairViewer pose) out.
27
  Weights: [naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt](https://huggingface.co/naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt) · Paper: [arXiv:2312.14132](https://arxiv.org/abs/2312.14132) · Upstream code: [naver/mast3r](https://github.com/naver/mast3r) · Port: [changh95/tt-mast3r](https://github.com/changh95/tt-mast3r)
28
 
29
  Runs on **p150** (mesh `P150`).
@@ -37,7 +37,7 @@ tt-model pull changh95/mast3r-p150 --with-weights
37
  tt-model serve changh95/mast3r-p150
38
  ```
39
 
40
- - Weights [`naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt`](https://huggingface.co/naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt) at `61c57447d7b0` go to your HF cache; the image does not contain them.
41
  - Serves on port 20000 (or the next free port); ready when the log says `Application startup complete`.
42
 
43
  ### Run with tt-cli
@@ -68,6 +68,7 @@ tt model stop changh95/mast3r-p150
68
 
69
  - `npz_b64` is a base64 `.npz` holding `pts3d1`/`pts3d2` float32 (512,512,3) and `conf1`/`conf2` float16 (512,512); both pointmaps are in the camera-1 frame, depth is `pts3d[..., 2]` in DUSt3R's own scale (not metres), `conf = 1 + exp(c) >= 1`. `output_format: "png"` returns per-view 16-bit depth + 8-bit confidence PNGs instead.
70
  - `preprocess` maps original pixels to the 512 grid, `(u, v) = ((x + offx) * scale, (y + offy) * scale)` (`mode` = `pad`, the default gray pad-to-square, or `crop` = centre-crop to square via the `MAST3R_PREPROC` env; offsets are negative for `crop`). With `return_pose: true`, `pose` holds `R`, `t` (view 2 in the view-1 frame, `X_cam2 = R @ X_cam1 + t`), `focal_1`, `focal_2` (512-px units) and `used_known_intrinsics`; `null` when PnP fails.
 
71
 
72
  ### Demo
73
 
@@ -75,28 +76,49 @@ tt model stop changh95/mast3r-p150
75
 
76
  Kitchen frames 00 / 03 (VGGT example scene) → both predicted pointmaps in the camera-1 frame, coloured by the source pixels (served `npz` output, `code/make_demo.py`).
77
 
78
- ### Accuracy and speed
79
 
80
- | Metric | Value |
 
 
81
  |---|---:|
82
- | Torch fp32 reference vs upstream DUSt3R (`AsymmetricCroCo3DStereo`, same inputs), pts3d / conf PCC | 1.00000 / 1.00000 (max abs diff 7e-7) -- the reference is the upstream network since the 2026-09-14 decoder-tap fix |
83
- | End-to-end pointmap PCC vs fp32 torch reference (synthetic pair, `test_mast3r.py`) | 0.9987 (legacy graph `TT_FUSED=0`: 0.9989) |
84
- | Activated pts3d PCC vs reference, 4 real stereo pairs (CO3Dv2 apple, kitchen 00/03 and 00/08, LLFF fern 000/002), head 1 / head 2 range | 0.987-0.9995 / 0.991-0.9996 (2026-09-14, `logs/pointmap/mast3r/after_*_metrics.md`) |
85
- | Two-head coherence: median NN distance between the view-1 and view-2 pointmaps / scene scale, device (torch reference) | apple 0.012 (0.008) · kitchen 0.004 (0.003) · kitchen08 0.004 (0.0035); before the fix 0.39-0.65 (two separate sheets) |
86
- | PairViewer rotation vs the torch reference on the same pairs (est. focal) | 0.07-0.31° |
87
- | Inference, served over HTTP (warm, 30 requests, one 512×512 pair, npz; median / min / max; 2026-09-13, graph +2 `layer_norm` since) | 73.6 / 73.2 / 79.9 ms device forward · 241 / 240 / 280 ms end-to-end (160 ms of it is the host npz encode; ~4.1 pairs/s); 1.4 s with `return_pose` |
88
- | `test_mast3r.py --layer end_to_end`, best-of-25 (2026-09-14) | 73.1 ms per pair (legacy graph: 238.9 ms) |
89
- | Same forward on an RTX 5090 (same host, port's torch reference, eager PyTorch, batch 1, incl. H2D/D2H; bf16 / fp16 autocast) | 63.8 / 56.6 ms → GPU 1.2× / 1.3× faster than the p150a's 73.9 ms (both views); fp32-strict GPU 103.4 ms (p150a 1.4× faster); fp16 weights resident 39.5 ms (1.9×); best `torch.compile` 21.1 ms (3.5×) |
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
90
 
91
  ### Caveats
92
 
93
- - Fixed 512×512 input, exactly one image pair per request (batch 1); non-square images are gray-padded to square (`MAST3R_PREPROC=crop` centre-crops instead). The estimated focal ran ~5-8 % high on a CO3Dv2 pair with known K — pass `intrinsics` when you know them.
94
- - Point-map quality fix 2026-09-14: earlier revisions tapped the wrong decoder blocks for the DPT heads (two non-overlapping pointmaps, 0.98 PCC vs upstream); numbers published before that date (CO3Dv2 12-pair PCC / AUC@30, ETH3D) describe the old network and are not repeated here.
95
- - DUSt3R backbone only: no MASt3R matcher / descriptor head and no N-view global alignment; pose is single-pair PairViewer (estimated-focal PnP-RANSAC on one pair, no CO3Dv2 AUC re-measured on this host).
96
- - bf16 on device: activated pointmaps 0.987-0.9996 PCC vs fp32 per head on real pairs, confidence 0.99 PCC (raw confidence channel 0.85-0.96) — loosen thresholds tuned on the reference. Dense outputs come base64-encoded (`.npz`, or 16-bit/8-bit PNG).
 
 
 
97
  - Weights are NAVER's DUSt3R checkpoint under CC-BY-NC-SA-4.0: non-commercial use only, share-alike.
98
- - Not an OpenAI-compatible API; `GET /v1/models` is a stub so the tt-model ready card does not 404. Validated on tt-metal `v0.78.0-dev20260820` (main `8b98410e730`), single p150a only.
99
- - GPU comparison: GPU fp32-strict eager is 1.4× slower than the p150a; GPU bf16/fp16 autocast 1.2–1.3× faster, fp16 weights resident 1.9×, compiled 3.5×; the served `npz` request is host-bound on both sides (~160 ms compression). RTX 5090 rows (2026-09-14): same host, the port's own torch reference (same weights) run eagerly in PyTorch 2.11 cu128 with fp32 weights + autocast unless stated, no TensorRT; medians of 50 iterations after warm-up, H2D/D2H included; the p150a rows are the served bf16 fused path incl. upload/readback. p150a power was not measured, so no efficiency comparison is made. Full table: [`GPU_COMPARISON.md`](GPU_COMPARISON.md).
100
 
101
  ### Licensing
102
 
@@ -105,11 +127,11 @@ Kitchen frames 00 / 03 (VGGT example scene) → both predicted pointmaps in the
105
 
106
  ## Provenance
107
 
108
- The exact sources the image was built from — `code/` in this repo is byte-identical to the model code inside the image:
109
 
110
  | component | built from |
111
  | --- | --- |
112
  | tt-metal | [`8b98410e730bb504fea43a88609756e34821d91d`](https://github.com/tenstorrent/tt-metal/commit/8b98410e730bb504fea43a88609756e34821d91d) |
113
- | `code/` digest | `b45444934e1d76bf` (sha256, first 16 hex digits) |
114
  | built | 2026-09-14T06:35:13+00:00 by tt-model 0.1.0 |
115
 
 
23
 
24
  # mast3r-p150
25
 
26
+ NAVER DUSt3R ViT-L two-view 3D reconstruction (the MASt3R backbone) port on one Tenstorrent Blackhole p150a.
27
  Weights: [naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt](https://huggingface.co/naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt) · Paper: [arXiv:2312.14132](https://arxiv.org/abs/2312.14132) · Upstream code: [naver/mast3r](https://github.com/naver/mast3r) · Port: [changh95/tt-mast3r](https://github.com/changh95/tt-mast3r)
28
 
29
  Runs on **p150** (mesh `P150`).
 
37
  tt-model serve changh95/mast3r-p150
38
  ```
39
 
40
+ - Weights [`naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt`](https://huggingface.co/naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt) go to your HF cache; the image does not contain them.
41
  - Serves on port 20000 (or the next free port); ready when the log says `Application startup complete`.
42
 
43
  ### Run with tt-cli
 
68
 
69
  - `npz_b64` is a base64 `.npz` holding `pts3d1`/`pts3d2` float32 (512,512,3) and `conf1`/`conf2` float16 (512,512); both pointmaps are in the camera-1 frame, depth is `pts3d[..., 2]` in DUSt3R's own scale (not metres), `conf = 1 + exp(c) >= 1`. `output_format: "png"` returns per-view 16-bit depth + 8-bit confidence PNGs instead.
70
  - `preprocess` maps original pixels to the 512 grid, `(u, v) = ((x + offx) * scale, (y + offy) * scale)` (`mode` = `pad`, the default gray pad-to-square, or `crop` = centre-crop to square via the `MAST3R_PREPROC` env; offsets are negative for `crop`). With `return_pose: true`, `pose` holds `R`, `t` (view 2 in the view-1 frame, `X_cam2 = R @ X_cam1 + t`), `focal_1`, `focal_2` (512-px units) and `used_known_intrinsics`; `null` when PnP fails.
71
+ - The `timing_ms` values above come from the container image. With the optimized `code/`, `forward` is about 28 ms (see below). The host `encode` code did not change.
72
 
73
  ### Demo
74
 
 
76
 
77
  Kitchen frames 00 / 03 (VGGT example scene) → both predicted pointmaps in the camera-1 frame, coloured by the source pixels (served `npz` output, `code/make_demo.py`).
78
 
79
+ ### Demo & Performances
80
 
81
+ Warm, batch 1, one 512×512 pair, fused traced graph (`MAST3R_OPT=all`), `bench_breakdown.py` with 30 iterations per run (median / min). The served path config is worker dispatch, 2 CQ and uint8 input (the `tt-model.yaml` serve env).
82
+
83
+ | Metric | Performance |
84
  |---|---:|
85
+ | Model call `model(img1, img2)` (H2D + trace + D2H of both maps), served path config | **28.1 ms median** (27.99–28.23, 2 runs) · 27.55 ms min |
86
+ | Device trace (ViT-L encoder ×2, decoder, two DPT heads), served path config | **25.2–25.4 ms median** · 24.52 ms min |
87
+ | Model call, ETH dispatch, 1 CQ, float input | 29.5–30.0 ms median · 29.27 ms min |
88
+ | Device trace, ETH dispatch, 1 CQ | **24.7–25.0 ms median** · 24.41 ms min |
89
+ | Device span (profiler cycles at 1.35 GHz, ETH dispatch, 1 CQ, 3 replays) | 23.47–23.49 ms · 688 programs |
90
+ | Host prep · H2D, served path config (uint8 pixels) | 0.52 ms · 0.60 ms |
91
+ | D2H of both maps, ETH dispatch, 1 CQ | 2.95 ms |
92
+ | `return_pose` forward on the real pair, served path config: symmetric graph (`sym`) · two-pass | **38.0 ms median** · 55.8 ms median |
93
+
94
+ The measurement hardware is a Blackhole chip with a 12×10 compute grid and dispatch on ETH cores. The "served path config" rows use the same 12×10 compute grid with worker dispatch and 2 CQ. Accuracy against the torch fp32 reference on the real kitchen pair (served path config): pts3d PCC 0.99021 / 0.99086 (gate 0.989), conf PCC 0.99232 / 0.99497 (gate 0.99); synthetic pair `test_mast3r.py` end_to_end PCC 0.9985. Details: [`VERIFICATION_2026-10-03.md`](VERIFICATION_2026-10-03.md).
95
+
96
+ RTX 5090 reference measurements (2026-09-14) are unchanged. They use the port's torch reference in PyTorch 2.11, batch 1, on the kitchen demo pair. The "incl. H2D/D2H" column compares with our model call (28.1 ms). The "forward only" column compares with our device trace (25.3 ms). Full table: [`GPU_COMPARISON.md`](GPU_COMPARISON.md).
97
+
98
+ | RTX 5090 precision | GPU incl. H2D/D2H | vs current build (28.1 ms) | GPU forward only vs ours (25.3 ms) |
99
+ |---|---:|---:|---:|
100
+ | fp32 strict, eager | 103.4 ms | **Blackhole 3.68× faster** | 102.3 ms: Blackhole 4.04× faster |
101
+ | tf32, eager | 70.8 ms | **Blackhole 2.52× faster** | 70.0 ms: Blackhole 2.77× faster |
102
+ | bf16 autocast, eager | 63.8 ms | **Blackhole 2.27× faster** | 62.6 ms: Blackhole 2.47× faster |
103
+ | fp16 autocast, eager | 56.6 ms | **Blackhole 2.01× faster** | 55.5 ms: Blackhole 2.19× faster |
104
+ | bf16 weights, eager | 42.3 ms | **Blackhole 1.51× faster** | 41.3 ms: Blackhole 1.63× faster |
105
+ | bf16 weights + SDPA, eager | 35.3 ms | **Blackhole 1.26× faster** | 34.3 ms: Blackhole 1.36× faster |
106
+ | bf16 weights + `torch.compile` (CUDA graphs) | 21.1 ms | GPU 1.33× faster | 20.1 ms: GPU 1.26× faster |
107
+
108
+ The Blackhole lead comes from one metal trace that runs fused, block-sharded bf16 kernels on all 120 compute cores. The eager GPU rows pay for kernel launches and explicit softmax attention. The autocast rows also cast the fp32 weights again on each forward. `torch.compile` with CUDA graphs removes these costs, and then the GPU is about 1.3× faster. The served `/predict` total stays host-bound on both sides, because the npz encode takes about 160 ms.
109
 
110
  ### Caveats
111
 
112
+ - Does not scale to multiple p150a in a mesh configuration. The current build uses a 12x10 compute grid of Tensix cores. To get this grid, the dispatch functions move from 10 Tensix cores to ETH cores (`patches/tt-metal-eth-dispatch.patch`). Thus, this build assumes that you do not need chip-to-chip ethernet communication.
113
+ - The serve config sets `MAST3R_CQS=2`. Two CQs need worker (Tensix) dispatch. With ETH dispatch, set `MAST3R_CQS=1`. The outputs are bit-identical.
114
+ - Fixed 512×512 input, exactly one image pair per request (batch 1). Non-square images are gray-padded to square (`MAST3R_PREPROC=crop` centre-crops instead). The estimated focal can be 5-8 % high, so pass `intrinsics` when you know them.
115
+ - Numbers published before 2026-09-14 describe an older network with wrong decoder taps. This card does not repeat them.
116
+ - DUSt3R backbone only: no MASt3R matcher / descriptor head and no N-view global alignment. The pose is single-pair PairViewer (estimated-focal PnP-RANSAC). With `return_pose`, one symmetric graph (`sym`) computes the pose maps. Its outputs are bit-identical to the two-pass path.
117
+ - bf16 on device: on the kitchen pair, pts3d PCC vs fp32 is 0.990 / 0.991 per head and conf PCC is 0.992 / 0.995. Loosen thresholds that you tuned on the reference. Some optimizations change numerics (SDPA chunks, GELU polynomial, HiFi3 head convs, upsample kernel). [`OPT_REPORT.md`](OPT_REPORT.md) lists each change.
118
+ - Dense outputs come base64-encoded (`.npz`, or 16-bit/8-bit PNG).
119
  - Weights are NAVER's DUSt3R checkpoint under CC-BY-NC-SA-4.0: non-commercial use only, share-alike.
120
+ - Not an OpenAI-compatible API; `GET /v1/models` is a stub so the tt-model ready card does not 404.
121
+ - p150a power was not measured, so no efficiency comparison is made.
122
 
123
  ### Licensing
124
 
 
127
 
128
  ## Provenance
129
 
130
+ These are the exact sources the container image was built from. **`code/` has since been updated (2026-10-03 optimized build; see [`OPT_REPORT.md`](OPT_REPORT.md)) and is newer than the image.** `tt-model serve` runs the image's code until the image is rebuilt. `tt-model.yaml` and `SERVING.md` still describe the image:
131
 
132
  | component | built from |
133
  | --- | --- |
134
  | tt-metal | [`8b98410e730bb504fea43a88609756e34821d91d`](https://github.com/tenstorrent/tt-metal/commit/8b98410e730bb504fea43a88609756e34821d91d) |
135
+ | `code/` digest (image) | `b45444934e1d76bf` (sha256, first 16 hex digits; the current `code/` differs) |
136
  | built | 2026-09-14T06:35:13+00:00 by tt-model 0.1.0 |
137
 
VERIFICATION_2026-10-03.md ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # mast3r-p150: independent verification, 2026-10-03
2
+
3
+ ## Verdict
4
+
5
+ **PASS.** An independent verifier measured commit `ad39150` again. Every claimed speed number matches within 1-2 %. Every accuracy number matches the claim to the printed digit. All accuracy gates pass. The verifier found no cached outputs, no skipped compute and no loosened gates in the diff `5b918db..ad39150`. No chip faulted.
6
+
7
+ ## Setup
8
+
9
+ | item | value |
10
+ |---|---|
11
+ | Hardware | a Blackhole chip with a 12×10 compute grid (x1 PCIe link) |
12
+ | Dispatch configurations | ETH dispatch with 1 CQ ("eth12"), and worker (Tensix) dispatch with 2 CQ (the served path config). Both use the 12×10 compute grid. |
13
+ | tt-metal | `8b98410e730` with `patches/tt-metal-eth-dispatch.patch` |
14
+ | Verified commit | `ad39150` (the `code/` in this repo) |
15
+ | Baseline commit | `81ff339`. Its model code is identical to the previous release `20bdff0`. It is the earliest commit with `bench_breakdown.py`. |
16
+ | Workload | one 512×512 image pair, batch 1, warm, traced fused graph (`MAST3R_OPT=all`). The real-pair checks use the kitchen pair `media/source_1.png` + `media/source_2.png`. |
17
+ | Method | The verifier checked out baseline and final in separate git worktrees. The verifier ran the same `bench_breakdown.py` method on both. One workload ran at a time. |
18
+ | Scripts | `code/bench_breakdown.py`, `code/tools_prof/r3_validate.sh`, `code/tools_prof/real_pair_acc.py`, `code/tools_prof/sym_check.py`, `code/test_mast3r.py`, `code/tools_prof/trace_span.py` (device span under tracy) |
19
+
20
+ ## Measured baseline vs final
21
+
22
+ Median / min in ms. Each bench run has 30 iterations. Where the verifier ran a configuration twice, the table shows both runs.
23
+
24
+ | measurement | baseline `81ff339` | final `ad39150` | change |
25
+ |---|---:|---:|---:|
26
+ | device trace, ETH dispatch, 1 CQ, float | 62.21 / 62.09 and 62.23 / 62.17 | **24.72 / 24.47** and **25.03 / 24.41** | 2.5× faster |
27
+ | model call (H2D + trace + D2H), ETH dispatch, 1 CQ, float | 97.02 / 96.55 and 97.19 / 91.57 | **29.99 / 29.84** and **29.54 / 29.27** | about 3.3× faster |
28
+ | D2H of both maps, ETH dispatch, 1 CQ | 29.5 | **2.95** | |
29
+ | device trace, worker dispatch | 61.84 / 61.77 (1 CQ) | 24.95 / 24.48 (2 CQ, float) | |
30
+ | model call, worker dispatch | 97.77 / 97.49 (1 CQ) | 28.42 / 27.83 (2 CQ, float) | |
31
+ | device trace, served path config (worker, 2 CQ, uint8 input) | | **25.39 / 24.82** and **25.23 / 24.52** | |
32
+ | model call, served path config | | **28.23 / 27.55** and **27.99 / 27.68** | about 3.5× faster than the baseline (about 98 ms) |
33
+ | host prep / H2D, served path config | | 0.52 / 0.60 | |
34
+ | `test_mast3r.py --layer end_to_end` | | 27.71 | |
35
+ | real pair, served path config: single pair / `sym` pose / two-pass pose | | 27.90 / 38.00 / 55.82 (min 27.43 / 37.84 / 54.55) | |
36
+ | device span (tracy, ETH dispatch, 1 CQ, 3 replays) | | **23.484 / 23.473 / 23.485** (688 programs, kernel sum 22.85, FW gaps 0.244) | |
37
+
38
+ `MAST3R_OPT=none` at the final commit gives a device trace of 62.19 / 62.13 ms and a model call of 91.76 / 86.78 ms. This confirms the baseline.
39
+
40
+ The claimed device span is 23.491 / 23.479 / 23.490 ms. The verified span matches it.
41
+
42
+ ## Accuracy gate results
43
+
44
+ All gates pass on `ad39150`. The gate code did not change.
45
+
46
+ | metric | gate | final `ad39150` | baseline `81ff339` |
47
+ |---|---:|---:|---:|
48
+ | synthetic randn pair, eth12 float, PCC head 1 / head 2 / all | ≥ 0.998 | **0.99848 / 0.99855 / 0.99875** | 0.99836 / 0.99895 / 0.99909 |
49
+ | `test_mast3r.py --layer end_to_end` PCC | ≥ 0.998 | **0.9985** (PASS) | |
50
+ | real pair, float: pts3d PCC head 1 / head 2 | ≥ 0.989 | **0.99035 / 0.99081** | 0.98985 / 0.99078 |
51
+ | real pair, float: conf PCC head 1 / head 2 | ≥ 0.99 | **0.99262 / 0.99511** | 0.99237 / 0.99438 |
52
+ | real pair, served uint8: pts3d PCC head 1 / head 2 | ≥ 0.989 | **0.99021 / 0.99086** | |
53
+ | real pair, served uint8: conf PCC head 1 / head 2 | ≥ 0.99 | **0.99232 / 0.99497** | |
54
+ | real pair, raw PCC vs the half-pixel reference, float / uint8 | | 0.999882 / 0.999814 · 0.999886 / 0.999782 | |
55
+ | synthetic uint8 pair (not gated) | | 0.99733 / 0.99941 | |
56
+
57
+ Other checks:
58
+
59
+ - The e2e outputs are equal to the synchronized readback and to the first call.
60
+ - `sym_check`: `out_ii`, `out_ji` and `out_jj` of the `sym` graph are bit-identical to the two-pass path.
61
+ - On the real pair, the final build is equal to or better than the baseline on every gated metric.
62
+
63
+ ## Disclosed numerics changes
64
+
65
+ The optimization rounds changed some numerics. [`OPT_REPORT.md`](OPT_REPORT.md) lists each change with its measured effect. The main changes are these:
66
+
67
+ - `sdpa2`: larger SDPA chunks (encoder q160 / k512, decoder q224 / k512).
68
+ - `mc2d` and `mm32`: the plain and residual linears use 2-D multicast matmuls with fp32 dest accumulation.
69
+ - `gpoly`: the fc1 GELU uses a model-local erf-GELU polynomial (max abs error 6.2e-5 vs exact erf GELU).
70
+ - `hf3`: HiFi3 instead of HiFi4 for the two output-resolution head convs.
71
+ - `tups`: a model-local bilinear upsample kernel. It is closer to the exact bilinear result than the stock upsample.
72
+ - `hostcol`: the served path uploads uint8 pixels. The device applies the normalization.
73
+ - `resmm` and `resmmd`: the residual adds move into the matmul epilogues.
74
+
75
+ The synthetic head 2 PCC is 0.99855, below the baseline 0.99895. It stays above the 0.998 gate. The three steps of the last round (`esplit`, `demb`, `hrqk`) are bit-identical. `gelua` and `lofi` stay off because of accuracy.
76
+
77
+ ## Review findings
78
+
79
+ - The verifier reviewed the diff `5b918db..ad39150`. The gate and bench scripts did not change (`real_pair_acc`, `test_mast3r`, `bench_breakdown`, `sym_check`, `r3_validate`).
80
+ - `hrqk` runs the same per-tile op sequence for q and k. `esplit` changes only the output base address and start tile of each writer. `demb` uses the existing dual-matmul path.
81
+ - New diagnostic defines (`D_NOCOMP`, `D_NOREAD`, `D_NOWRITE`) and a reader override switch on only with `MAST3R_HROPE_DIAG` / `MAST3R_LN_READER*`. These variables are empty by default. `tt-model.yaml` and the server do not set them.
82
+ - Minor: the verifier did not measure the round-start device span (23.531 ms) again. The last round gained only 0.044 ms (0.19 %). Host wall time cannot show this gain. The final absolute span reproduces.
83
+ - Minor: `trace_span.py` converts cycles to time at the nominal 1.35 GHz. Under load, the clock drops to about 1.26-1.30 GHz. Thus, the true device time is about 4 % above the printed span.
84
+ - Context: the RTX 5090 with `torch.compile` and CUDA graphs (21.1 ms with H2D/D2H, 20.1 ms forward only) is still faster than this build (about 28 ms model call, about 25 ms trace). This build is faster than every eager GPU variant, including bf16 weights with SDPA (35.3 ms). See [`GPU_COMPARISON.md`](GPU_COMPARISON.md).
85
+ - The verifier did not run the server smoke test. The developer reports PASS 4 / 4 with the `tt-model.yaml` serve env (worker dispatch, 2 CQ).
code/bench_breakdown.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Per-phase timing of the traced fused DUSt3R forward (baseline profiling helper).
3
+
4
+ Phases per call (same work as TtDust3r.__call__ / test_mast3r.py end_to_end):
5
+ host_prep : torch.cat + bf16 cast + host ttnn tensor (ROW_MAJOR)
6
+ h2d : copy_host_to_device_tensor into the persistent trace input (blocking)
7
+ trace : execute_trace + synchronize_device (pure device time of the whole graph)
8
+ d2h : two to_torch readbacks of the [1,1,H*W,4] head outputs (+ host untilize/reshape)
9
+ e2e : the public model call (host_prep + h2d + trace + d2h, non-blocking trace)
10
+
11
+ Usage (after chipenv.sh):
12
+ python bench_breakdown.py --mode worker11 # p150-equivalent: Tensix dispatch, grid capped 11x10
13
+ python bench_breakdown.py --mode worker # Galaxy default Tensix dispatch (12x10, Galaxy-only A/B)
14
+ python bench_breakdown.py --mode eth12 # ETH dispatch, capped 12x10
15
+ MAST3R_CQS=2 python bench_breakdown.py --mode worker # Tensix dispatch 12x10, 2 CQs (split trace, overlapped head-1 read)
16
+ add --u8 for uint8 pixel views (HWC-backed, as the server's PIL preprocess produces them)
17
+
18
+ With MAST3R_CQS=2 "trace" is both traces back to back (A + B, synchronized), and "d2h" is the two reads
19
+ after that sync (no overlap); only "e2e" shows the overlap.
20
+ """
21
+ from __future__ import annotations
22
+
23
+ import argparse
24
+ import os
25
+ import statistics
26
+ import sys
27
+ import time
28
+
29
+ _CODE_ROOT = os.path.dirname(os.path.abspath(__file__))
30
+ sys.path.insert(0, _CODE_ROOT)
31
+ sys.path.insert(0, "/home/ttuser/experiments/tt-models/tools")
32
+
33
+ ap = argparse.ArgumentParser()
34
+ ap.add_argument("--mode", default="worker11", choices=["worker11", "worker", "eth12"])
35
+ ap.add_argument("--iters", type=int, default=30)
36
+ ap.add_argument("--no-ref", action="store_true")
37
+ ap.add_argument("--u8", action="store_true", help="uint8 pixel views (exact reference input (v/255-0.5)/0.5)")
38
+ ap.add_argument("--dump", default="", help="save the first-call outputs (torch.save) for bit-exact A/B")
39
+ ap.add_argument("--cmp", default="", help="compare the first-call outputs with a --dump file")
40
+ ap.add_argument("--u8-float", action="store_true", help="the --u8 images, fed to the model as float (A/B)")
41
+ args = ap.parse_args()
42
+
43
+ # grid cap must be set before ttnn creates its runtime options
44
+ if args.mode == "worker11":
45
+ os.environ["TT_METAL_CORE_GRID_OVERRIDE_TODEPRECATE"] = "10,9"
46
+
47
+ import torch # noqa: E402
48
+ import ttnn # noqa: E402
49
+ from models.demos.mast3r.reference.torch_dust3r import load_checkpoint, load_dust3r # noqa: E402
50
+ from models.demos.mast3r.tt.ttnn_dust3r import fused_config, get_model, release_device_caches # noqa: E402
51
+
52
+
53
+ def pcc(a, b):
54
+ a = a.float().flatten(); b = b.float().flatten()
55
+ a = a - a.mean(); b = b - b.mean()
56
+ return float((a @ b) / (a.norm() * b.norm()))
57
+
58
+
59
+ cfg = fused_config()
60
+ kw = cfg.open_device_kwargs(l1_small_size=32 * 1024)
61
+ print(f"# MAST3R_CQS={cfg.num_cqs} opts={','.join(cfg.opts)}")
62
+ if args.mode == "eth12":
63
+ from open_device_12x10 import open_device
64
+ device = open_device(grid="12x10", **kw)
65
+ else:
66
+ device = ttnn.open_device(device_id=0, **kw)
67
+ g = device.compute_with_storage_grid_size()
68
+ print(f"# mode={args.mode} grid={g.x}x{g.y}")
69
+
70
+ torch.manual_seed(0)
71
+ state = load_checkpoint()
72
+ img1 = torch.randn(1, 3, 512, 512)
73
+ img2 = torch.randn(1, 3, 512, 512)
74
+ rin1, rin2 = img1, img2
75
+ if args.u8 or args.u8_float:
76
+ # smooth-ish random uint8 images (HWC arrays viewed as (1, 3, H, W), like the server's preprocess)
77
+ def u8img():
78
+ lo = torch.rand(1, 3, 32, 32) * 255
79
+ im = torch.nn.functional.interpolate(lo, size=(512, 512), mode="bicubic", align_corners=False)
80
+ im = (im + torch.randn_like(im) * 20).clamp(0, 255).round().to(torch.uint8)
81
+ return im[0].permute(1, 2, 0).contiguous().permute(2, 0, 1).unsqueeze(0) # HWC storage
82
+ img1, img2 = u8img(), u8img()
83
+ rin1, rin2 = ((img1.float() / 255 - 0.5) / 0.5), ((img2.float() / 255 - 0.5) / 0.5)
84
+ if args.u8_float:
85
+ img1, img2 = rin1, rin2
86
+ ref = None
87
+ if not args.no_ref:
88
+ t0 = time.perf_counter()
89
+ with torch.no_grad():
90
+ ref = load_dust3r(state)(rin1, rin2)
91
+ print(f"# torch fp32 reference: {(time.perf_counter() - t0):.1f} s")
92
+
93
+ try:
94
+ model = get_model(state, device)
95
+ t0 = time.perf_counter()
96
+ out = model(img1, img2) # eager warm + capture + first replay
97
+ print(f"# first call (compile + warm + capture): {(time.perf_counter() - t0):.1f} s")
98
+ if ref is not None:
99
+ p1, p2 = pcc(ref[0], out[0]), pcc(ref[1], out[1])
100
+ pall = pcc(torch.stack(ref), torch.stack(out))
101
+ print(f"# PCC head1={p1:.5f} head2={p2:.5f} all={pall:.5f}")
102
+
103
+ if args.dump:
104
+ torch.save(out, args.dump)
105
+ if args.cmp:
106
+ prev = torch.load(args.cmp)
107
+ for h in range(2):
108
+ d = (prev[h] - out[h]).abs()
109
+ print(f"# vs {args.cmp}: head{h+1} bit-identical={torch.equal(prev[h], out[h])} n_diff={int((d > 0).sum())} "
110
+ f"max_abs={float(d.max()):.3g} PCC(prev)={pcc(prev[h], out[h]):.6f}")
111
+ B, _, H, W = img1.shape
112
+ ph = {k: [] for k in ("host_prep", "h2d", "trace", "d2h", "e2e")}
113
+ kind = model.input_kind(img1)
114
+ kd = model._kinds[kind]
115
+ for _ in range(args.iters):
116
+ t0 = time.perf_counter()
117
+ host_in = model._host_input(img1, img2)
118
+ t1 = time.perf_counter()
119
+ ttnn.copy_host_to_device_tensor(host_in, kd["in"], cq_id=0)
120
+ ttnn.synchronize_device(device)
121
+ t2 = time.perf_counter()
122
+ for tid in kd["tids"]:
123
+ ttnn.execute_trace(device, tid, cq_id=0, blocking=False)
124
+ ttnn.synchronize_device(device)
125
+ t3 = time.perf_counter()
126
+ _ = model.read_outputs(kind, B, H, W) # synchronized (no events pending)
127
+ t4 = time.perf_counter()
128
+ ph["host_prep"].append((t1 - t0) * 1e3); ph["h2d"].append((t2 - t1) * 1e3)
129
+ ph["trace"].append((t3 - t2) * 1e3); ph["d2h"].append((t4 - t3) * 1e3)
130
+ for _ in range(args.iters):
131
+ t0 = time.perf_counter()
132
+ _ = model(img1, img2)
133
+ ph["e2e"].append((time.perf_counter() - t0) * 1e3)
134
+ for k, v in ph.items():
135
+ print(f"{k:10s} median {statistics.median(v):8.2f} ms min {min(v):8.2f} max {max(v):8.2f}")
136
+ o = kd["outs"][0]
137
+ print(f"# output tensor: shape={tuple(o.shape)} padded={tuple(o.padded_shape)} layout={o.layout} dtype={o.dtype}")
138
+ print(f"# input tensor: shape={tuple(kd['in'].shape)} layout={kd['in'].layout} dtype={kd['in'].dtype} traces={len(kd['tids'])}")
139
+ # bit-exactness of the e2e path (overlapped reads with MAST3R_CQS=2) vs the synchronized reads above
140
+ a = model(img1, img2)
141
+ ttnn.synchronize_device(device)
142
+ b = model.read_outputs(kind, B, H, W)
143
+ print(f"# e2e outputs == synchronized readback: {all(torch.equal(x, y) for x, y in zip(a, b))}; "
144
+ f"== first call: {all(torch.equal(x, y) for x, y in zip(a, out))}")
145
+ finally:
146
+ release_device_caches()
147
+ ttnn.close_device(device)
code/models/demos/mast3r/postprocess.py CHANGED
@@ -58,7 +58,7 @@ def preprocess_mode(env: Optional[Mapping[str, str]] = None) -> str:
58
  return v
59
 
60
 
61
- def preprocess_image(img: Image.Image, size: int = IMG_SIZE, mode: str = DEFAULT_PREPROC):
62
  """Square ``size x size`` network input in [-1, 1] plus a preprocess record.
63
 
64
  ``mode="pad"`` (default): pad to a square (gray 128) then bicubic-resize -- preserves
@@ -69,6 +69,10 @@ def preprocess_image(img: Image.Image, size: int = IMG_SIZE, mode: str = DEFAULT
69
  with ``scale = out_size / pad_side``; ``pad_side`` is the side of the square canvas
70
  (``max(W, H)`` for pad, ``min(W, H)`` for crop) and the offsets are >= 0 for pad,
71
  <= 0 for crop. :func:`intrinsics_to_canonical` applies the same map to ``K``.
 
 
 
 
72
  """
73
  if mode not in PREPROC_MODES:
74
  raise ValueError(f"preprocess mode must be one of {PREPROC_MODES}, got {mode!r}")
@@ -85,6 +89,9 @@ def preprocess_image(img: Image.Image, size: int = IMG_SIZE, mode: str = DEFAULT
85
  canvas = img.crop((x0, y0, x0 + s, y0 + s))
86
  offx, offy = -x0, -y0
87
  canvas = canvas.resize((size, size), Image.BICUBIC)
 
 
 
88
  arr = np.asarray(canvas, dtype=np.float32) / 255.0
89
  arr = (arr - 0.5) / 0.5
90
  tensor = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).contiguous()
@@ -92,10 +99,10 @@ def preprocess_image(img: Image.Image, size: int = IMG_SIZE, mode: str = DEFAULT
92
  "pad_side": s, "out_size": size, "mode": mode}
93
 
94
 
95
- def load_image_for_dust3r(path, size: int = IMG_SIZE, mode: str = DEFAULT_PREPROC):
96
  """File-path convenience wrapper around :func:`preprocess_image` (eval scripts)."""
97
  with Image.open(Path(path)) as im:
98
- return preprocess_image(im, size, mode)
99
 
100
 
101
  def intrinsics_to_canonical(K: np.ndarray, preproc: dict) -> np.ndarray:
 
58
  return v
59
 
60
 
61
+ def preprocess_image(img: Image.Image, size: int = IMG_SIZE, mode: str = DEFAULT_PREPROC, uint8: bool = False):
62
  """Square ``size x size`` network input in [-1, 1] plus a preprocess record.
63
 
64
  ``mode="pad"`` (default): pad to a square (gray 128) then bicubic-resize -- preserves
 
69
  with ``scale = out_size / pad_side``; ``pad_side`` is the side of the square canvas
70
  (``max(W, H)`` for pad, ``min(W, H)`` for crop) and the offsets are >= 0 for pad,
71
  <= 0 for crop. :func:`intrinsics_to_canonical` applies the same map to ``K``.
72
+
73
+ ``uint8=True``: return the resized pixels themselves as a ``(1, 3, size, size)`` torch.uint8
74
+ view of the HWC array (no copy); the network input is exactly ``(v/255 - 0.5)/0.5`` of it,
75
+ which the fused device graph applies itself (``TtDust3r`` uint8 path).
76
  """
77
  if mode not in PREPROC_MODES:
78
  raise ValueError(f"preprocess mode must be one of {PREPROC_MODES}, got {mode!r}")
 
89
  canvas = img.crop((x0, y0, x0 + s, y0 + s))
90
  offx, offy = -x0, -y0
91
  canvas = canvas.resize((size, size), Image.BICUBIC)
92
+ rec = {"orig_W": W, "orig_H": H, "offx": offx, "offy": offy, "pad_side": s, "out_size": size, "mode": mode}
93
+ if uint8:
94
+ return torch.from_numpy(np.array(canvas, dtype=np.uint8)).permute(2, 0, 1).unsqueeze(0), rec
95
  arr = np.asarray(canvas, dtype=np.float32) / 255.0
96
  arr = (arr - 0.5) / 0.5
97
  tensor = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).contiguous()
 
99
  "pad_side": s, "out_size": size, "mode": mode}
100
 
101
 
102
+ def load_image_for_dust3r(path, size: int = IMG_SIZE, mode: str = DEFAULT_PREPROC, uint8: bool = False):
103
  """File-path convenience wrapper around :func:`preprocess_image` (eval scripts)."""
104
  with Image.open(Path(path)) as im:
105
+ return preprocess_image(im, size, mode, uint8=uint8)
106
 
107
 
108
  def intrinsics_to_canonical(K: np.ndarray, preproc: dict) -> np.ndarray:
code/models/demos/mast3r/tt/fattn.py ADDED
@@ -0,0 +1,130 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """MAST3R_OPT "fattn": model-local SDPA program (``ttnn.generic_op``; tt-metal untouched).
2
+
3
+ ttnn's own streaming SDPA compute kernel (``sdpa.cpp`` -> ``sdpa_standard_v2``) and its own writer, with
4
+ * a single K chunk = the whole key sequence (exact softmax: no online max / sum rescaling between K chunks),
5
+ * a model-local reader (``kernels/fattn_reader.cpp``) that keeps each head's K^T / V RESIDENT in L1 (loaded once
6
+ per core and head instead of once per q chunk, no inter-core KV chains),
7
+ * the flat B*H*q-chunk work split to the tile row (enc 9 / 8 rows per core instead of 2 x 5 = 10, no padded rows).
8
+ The compile-time configuration mirrors ttnn's SDPAProgramFactory for this case (non-causal, no mask, bf16, default
9
+ compute config HiFi2 / bf16 dest / approx, exp_approx_mode False).
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import math
14
+ import os
15
+ import struct
16
+
17
+ import ttnn
18
+
19
+ _KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
20
+ _SDPA_K = "ttnn/cpp/ttnn/operations/transformer/sdpa/device/kernels/"
21
+ _INACTIVE = 0xFFFFFFFF
22
+
23
+
24
+ def _largest_subblock(bh, bw, dst, max_h=1 << 30, max_w=1 << 30):
25
+ for h, w in ((2, 4), (4, 2), (1, 8), (8, 1), (1, 7), (7, 1), (2, 3), (3, 2), (1, 6), (6, 1),
26
+ (1, 5), (5, 1), (2, 2), (1, 4), (4, 1), (1, 3), (3, 1), (1, 2), (2, 1), (1, 1)):
27
+ if h * w > dst or h > max_h or w > max_w or bh % h or bw % w:
28
+ continue
29
+ return h, w
30
+ return 1, 1
31
+
32
+
33
+ def _gran(n, mx):
34
+ g = min(n, mx)
35
+ while g > 1 and n % g:
36
+ g -= 1
37
+ return g
38
+
39
+
40
+ def fattn(q, k, v, device, q_chunk_t=None, out_mc=None, ncores=None):
41
+ """softmax(q k^T / sqrt(dh)) v for q [B, H, Sq, dh], k / v [B, H, Sk, dh] (TILE bf16, interleaved)."""
42
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
43
+ B, NQH, Sq, DH = (int(d) for d in q.shape)
44
+ NKH, Sk = int(k.shape[1]), int(k.shape[2])
45
+ assert Sq % 32 == 0 and Sk % 32 == 0 and DH % 32 == 0 and NQH % NKH == 0
46
+ Sqt, Skt, DHt = Sq // 32, Sk // 32, DH // 32
47
+ vDHt = DHt
48
+ Sq_chunk_t = q_chunk_t or int(os.environ.get("MAST3R_FATTN_Q", "1"))
49
+ assert Sqt % Sq_chunk_t == 0
50
+ q_num_chunks = Sqt // Sq_chunk_t
51
+ Sk_chunk_t, k_num_chunks = Skt, 1
52
+ grid = device.compute_with_storage_grid_size()
53
+ gx, gy = grid.x, grid.y
54
+ n = ncores or gx * gy
55
+ total = B * NQH * q_num_chunks
56
+ base, extra = divmod(total, n)
57
+ max_chunks = base + (1 if extra else 0)
58
+ dst = 8
59
+ qk_h, qk_w = _largest_subblock(Sq_chunk_t, Sk_chunk_t, dst)
60
+ o_h, o_w = _largest_subblock(Sq_chunk_t, vDHt, dst, max_h=int(os.environ.get("MAST3R_FATTN_OH", "2")))
61
+ qktv_h = 2 if (o_h == 1 and 2 * o_w <= dst and Sq_chunk_t >= 2) else o_h
62
+ out0_t = 2 * qktv_h * vDHt
63
+ scale_bits = struct.unpack("<I", struct.pack("<f", 1.0 / math.sqrt(DH)))[0]
64
+ ident = 0x3F803F80
65
+
66
+ out = ttnn.allocate_tensor_on_device(ttnn.Shape([B, NQH, Sq, DH]), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
67
+
68
+ # CB ids in ttnn's allocation order (no mask / sink / paging CBs)
69
+ CB_Q, CB_K, CB_V, CB_ID, CB_COL, CB_RECIP, CB_QK, CB_OA, CB_OB, CB_MA, CB_MB, CB_SA, CB_SB, CB_EMD, CB_OUT = range(15)
70
+ q_tiles = Sq_chunk_t * DHt * (2 if max_chunks > 1 else 1)
71
+ sizes = {CB_Q: q_tiles, CB_K: Skt * DHt, CB_V: Skt * vDHt, CB_ID: 1, CB_COL: 1, CB_RECIP: 1,
72
+ CB_QK: Sq_chunk_t * Sk_chunk_t, CB_OA: Sq_chunk_t * vDHt, CB_OB: Sq_chunk_t * vDHt,
73
+ CB_MA: Sq_chunk_t, CB_MB: Sq_chunk_t, CB_SA: Sq_chunk_t, CB_SB: Sq_chunk_t, CB_EMD: Sq_chunk_t,
74
+ CB_OUT: out0_t}
75
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
76
+ cbs = [ttnn.CBDescriptor(total_size=nt * 2048, core_ranges=cores,
77
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=i, data_format=ttnn.bfloat16,
78
+ page_size=2048)])
79
+ for i, nt in sizes.items()]
80
+
81
+ ta_out = list(ttnn.TensorAccessorArgs(out).get_compile_time_args())
82
+ num_cores = gx * gy
83
+ cp_ct = [B, NQH, NKH, Skt, DHt, vDHt, Sq_chunk_t, q_num_chunks, Sk_chunk_t, k_num_chunks,
84
+ DHt, qk_w, qk_h, Sq_chunk_t // qk_h, Sk_chunk_t // qk_w, 1,
85
+ Sk_chunk_t, o_w, o_h, Sq_chunk_t // o_h, vDHt // o_w, 1,
86
+ num_cores, 0, 0, 0, 0, scale_bits, 0, 0, 1, Skt, 0, 0,
87
+ CB_Q, CB_K, CB_V, _INACTIVE, _INACTIVE, CB_ID, CB_COL, _INACTIVE, CB_RECIP, CB_OUT, CB_QK,
88
+ CB_OA, CB_OB, CB_MA, CB_MB, CB_SA, CB_SB, CB_EMD] + ta_out
89
+ wr_ct = [B, NQH, NKH, Sqt, Sqt, Sk, DHt, vDHt, Sq_chunk_t, q_num_chunks, Sk_chunk_t, k_num_chunks,
90
+ ident, scale_bits, num_cores, 0, 0, 0, 0, 0, 0, 1, o_h, 0, 0, 0] + ta_out + [0, 0] + \
91
+ [_INACTIVE, CB_ID, CB_COL, _INACTIVE, CB_OUT, CB_Q]
92
+ rd_ct = [q_num_chunks, Sq_chunk_t, Sqt, Skt, DHt, CB_Q, CB_K, CB_V, NQH // NKH] + \
93
+ list(ttnn.TensorAccessorArgs(q).get_compile_time_args()) + \
94
+ list(ttnn.TensorAccessorArgs(k).get_compile_time_args()) + \
95
+ list(ttnn.TensorAccessorArgs(v).get_compile_time_args())
96
+
97
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
98
+ for i in range(num_cores):
99
+ x, y = i % gx, i // gx
100
+ if i < n:
101
+ gs = i * base + min(i, extra)
102
+ gc = base + (1 if i < extra else 0)
103
+ else:
104
+ gs, gc = total, 0
105
+ rd_rt[x][y] = [q.buffer_address(), k.buffer_address(), v.buffer_address(), gs, gc]
106
+ wr_rt[x][y] = [out.buffer_address(), i, 1, 0, 0, 0, 0, 0, gs, gc, 0, 0]
107
+ cp_rt[x][y] = [i, 1, 0, 0, 0, gs, gc]
108
+ defines = [("STATS_GRANULARITY", str(_gran(Sq_chunk_t, dst))), ("SUB_EXP_GRANULARITY", str(_gran(Sk_chunk_t, dst))),
109
+ ("MUL_BCAST_GRANULARITY", str(_gran(Sq_chunk_t * Sk_chunk_t, dst))),
110
+ ("DHT_GRANULARITY", str(_gran(DHt, dst))), ("REDUCE_GRANULARITY", str(_gran(Sq_chunk_t, dst // 2))),
111
+ ("EXP_APPROX_MODE", "0")]
112
+ rd = ttnn.KernelDescriptor(
113
+ kernel_source=os.path.join(_KDIR, "fattn_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
114
+ core_ranges=cores, compile_time_args=rd_ct, runtime_args=rd_rt, config=ttnn.ReaderConfigDescriptor(),
115
+ defines=[(d, "1") for d in os.environ.get("MAST3R_FATTN_DIAG", "").split(",") if d])
116
+ wr = ttnn.KernelDescriptor(
117
+ kernel_source=_SDPA_K + "dataflow/writer_interleaved.cpp", source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
118
+ core_ranges=cores, compile_time_args=wr_ct, runtime_args=wr_rt, defines=defines,
119
+ config=ttnn.WriterConfigDescriptor())
120
+ ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi2, math_approx_mode=True,
121
+ fp32_dest_acc_en=False)
122
+ local = os.environ.get("MAST3R_FATTN_LOCAL", "0") == "1"
123
+ if local:
124
+ defines = defines + [("FSDPA_PROF", os.environ.get("MAST3R_FATTN_PROF", "0"))]
125
+ cp = ttnn.KernelDescriptor(
126
+ kernel_source=os.path.join(_KDIR, "fsdpa", "sdpa.cpp") if local else _SDPA_K + "compute/sdpa.cpp", source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
127
+ core_ranges=cores, compile_time_args=cp_ct, runtime_args=cp_rt, defines=defines, config=ccfg)
128
+ prog = ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs)
129
+ ttnn.generic_op([q, k, v, out], prog)
130
+ return out
code/models/demos/mast3r/tt/fused.py CHANGED
@@ -55,6 +55,171 @@ DEFAULT_TRACE_REGION_SIZE = 512 * 1024 * 1024 # bytes; measured: the 944-op tra
55
 
56
  _TRUE = ("1", "true", "yes", "on")
57
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
58
 
59
  def _truthy(v: Optional[str], default: bool) -> bool:
60
  if v is None or v.strip() == "":
@@ -94,6 +259,8 @@ class FusedConfig:
94
  ``MAST3R_TRACE`` ``1`` (default) | ``0``. Capture the whole device graph into one
95
  metal trace (persistent input buffer, replay per call).
96
  ``MAST3R_TRACE_REGION`` trace region bytes for ``ttnn.open_device`` (default 512 MiB).
 
 
97
  ``MAST3R_DPT_FUSE`` ``1`` (default) | ``0``. Exact DPT cleanups: TILE reshape of the
98
  four taps, relu fused into the resconv conv1 / head2 convs.
99
 
@@ -110,6 +277,8 @@ class FusedConfig:
110
  trace_region_size: int = DEFAULT_TRACE_REGION_SIZE
111
  dpt_fuse: bool = False
112
  conv_weight_cache: bool = False
 
 
113
 
114
  @classmethod
115
  def from_env(cls, env: Optional[Mapping[str, str]] = None) -> "FusedConfig":
@@ -131,11 +300,20 @@ class FusedConfig:
131
  if region <= 0:
132
  raise ValueError(f"MAST3R_TRACE_REGION must be > 0, got {region}")
133
  dpt_fuse = _truthy(env.get("MAST3R_DPT_FUSE"), True)
 
 
 
 
134
  return cls(
135
  enabled=True, rope=rope, rope_lut=rope_lut, mm=mm, sdpa_chunks=sdpa_chunks,
136
  trace=trace, trace_region_size=region, dpt_fuse=dpt_fuse, conv_weight_cache=True,
 
137
  )
138
 
 
 
 
 
139
  @property
140
  def fused_rope(self) -> bool:
141
  """True when q/k weights, biases and LUTs must be channel-permuted."""
@@ -167,8 +345,36 @@ class FusedConfig:
167
  kw = dict(base)
168
  if self.enabled and self.trace:
169
  kw["trace_region_size"] = self.trace_region_size
 
 
170
  return kw
171
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
172
 
173
  def parse_sdpa_chunks(value: Optional[str],
174
  default: Optional[Tuple[int, int]] = None) -> Optional[Tuple[int, int]]:
@@ -219,6 +425,15 @@ def expand_head_perm(perm: torch.Tensor, heads: int) -> torch.Tensor:
219
  return torch.cat([perm + h * dh for h in range(heads)])
220
 
221
 
 
 
 
 
 
 
 
 
 
222
  def permute_out_rows(weight: torch.Tensor, bias: Optional[torch.Tensor], idx: torch.Tensor):
223
  """Row gather of an ``nn.Linear`` weight ``(out, in)`` and bias ``(out,)``: exact."""
224
  w = weight[idx].contiguous()
 
55
 
56
  _TRUE = ("1", "true", "yes", "on")
57
 
58
+ #: Galaxy/p150 12x10 optimisation pass (2026-10-02, OPT_REPORT.md). ``MAST3R_OPT`` selects
59
+ #: them: ``all`` (default) | ``none`` | comma list. Each is an independent A/B knob.
60
+ #: out compact channel-first ROW_MAJOR head output (D2H 29 -> 3 ms)
61
+ #: mm per-shape matmul program configs tuned on 12x10 + exact GELU fused into fc1
62
+ #: gelua approximate GELU in the fused fc1 epilogue (needs "mm"; accuracy-gated, default off)
63
+ #: ``OPT_DEFAULT_OFF`` entries are only enabled when named explicitly (accuracy-gated A/B).
64
+ #: dpt DPT layout cleanup: refinenet out_conv before the upsample, RM conv in/out (no tilize)
65
+ #: l1 transient per-block activations (LN / qkv / heads / RoPE / SDPA / fc1) in L1 interleaved
66
+ #: lofi LoFi math fidelity for the transformer matmuls (accuracy-gated A/B, default off)
67
+ #: dptf HiFi2 (fp32 acc) math fidelity for the DPT convs except the two head convs (was HiFi4)
68
+ #: gelut tanh-form GELU (UnaryOpType.GELU_TANH, fp32 SFPU tanh; |tanh-form - erf-form| <= 4.7e-4)
69
+ #: in the fused fc1 epilogue instead of the erf-form piecewise polynomial (needs "mm")
70
+ #: lnfold fold the per-block LayerNorm affine (gamma, beta) into the following linear
71
+ #: (W' = W * gamma, b' = b + W @ beta, fp32 on host, then bf16): affine-free LN on device
72
+ #: phase DPT head tail (bilinear x2 -> 3x3 conv + relu -> 1x1 to 4) as 4 polyphase 3x3 convs on the
73
+ #: 256^2 grid with the bilinear weights folded into the kernels (exact algebra, fp64 fold), the
74
+ #: 2-pixel output ring recomputed exactly from 4 thin strips; no 512^2 x 128 intermediate.
75
+ #: Output [1, 1, 16, hw/4 + 2(H+W)] ROW_MAJOR; the host interleaves phases (needs "out", "dpt")
76
+ #: hostcol the per-call host cast of the two views writes the patch-embed im2col layout [2, N, 3*16*16]
77
+ #: bf16 directly (one strided copy, replaces cat + cast); the device skips its reshape/permute im2col
78
+ #: phase0 head.0 too: refinenet1's bilinear x2 + head.0 3x3 conv as 4 phase convs on the 128^2 grid,
79
+ #: phases interleaved on device (RM page concats) into the 256^2 head.0 output; the final
80
+ #: 6-px ring is recomputed exactly by a thin-strip run of the original ops (needs "phase")
81
+ #: decb decoder: the two branch blocks of a step share one B=2 split-heads / RoPE / SDPA / concat-heads
82
+ #: per attention (per-branch matmuls stay separate; outputs concatenated / sliced on batch)
83
+ #: resl1 encoder residual stream (patch-embed output and every residual matmul output) in L1 interleaved
84
+ #: (the fused residual op's ones vector follows it into L1); enc_norm output back in DRAM
85
+ #: ropes RoPE trans_mat HEIGHT_SHARDED on every core (one 32x32 tile each, per-pass L1 copy) -> the
86
+ #: prefill-sharded rotary_embedding_llama factory reads it locally (bit-identical; 26.7 -> 21.3 us enc)
87
+ #: sym return_pose (B-mode): one symmetric graph -- encoder once, decoder for both view orders, DPT
88
+ #: head 1 + head 2 of (1, 2) and head 1 of (2, 1); head 2 of (2, 1) (unused by the PairViewer pose)
89
+ #: skipped. Outputs bit-identical to the two-pass path. Server-side switch only.
90
+ #: tsplit with MAST3R_CQS=2: head 2's main rows and its ring strips are separate trace segments, so the
91
+ #: 2 MB main-rows readback (CQ1) overlaps the ~0.5 ms strip computation (CQ0); host assembles both
92
+ #: sdpa2 per-shape SDPA chunks balanced for 120 cores (flat B*H*q-chunk scheduling): encoder (B*H=32) q160/k512
93
+ #: (7 q chunks/head -> 224 chunks <= 2 per core), decoder (B*H=24) q224/k512 (5 chunks/head -> 120 chunks,
94
+ #: 1 per core); the kernel pads the partial last chunk. Probe (sdpa_probe3.py): 133 -> 94 us, 93 -> 67 us
95
+ #: hrope model-local generic_op kernel (tt/heads_rope.py + tt/kernels/heads_rope_*.cpp): split heads + RoPE of q and k
96
+ #: in one program on all 120 cores, reading the per-branch qkv / q / kv matmul outputs directly (no decoder
97
+ #: batch concat); same compute sequence as rotary_embedding_llama -> bit-identical (heads_rope_probe.py:
98
+ #: enc 72.4 -> 26.5 us, dec self 74.3 -> 24.2 us, dec cross 75.3 -> 24.0 us). Needs rope=llama
99
+ #: hcat model-local generic_op concat-heads kernel (data movement only, bit-exact) on all 120 cores; in the decoder
100
+ #: it writes the two branches' [1, N, D] outputs directly (no batch slices). Probe (heads_concat_probe.py):
101
+ #: enc 13.5 -> 9.6 us, dec 19.4 -> 8.3 us
102
+ #: mln decoder: the two branches' affine-free LayerNorms of a step run as ONE program (stock LN kernels and
103
+ #: exact stock compile-time config via generic_op, each core's runtime args point at its branch's tensor):
104
+ #: bit-identical, 64 cores instead of 2 x 32 (multi_ln_probe.py: 30.9 -> 21.3 us per pair)
105
+ #: mc2d plain linears (enc qkv, dec qkv / cq / ckv / embed) as full-grid 2-D mcast ttnn.linear with a per-shape
106
+ #: in0_block_w (8/16 for the K=768/1024/3072 decoder shapes) instead of minimal_matmul (mc2d_half_probe.py:
107
+ #: dec cq 22 -> 16 us, ckv 27 -> 22, qkv 33 -> 31, embed 27 -> 19); not bit-identical (different K blocking)
108
+ #: mc2dr the residual linears (proj / cproj / fc2) as 2-D mcast linear + separate add instead of the fused dit op
109
+ #: mm32 mc2d / mc2dr linears with fp32 dest accumulation (out subblocks <= 4): with bf16 dest the larger
110
+ #: in0_block_w accumulates more K tiles in 16-bit dest (real-pair depth error 1.13 -> 1.43 %); fp32 dest is
111
+ #: more accurate than the original dit / minimal_matmul (resid_acc_probe.py) at the same speed
112
+ #: addln every residual add that feeds a LayerNorm runs inside the LN program (stock FUSE_PRE_ADD reader/compute
113
+ #: path, model-local copy of layernorm.cpp that also packs the bf16 sum, model-local writer; cb_x kept bf16
114
+ #: so the LN sees the same bf16 sum ttnn.add produces): bit-identical, 2 programs -> 1 per residual
115
+ #: (multi_ln_probe.py add: dec pair 41.9 -> 24.6 us, enc 54.1 -> 31.2 us). Needs mc2dr, mln
116
+ #: cdbl DPT 3x3 convs at <= 64x64: activation + weight double buffering (64^2 height-sharded, 32^2 / 16^2
117
+ #: block-sharded); bit-identical (conv_cfg_probe.py: 64^2 96 -> 76 us, 32^2 51 -> 40, 16^2 45 -> 39)
118
+ #: pil phase0 interleave of the 4 head.0 phase convs by a model-local data-movement kernel on ROW_MAJOR pixel rows
119
+ #: (convs emit ROW_MAJOR; tt/kernels/phase_il_rm_*.cpp) instead of concat + transpose + 0/1 matmul + transpose
120
+ #: (exact; phase_il_probe.py: 30 us vs ~200 us per head)
121
+ #: til head-tail phase interleave (rows (a,b,c) x (i,j) -> NCHW) by a model-local data-movement kernel
122
+ #: (tt/kernels/tail_il_*.cpp) instead of RM reshape + permute + reshape + tilize + 0/1 matmul + untilize (exact;
123
+ #: tail_il_probe.py: 35 us vs ~130 us per head)
124
+ #: dmm decoder: the two branches' same-shape linears (qkv, proj, cq, ckv, cproj, fc1+GELU, fc2) as ONE program
125
+ #: each: ttnn's own 2-D mcast matmul descriptors placed on the top / bottom 12x5 halves (allowed_worker_cores)
126
+ #: and merged (ttnn.merge_program_descriptors), run through generic_op; bit-identical to full-grid linears
127
+ #: (dual_mm_probe.py rows: per pair 60 -> 51 us qkv, 30 -> 25 proj/cq/cproj, 44 -> 37 ckv, 154 -> 132 fc1,
128
+ #: 72 -> 63 fc2). Needs mc2d, mc2dr, addln
129
+ #: dadd DPT refinenet with skip: x + resConfUnit1(skip) and resConfUnit2's leading relu as ONE model-local kernel
130
+ #: (two bf16-rounded FPU adds in ttnn.add's order + SFPU relu; bit-identical; add3_probe.py: 128^2 165 -> 104 us,
131
+ #: 64^2 56 -> 36, 32^2 21 -> 11)
132
+ #: addln2 addln at the stack ends too: the encoder's last add + enc_norm and the decoder's last add + dec_norm as the
133
+ #: fused add + LN program with gamma/beta (stock FUSE_GAMMA/FUSE_BETA path); the block-5/8 taps are copied from
134
+ #: the next step's fused add + LN sums instead of separate adds (bit-identical)
135
+ #: rchain DPT resConfUnit conv1 -> conv2 on ttnn's L1 conv path: conv1's sharded output feeds conv2 directly (no
136
+ #: S2I to DRAM + I2S back; no DRAM-path NHWC reshape copy at 16^2); >= 64^2 the relu input is height-sharded
137
+ #: into the conv's own spec and freed after halo (L1 CB budget at 128^2), < 64^2 it lives in L1 interleaved
138
+ #: (bit-identical; rchain_probe.py per resconv: 128^2 376 -> 322 us, 64^2 148 -> 128, 32^2 75 -> 68, 16^2 76 -> 47)
139
+ #: dlin DPT refinenet out_conv (1x1) with its input (resConfUnit2's sum) in L1 and an explicit full-grid program
140
+ #: config with the auto config's in0_block_w / compute config (bit-identical); output in L1 below 128^2
141
+ #: (dlin_probe.py: 16384x256x256 96 -> 30 us, 4096: 33 -> 16, 1024: 15 -> 10). Needs rchain
142
+ #: dfront DPT front end: the 4 tap projections (dlin configs, auto in0_block_w) write L1, so the ConvTranspose /
143
+ #: ap3_down / layer_rn convs run on ttnn's L1 conv path (no DRAM-path NHWC reshape copies at 16^2, sharded
144
+ #: ap3_down output feeds layer4_rn); layer_rn outputs and the ap1 ConvTranspose output go to DRAM (l2_rn on the
145
+ #: sharded convT output would pick another config). Bit-identical. Needs dlin
146
+ #: dshard DPT >= 64^2: the refinenet adds read resConfUnit conv2's height shards in place and write relu(s) straight
147
+ #: into the height-shard spec of the next conv1 (model-local add3s / add2 kernels, same FPU adds -> bit-
148
+ #: identical); the leading relu of resConfUnit1 writes that spec directly. add3s_probe.py at 128^2: S2I + add3
149
+ #: + I2S 156 -> 72 us, S2I + add 77 -> 29 us. Needs rchain
150
+ #: ups1 refinenet upsample output staged in L1 (S2I to L1 + tilize L1 -> DRAM instead of S2I to DRAM + tilize
151
+ #: DRAM -> DRAM; exact)
152
+ #: tailf head tail after the 4 per-phase 1x1s: phase sum + head.4 bias + transpose + slice + untilize + NCHW
153
+ #: interleave as ONE model-local program (tt/kernels/tailf_*.cpp; FPU adds, exact transpose and 16-bit
154
+ #: interleave; bit-identical; tailf_probe.py at 256^2: 113 -> 55 us). Needs til
155
+ #: pemm patch-embed linear with an explicit full-grid 2-D config (auto config's in0_block_w / compute config:
156
+ #: bit-identical) writing the L1 residual directly (no 4 MB copy)
157
+ #: pcat attention output projections (encoder proj, decoder self / cross proj) read the SDPA output [B, H, N, dh]
158
+ #: directly as the concatenated-heads matrix: ttnn's own 2-D mcast matmul descriptor with the in0 sender reader
159
+ #: swapped for a model-local copy that remaps the tile ids (tt/kernels/mm_in0_heads_reader.cpp); no
160
+ #: concat-heads program (pcat_probe.py: enc 39.9 -> 33.3 us, dec pair 31.1 -> 26.7 us; bit-identical).
161
+ #: Needs hcat, mc2d (and dmm for the decoder)
162
+ #: p1hs refinenet1's out_conv writes the height-shard spec of head.0's phase convs directly; its DRAM copy for the
163
+ #: ring strips is one S2I (instead of a DRAM write + an I2S). Bit-identical. Needs dlin
164
+ #: hf3 HiFi3 (fp32 acc) instead of HiFi4 for the two output-resolution head convs (head.0 / head.2 phase convs and
165
+ #: their ring-strip convs). Not bit-identical (drops only the lo x lo mantissa partial product): real pair
166
+ #: gates unchanged to 4-5 digits (conf h1 0.99149 -> 0.99150 float); eth12 trace 28.72 -> 28.32 ms
167
+ #: gpoly fc1 GELU epilogue (encoder fc1 and the decoder fc1 pairs) as a model-local SFPU erf-GELU minimax polynomial
168
+ #: (kernels/mast3r_gelu_poly.h: 0.5x + |x| s q(s^2), s = min(|x|, 4.25), 9 coefficients, tools_prof/gelu_fit.py)
169
+ #: in the GELU_TANH slot of a model-local copy of ttnn's matmul compute kernel (kernels/mm_gelu_compute.cpp, same
170
+ #: program otherwise). Max |err| vs the erf GELU of the reference 6e-5 (gelut: 4.7e-4); ~30 instead of ~40 SFPU
171
+ #: instructions per row (rejected_gelu/gelu_fast_probe.py EXP=5: enc fc1 198 -> 155 us). Needs gelut
172
+ #: qkx decoder: n_b = LN(x_b) feeds branch b's qkv AND the other branch's cross k/v (norm_y(x_b) == n_b with the
173
+ #: affine folded), so each branch runs ONE linear with concatenated weights [qkv_b | ckv_other] (N = 3840, same
174
+ #: in0_block_w -> bit-identical); the self-attention RoPE reads q/k/v at row stride 5*Ht, the cross RoPE reads
175
+ #: k/v from the other branch's output (cols 3Ht..5Ht). One dual program per step instead of two. Needs dmm
176
+ #: sdbl DPT ring-strip convs (head.0 / head.2 on the thin strips) with act + weight double buffering (height-sharded,
177
+ #: bit-identical; strip_conv_probe.py: 68 -> 60 us and 74 -> 69 us per strip conv)
178
+ #: pilhs the phase0 interleave (pil) writes the head.2 phase convs' HEIGHT_SHARDED input spec directly (one contiguous
179
+ #: NoC write per 32-row unit into its shard; tt/kernels/phase_il_rm_writer_hs.cpp) instead of L1 interleaved + an
180
+ #: interleaved-to-sharded copy (~21 us per head). Exact. Needs pil
181
+ #: tups DPT refinenet upsample (32^2 -> 64^2, 64^2 -> 128^2): bilinear x2 (half-pixel, clamped == ttnn.upsample) straight
182
+ #: from the TILE out_conv output to the TILE DRAM tensor by a model-local kernel (tt/kernels/ups2_*.cpp; ups2h_* for the
183
+ #: 16-px-wide 16^2 -> 32^2 stage, two image rows per tile row): each output
184
+ #: tile = sum of <= 4 tile matmuls with constant interpolation tiles (entries exact in bf16, HiFi4, fp32 dest: exact
185
+ #: products and sums, one bf16 rounding). Replaces untilize + I2S + halo + upsample + S2I + tilize (tups_probe.py:
186
+ #: 64^2 x 256 136 -> 43 us, 32^2 108 -> 20 us). Not bit-identical: the stock upsample rounds in a bf16 dest
187
+ #: (23 % of its outputs differ from the correctly rounded value, 3.5 % for this kernel)
188
+ #: sgat DPT ring strips: their inputs (top/bottom 3 rows, left/right 3 columns of refinenet1's 128^2 output, the latter
189
+ #: already transposed) gathered straight from the height shards by a model-local data-movement kernel
190
+ #: (tt/kernels/strip_gather.cpp; 32-byte face-row reads) and carried to the strip segment, instead of an 8 MB S2I to
191
+ #: DRAM + untilize + 4 slices + 2 concats + permute. Exact. Needs p1hs
192
+ #: tapm DPT front: the 4 tap projections (1x1 linears, DRAM-latency bound, 19-22 us each on 32-96 cores) as ONE merged
193
+ #: program on disjoint column blocks (6 / 3 / 2 / 1 columns for ap3 / ap2 / ap1 / ap0), dlin's in0_block_w ->
194
+ #: bit-identical. Needs dfront
195
+ #: obs encoder qkv / proj / fc2 write BLOCK_SHARDED outputs (= the 2-D mcast config's per-core block: the matmul
196
+ #: packs into its own shard, no output NoC writes; r6_mm_sweep.py OUTM=bs: qkv 75 -> 60 us, proj 31 -> 26, fc2
197
+ #: 92 -> 87). hrope reads the qkv shards through a TensorAccessor, the fused add + LN reads proj / fc2 through its
198
+ #: stock TensorAccessor. Bit-identical. Needs hrope, addln, addln2
199
+ #: obsd decoder: the dual (top / bottom half) qkx, proj, cproj and fc2 linears write per-half BLOCK_SHARDED outputs
200
+ #: (qkx 74 -> 58 us); hrope takes up to two sharded layouts, the fused add + LN one reader kernel per layout.
201
+ #: Bit-identical. Needs dmm, qkx, pcat
202
+ #: resmm encoder: the residual adds folded into the proj / fc2 matmul FUSE_BIAS epilogues (model-local compute kernel
203
+ #: mm_gelu_compute.cpp with MAST3R_RESID_CB: dst += resid tile via dest-reuse ELWADD; the residual stream stays
204
+ #: BLOCK_SHARDED in the obs output spec and is read through a CB on its shard). The LayerNorms read the sum only
205
+ #: (no add, no second write). Not bit-identical (one bf16 rounding of the sum instead of two). Needs obs, pcat,
206
+ #: mm32
207
+ #: resmmd decoder: as resmm for the dual proj / cproj / fc2 linears (per-half block-sharded residual streams, obsd
208
+ #: specs); the LayerNorms of both branches read the sums only. Needs obsd, pcat, mm32, qkx, dmm
209
+ #: esplit the final enc_norm LayerNorm writes the two views' encoder outputs as two tensors (rows of view 1 / view 2,
210
+ #: each core's writer points at its view's tensor) instead of one [2, N, 1024] tensor + two DRAM slices.
211
+ #: Bit-identical. Needs resmm
212
+ #: demb the two views' decoder embeds as one dual (top / bottom 12x5) program writing decoder step 0's per-half
213
+ #: block-sharded residual specs directly (replaces 2 full-grid linears + 2 I2S). Bit-identical. Needs dmm, obsd,
214
+ #: resmmd, resl1d, decb
215
+ #: hrqk hrope compute: q and k of a unit in one pass per stage (4 tiles per fp32 dest half instead of 2 x 2).
216
+ #: Same per-tile ops -> bit-identical (r9_hrope_probe.py: enc 28.9 -> 28.2 us)
217
+ #: resl1d decoder residual streams in L1 (decoder-embed output in L1; the block-5/8 taps copied to DRAM, dec_norm
218
+ #: outputs in DRAM, so nothing in L1 survives into the DPT heads)
219
+ OPT_NAMES = ("out", "mm", "gelua", "dpt", "l1", "lofi", "dptf", "gelut", "lnfold", "phase", "hostcol", "phase0",
220
+ "decb", "resl1", "ropes", "sym", "tsplit", "sdpa2", "hrope", "hcat", "mln", "mc2d", "mc2dr", "resl1d", "addln", "cdbl", "pil", "til", "mm32", "dmm", "dadd", "addln2", "rchain", "dlin", "dfront", "dshard", "ups1", "tailf", "pemm", "pcat", "p1hs", "hf3", "gpoly", "qkx", "sdbl", "pilhs", "tups", "sgat", "tapm", "obs", "obsd", "resmm", "resmmd", "esplit", "demb", "hrqk")
221
+ OPT_DEFAULT_OFF: Tuple[str, ...] = ("gelua", "lofi")
222
+
223
 
224
  def _truthy(v: Optional[str], default: bool) -> bool:
225
  if v is None or v.strip() == "":
 
259
  ``MAST3R_TRACE`` ``1`` (default) | ``0``. Capture the whole device graph into one
260
  metal trace (persistent input buffer, replay per call).
261
  ``MAST3R_TRACE_REGION`` trace region bytes for ``ttnn.open_device`` (default 512 MiB).
262
+ ``MAST3R_CQS`` ``1`` (default) | ``2``. With 2 command queues the trace is split after
263
+ DPT head 1 and head 1's readback (CQ1) overlaps head 2 (CQ0).
264
  ``MAST3R_DPT_FUSE`` ``1`` (default) | ``0``. Exact DPT cleanups: TILE reshape of the
265
  four taps, relu fused into the resconv conv1 / head2 convs.
266
 
 
277
  trace_region_size: int = DEFAULT_TRACE_REGION_SIZE
278
  dpt_fuse: bool = False
279
  conv_weight_cache: bool = False
280
+ opts: Tuple[str, ...] = ()
281
+ num_cqs: int = 1
282
 
283
  @classmethod
284
  def from_env(cls, env: Optional[Mapping[str, str]] = None) -> "FusedConfig":
 
300
  if region <= 0:
301
  raise ValueError(f"MAST3R_TRACE_REGION must be > 0, got {region}")
302
  dpt_fuse = _truthy(env.get("MAST3R_DPT_FUSE"), True)
303
+ opts = parse_opts(env.get("MAST3R_OPT"))
304
+ num_cqs = int(env.get("MAST3R_CQS") or 1)
305
+ if num_cqs not in (1, 2):
306
+ raise ValueError(f"MAST3R_CQS must be 1 or 2, got {num_cqs}")
307
  return cls(
308
  enabled=True, rope=rope, rope_lut=rope_lut, mm=mm, sdpa_chunks=sdpa_chunks,
309
  trace=trace, trace_region_size=region, dpt_fuse=dpt_fuse, conv_weight_cache=True,
310
+ opts=opts, num_cqs=num_cqs,
311
  )
312
 
313
+ def has(self, name: str) -> bool:
314
+ """Is optimisation ``name`` (one of :data:`OPT_NAMES`) enabled?"""
315
+ return name in self.opts
316
+
317
  @property
318
  def fused_rope(self) -> bool:
319
  """True when q/k weights, biases and LUTs must be channel-permuted."""
 
345
  kw = dict(base)
346
  if self.enabled and self.trace:
347
  kw["trace_region_size"] = self.trace_region_size
348
+ if self.num_cqs == 2:
349
+ kw["num_command_queues"] = 2
350
  return kw
351
 
352
+ @property
353
+ def split_io(self) -> bool:
354
+ """``MAST3R_CQS=2`` (traced fused path): the forward is captured as two traces (A: encoder,
355
+ decoder, DPT head 1; B: DPT head 2); head 1 is read back on CQ1 (event-synced after A) while
356
+ CQ0 runs B. Needs a device opened with 2 command queues (Tensix dispatch on Galaxy)."""
357
+ return self.enabled and self.trace and self.num_cqs == 2
358
+
359
+
360
+ def parse_opts(value: Optional[str]) -> Tuple[str, ...]:
361
+ """``MAST3R_OPT``: unset / ``all`` -> every :data:`OPT_NAMES`; ``none`` / ``0`` -> ();
362
+ otherwise a comma list (``-name`` removes one from ``all``: ``all,-gelu``)."""
363
+ v = (value or "all").strip().lower()
364
+ if v in ("none", "0", "off", ""):
365
+ return ()
366
+ sel = set()
367
+ for tok in (t.strip() for t in v.split(",") if t.strip()):
368
+ if tok == "all":
369
+ sel |= set(n for n in OPT_NAMES if n not in OPT_DEFAULT_OFF)
370
+ elif tok.startswith("-"):
371
+ sel.discard(tok[1:])
372
+ elif tok in OPT_NAMES:
373
+ sel.add(tok)
374
+ else:
375
+ raise ValueError(f"MAST3R_OPT: unknown optimisation {tok!r}; known {OPT_NAMES}")
376
+ return tuple(n for n in OPT_NAMES if n in sel)
377
+
378
 
379
  def parse_sdpa_chunks(value: Optional[str],
380
  default: Optional[Tuple[int, int]] = None) -> Optional[Tuple[int, int]]:
 
425
  return torch.cat([perm + h * dh for h in range(heads)])
426
 
427
 
428
+ def fold_ln_affine(weight: torch.Tensor, bias: Optional[torch.Tensor], gamma: torch.Tensor, beta: torch.Tensor):
429
+ """``Linear(LN(x) * gamma + beta)`` == ``Linear'(LN(x))`` with ``W' = W * gamma`` (scale the input
430
+ columns of the ``(out, in)`` weight) and ``b' = b + W @ beta``. Computed in fp32 (exact up to the
431
+ final bf16 rounding of ``W'`` / ``b'`` at upload)."""
432
+ w = weight.float()
433
+ b = torch.zeros(w.shape[0]) if bias is None else bias.float()
434
+ return (w * gamma.float()[None, :]).contiguous(), (b + w @ beta.float()).contiguous()
435
+
436
+
437
  def permute_out_rows(weight: torch.Tensor, bias: Optional[torch.Tensor], idx: torch.Tensor):
438
  """Row gather of an ``nn.Linear`` weight ``(out, in)`` and bias ``(out,)``: exact."""
439
  w = weight[idx].contiguous()
code/models/demos/mast3r/tt/heads_rope.py ADDED
@@ -0,0 +1,1138 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Model-local fused split-heads + RoPE kernel (``ttnn.generic_op``; tt-metal untouched).
2
+
3
+ ``heads_rope(srcs, cos, sin, trans_mat, H, S, dh, device, ...) -> (q, k, v)``, each ``[B, H, S, dh]`` TILE bf16.
4
+
5
+ ``srcs`` lists, per output batch index b, ``(q_tensor, q_base, q_w, k_tensor, k_base, k_w, v_tensor, v_base, v_w)``
6
+ where ``*_base`` is the tile offset of head 0 / seq-tile 0 of that batch and ``*_w`` the row width in tiles of the
7
+ source tensor (all sources TILE bf16 interleaved, same buffer type). One program replaces
8
+ ``nlp_create_qkv_heads`` (+ the batch concat of per-branch sources in the decoder) + 2 x ``rotary_embedding_llama``;
9
+ the compute kernel runs the stock rotary op's exact op sequence (bit-identical at the same compute config).
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import math
14
+ import os
15
+
16
+ import ttnn
17
+
18
+ _KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
19
+ _TILE = 2048
20
+
21
+
22
+ def _cb(idx, ntiles, cores):
23
+ return ttnn.CBDescriptor(
24
+ total_size=ntiles * _TILE, core_ranges=cores,
25
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=idx, data_format=ttnn.bfloat16, page_size=_TILE)])
26
+
27
+
28
+ def _is_dram(t):
29
+ return t.memory_config().buffer_type == ttnn.BufferType.DRAM
30
+
31
+
32
+ def heads_rope(srcs, cos, sin, trans_mat, H, S, dh, device, out_mc=None, ncores=None, qk2=False):
33
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
34
+ B = len(srcs)
35
+ St, Wt = S // 32, dh // 32
36
+ grid = device.compute_with_storage_grid_size()
37
+ gx, gy = grid.x, grid.y
38
+ n = ncores or gx * gy
39
+ U = B * St * H
40
+ base, extra = divmod(U, n)
41
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
42
+ shape = ttnn.Shape([B, H, S, dh])
43
+ q = ttnn.allocate_tensor_on_device(shape, ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
44
+ k = ttnn.allocate_tensor_on_device(shape, ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
45
+ v = ttnn.allocate_tensor_on_device(shape, ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
46
+ src_dram = any(_is_dram(t) for s_ in srcs for t in (s_[0], s_[3], s_[6]))
47
+ # MAST3R_OPT "obs": sharded sources (block-sharded matmul outputs) are read through a TensorAccessor; up to two
48
+ # distinct sharded layouts (kind 1 / 2), interleaved sources keep the pow2 address generator (kind 0)
49
+ layouts = []
50
+
51
+ def kind(t):
52
+ if not t.memory_config().is_sharded():
53
+ assert _is_dram(t) == src_dram
54
+ return 0
55
+ ta = list(ttnn.TensorAccessorArgs(t).get_compile_time_args())
56
+ if ta not in layouts:
57
+ layouts.append(ta)
58
+ return 1 + layouts.index(ta)
59
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
60
+ src_args = []
61
+ for (qt, qb, qw, kt, kb, kw, vt, vb, vw) in srcs:
62
+ src_args += [qt.buffer_address(), qb, qw, kind(qt), kt.buffer_address(), kb, kw, kind(kt),
63
+ vt.buffer_address(), vb, vw, kind(vt)]
64
+ assert len(layouts) <= 2
65
+ ta_src = [v_ for l_ in layouts for v_ in l_]
66
+ start = 0
67
+ for i in range(gx * gy):
68
+ x, y = i % gx, i // gx
69
+ cnt = (base + (1 if i < extra else 0)) if i < n else 0
70
+ rd_rt[x][y] = [start, cnt]
71
+ wr_rt[x][y] = [start, cnt]
72
+ cp_rt[x][y] = [start, cnt]
73
+ start += cnt
74
+ assert start == U
75
+ ct = [H, St, Wt]
76
+ diag = [(d_, "1") for d_ in os.environ.get("MAST3R_HROPE_DIAG", "").split(",") if d_] # probes only
77
+ rd = ttnn.KernelDescriptor(
78
+ kernel_source=os.path.join(_KDIR, "heads_rope_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
79
+ core_ranges=cores, compile_time_args=ct + [B] + ta_src, runtime_args=rd_rt,
80
+ common_runtime_args=[cos.buffer_address(), sin.buffer_address(), trans_mat.buffer_address()] + src_args,
81
+ defines=[("COS_DRAM", "true" if _is_dram(cos) else "false"), ("TRANS_DRAM", "true" if _is_dram(trans_mat) else "false"),
82
+ ("SRC_DRAM", "true" if src_dram else "false")] + [("NTA", str(len(layouts)))] + diag,
83
+ config=ttnn.ReaderConfigDescriptor())
84
+ wr = ttnn.KernelDescriptor(
85
+ kernel_source=os.path.join(_KDIR, "heads_rope_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
86
+ core_ranges=cores, compile_time_args=ct, runtime_args=wr_rt,
87
+ common_runtime_args=[q.buffer_address(), k.buffer_address(), v.buffer_address()],
88
+ defines=[("OUT_DRAM", "true" if out_mc.buffer_type == ttnn.BufferType.DRAM else "false")] + diag,
89
+ config=ttnn.WriterConfigDescriptor())
90
+ cp = ttnn.KernelDescriptor(
91
+ kernel_source=os.path.join(_KDIR, "heads_rope_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
92
+ core_ranges=cores, compile_time_args=ct, runtime_args=cp_rt, defines=diag + ([("HROPE_QK2", "1")] if qk2 else []),
93
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
94
+ fp32_dest_acc_en=True))
95
+ m = 2 if qk2 else 1 # MAST3R_OPT "hrqk": q + k per stage -> the stage CBs hold 2 * Wt tiles
96
+ cbs = [_cb(0, 4 * Wt, cores), _cb(1, 2 * Wt, cores), _cb(2, 2 * Wt, cores), _cb(3, 1, cores), _cb(4, 2 * Wt, cores),
97
+ _cb(16, 4 * Wt, cores), _cb(24, m * Wt, cores), _cb(25, m * Wt, cores), _cb(26, m * Wt, cores)]
98
+ prog = ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs)
99
+ ins = []
100
+ for s in srcs:
101
+ for t in (s[0], s[3], s[6]):
102
+ if all(t is not u for u in ins):
103
+ ins.append(t)
104
+ ttnn.generic_op(ins + [cos, sin, trans_mat, q, k, v], prog)
105
+ return q, k, v
106
+
107
+
108
+ def heads_concat(ctx, outs, H, S, dh, device, ub=2):
109
+ """Concatenate heads: ``ctx`` ``[B, H, S, dh]`` TILE bf16 interleaved -> per batch b the rows of
110
+ ``outs[b] = (tensor, base_tile, row_w_tiles)`` (pre-allocated, e.g. ``[1, S, H*dh]`` per branch, or one
111
+ ``[B, S, H*dh]`` with base b*St*row_w). Pure data movement (bit-exact), all cores."""
112
+ B = len(outs)
113
+ St, Wt = S // 32, dh // 32
114
+ grid = device.compute_with_storage_grid_size()
115
+ gx, gy = grid.x, grid.y
116
+ n = gx * gy
117
+ U = B * St * H
118
+ base, extra = divmod(U, n)
119
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
120
+ rd_rt, wr_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
121
+ o_args = []
122
+ for (t, b0, w) in outs:
123
+ o_args += [t.buffer_address(), b0, w]
124
+ start = 0
125
+ for i in range(n):
126
+ x, y = i % gx, i // gx
127
+ cnt = base + (1 if i < extra else 0)
128
+ rd_rt[x][y] = [ctx.buffer_address(), start, cnt]
129
+ wr_rt[x][y] = [start, cnt] + o_args
130
+ start += cnt
131
+ ct = [H, St, Wt, ub]
132
+ rd = ttnn.KernelDescriptor(
133
+ kernel_source=os.path.join(_KDIR, "heads_concat_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
134
+ core_ranges=cores, compile_time_args=ct, runtime_args=rd_rt,
135
+ defines=[("IN_DRAM", "true" if _is_dram(ctx) else "false")], config=ttnn.ReaderConfigDescriptor())
136
+ wr = ttnn.KernelDescriptor(
137
+ kernel_source=os.path.join(_KDIR, "heads_concat_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
138
+ core_ranges=cores, compile_time_args=ct + [B], runtime_args=wr_rt,
139
+ defines=[("OUT_DRAM", "true" if _is_dram(outs[0][0]) else "false")], config=ttnn.WriterConfigDescriptor())
140
+ prog = ttnn.ProgramDescriptor(kernels=[rd, wr], semaphores=[], cbs=[_cb(0, 2 * ub * Wt, cores)])
141
+ ins = []
142
+ for o in outs:
143
+ if all(o[0] is not u for u in ins):
144
+ ins.append(o[0])
145
+ ttnn.generic_op([ctx] + ins, prog)
146
+ return [o[0] for o in outs]
147
+
148
+
149
+ _LN_K = "ttnn/cpp/ttnn/operations/normalization/layernorm/device/kernels/"
150
+
151
+
152
+ # probes only (tools_prof/r9_ln_probe.py): model-local reader for the plain (non-fused, affine-free) multi_layernorm
153
+ _LN_READER_OVERRIDE = [os.environ.get("MAST3R_LN_READER") or None]
154
+ _LN_READER_DEFS = [[(d_.split("=")[0], d_.split("=")[1]) for d_ in os.environ.get("MAST3R_LN_READER_DEFS", "").split(",") if d_]]
155
+
156
+
157
+ def multi_layernorm(xs, eps, device, out_mc=None, residual_ys=None, res_mc=None, gamma=None, beta=None, fast=False, split=False,
158
+ out_split=1):
159
+ """Affine-free LayerNorm of several same-shape TILE bf16 interleaved tensors in ONE program: the stock
160
+ ttnn.layer_norm kernels (reader_unary_interleaved_ln / layernorm.cpp / writer_unary_interleaved_start_id_blocked)
161
+ with the exact compile-time config ttnn's LayerNormMultiCoreProgramFactory builds for this case (default
162
+ compute config HiFi4 + fp32 dest, non-Welford, non-large-tensor), one tile row per core, but each core's runtime
163
+ args point at the tensor its row belongs to. Bit-identical to len(xs) separate ttnn.layer_norm calls.
164
+
165
+ ``residual_ys``: fused residual add -- normalises s_i = xs[i] + residual_ys[i] (stock FUSE_PRE_ADD path, the sum
166
+ kept in fp32 inside the kernel) and ALSO writes s_i (bf16) with the model-local ln_add_compute / ln_add_writer
167
+ kernels. Returns (normed list, sums list). Replaces ttnn.add + ttnn.layer_norm (not bit-identical: the LN sees
168
+ the unrounded sum).
169
+
170
+ ``out_split`` = k (one input, no residual): the normed output is written as k tensors of shape
171
+ ``[shape[0] // k, *shape[1:]]`` (consecutive row ranges; MAST3R_OPT "esplit": the two views' encoder outputs
172
+ without a batched tensor + slices). Same kernels and per-core math; only each writer's base address / start tile
173
+ differs -> bit-identical to slicing the single output."""
174
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
175
+ fuse = residual_ys is not None
176
+ x0 = xs[0]
177
+ shape = [int(d) for d in x0.shape]
178
+ W = int(shape[-1])
179
+ rows_per = 1
180
+ for d in shape[:-1]:
181
+ rows_per *= int(d)
182
+ rows_per //= 32
183
+ Wt = W // 32
184
+ grid = device.compute_with_storage_grid_size()
185
+ gx, gy = grid.x, grid.y
186
+ total = rows_per * len(xs)
187
+ assert total <= gx * gy, "one tile row per core"
188
+ if out_split > 1:
189
+ assert len(xs) == 1 and not fuse and shape[0] % out_split == 0 and rows_per % out_split == 0
190
+ oshape = [shape[0] // out_split] + shape[1:]
191
+ outs = [ttnn.allocate_tensor_on_device(ttnn.Shape(oshape), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
192
+ for _ in range(out_split)]
193
+ rows_out = rows_per // out_split
194
+ else:
195
+ outs = [ttnn.allocate_tensor_on_device(ttnn.Shape(shape), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc) for _ in xs]
196
+ rows_out = rows_per
197
+ sums = []
198
+ if fuse:
199
+ res_mc = res_mc or x0.memory_config()
200
+ sums = [ttnn.allocate_tensor_on_device(ttnn.Shape(shape), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, res_mc)
201
+ for _ in xs]
202
+ block = 4
203
+ import struct
204
+ eps_bits = struct.unpack("<I", struct.pack("<f", eps))[0]
205
+ packed_one = 0x3F803F80
206
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
207
+ rd_vals, wr_vals = {}, {}
208
+ for i in range(total):
209
+ x, y = i % gx, i // gx
210
+ t, r = divmod(i, rows_per)
211
+ rd_vals[(x, y)] = [xs[t].buffer_address(), 1, Wt, r, packed_one, eps_bits,
212
+ gamma.buffer_address() if gamma is not None else 0,
213
+ beta.buffer_address() if beta is not None else 0,
214
+ residual_ys[t].buffer_address() if fuse else 0]
215
+ to, ro = (divmod(r, rows_out) if out_split > 1 else (t, r))
216
+ wr_rt[x][y] = wr_vals[(x, y)] = [outs[to].buffer_address(), Wt, 1, ro * Wt] + ([sums[t].buffer_address()] if fuse else [])
217
+ cp_rt[x][y] = [1]
218
+ ranges = []
219
+ full_rows, rem = divmod(total, gx)
220
+ if full_rows:
221
+ ranges.append(ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, full_rows - 1)))
222
+ if rem:
223
+ ranges.append(ttnn.CoreRange(ttnn.CoreCoord(0, full_rows), ttnn.CoreCoord(rem - 1, full_rows)))
224
+ cores = ttnn.CoreRangeSet(ranges)
225
+ ta_in = ttnn.TensorAccessorArgs(x0).get_compile_time_args()
226
+ ta_b = ttnn.TensorAccessorArgs(residual_ys[0]).get_compile_time_args() if fuse else [0, 0]
227
+ ta_out = ttnn.TensorAccessorArgs(outs[0]).get_compile_time_args()
228
+ named = {"cb_in": 0, "cb_inb": 1, "cb_scaler": 2, "cb_eps": 3, "cb_gamma": 5, "cb_beta": 6, "cb_out": 16, "cb_ex": 18,
229
+ "cb_ex2": 19, "cb_xmm2": 20, "cb_ex2pe": 21, "cb_fusion": 22, "cb_x": 23, "cb_xmm": 24, "cb_reciprocals": 25,
230
+ "cb_accumulate": 26, "cb_in_rm": 27, "cb_out_rm": 28, "cb_x_welford": 23 if fuse else 0,
231
+ "welford_fp32_alias": 0, "cb_ex_welford": 18, "cb_ex2_welford": 19, "welford_state_fp32_alias": 0}
232
+ named = list(named.items())
233
+ pre = [("FUSE_PRE_ADD", "1")] if fuse else []
234
+ ta_g = ttnn.TensorAccessorArgs(gamma).get_compile_time_args() if gamma is not None else [0, 0]
235
+ ta_be = ttnn.TensorAccessorArgs(beta).get_compile_time_args() if beta is not None else [0, 0]
236
+ rdef = list(pre) + ([("FUSE_GAMMA", "1")] if gamma is not None else []) + ([("FUSE_BETA", "1")] if beta is not None else [])
237
+ # MAST3R_OPT "obsd": the residual_ys may be block-sharded on different core grids (dual linears on the two grid
238
+ # halves) -> one reader kernel per distinct residual layout, each on the cores of its tensor's rows
239
+ groups = [] # [(ta_b, [core coords])]
240
+ for i in range(total):
241
+ t = i // rows_per
242
+ tb = list(ttnn.TensorAccessorArgs(residual_ys[t]).get_compile_time_args()) if fuse else list(ta_b)
243
+ tb = (tuple(ttnn.TensorAccessorArgs(xs[t]).get_compile_time_args()), tuple(tb))
244
+ if not groups or groups[-1][0] != tb:
245
+ if any(g_[0] == tb for g_ in groups):
246
+ gi = [g_[0] for g_ in groups].index(tb)
247
+ groups[gi][1].append((i % gx, i // gx))
248
+ continue
249
+ groups.append((tb, []))
250
+ groups[-1][1].append((i % gx, i // gx))
251
+ split = split and fuse
252
+ rds, wrs = [], []
253
+ for (ta_x, tb), crs in groups:
254
+ ta_x, tb = list(ta_x), list(tb)
255
+ rd_rt = ttnn.RuntimeArgs()
256
+ for (x_, y_) in crs:
257
+ rd_rt[x_][y_] = rd_vals[(x_, y_)]
258
+ rcores = cores if len(groups) == 1 else ttnn.CoreRangeSet(
259
+ [ttnn.CoreRange(ttnn.CoreCoord(x_, y_), ttnn.CoreCoord(x_, y_)) for (x_, y_) in crs])
260
+ rds.append(ttnn.KernelDescriptor(
261
+ kernel_source=os.path.join(_KDIR, "ln_reader_split.cpp") if split else
262
+ (_LN_READER_OVERRIDE[0] if (_LN_READER_OVERRIDE[0] and not fuse and gamma is None and beta is None
263
+ and Wt % block == 0) else _LN_K + "dataflow/reader_unary_interleaved_ln.cpp"),
264
+ source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
265
+ core_ranges=rcores, compile_time_args=[block, 0, W] + ta_x + tb + ta_g + ta_be + [2048],
266
+ named_compile_time_args=named, defines=rdef + ([("LN_SPLIT_B", "1")] if split else []) + _LN_READER_DEFS[0],
267
+ runtime_args=rd_rt, config=ttnn.ReaderConfigDescriptor()))
268
+ if split:
269
+ # MAST3R_OPT "lnsplit": the writer reads b on the other RISC / NoC (kernels/ln_add_writer_rb.cpp)
270
+ w_rt = ttnn.RuntimeArgs()
271
+ for (x_, y_) in crs:
272
+ w_rt[x_][y_] = wr_vals[(x_, y_)] + [rd_vals[(x_, y_)][8]]
273
+ wrs.append(ttnn.KernelDescriptor(
274
+ kernel_source=os.path.join(_KDIR, "ln_add_writer_rb.cpp"),
275
+ source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
276
+ core_ranges=rcores, compile_time_args=[block] + tb, runtime_args=w_rt,
277
+ defines=[("OUT_DRAM", "true" if _is_dram(outs[0]) else "false"),
278
+ ("RES_DRAM", "true" if _is_dram(sums[0]) else "false")],
279
+ config=ttnn.WriterConfigDescriptor()))
280
+ if split:
281
+ wr = None
282
+ elif fuse:
283
+ wr = ttnn.KernelDescriptor(
284
+ kernel_source=os.path.join(_KDIR, "ln_add_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
285
+ core_ranges=cores, compile_time_args=[block], runtime_args=wr_rt,
286
+ defines=[("OUT_DRAM", "true" if _is_dram(outs[0]) else "false"),
287
+ ("RES_DRAM", "true" if _is_dram(sums[0]) else "false")],
288
+ config=ttnn.WriterConfigDescriptor())
289
+ else:
290
+ wr = ttnn.KernelDescriptor(
291
+ kernel_source=_LN_K + "dataflow/writer_unary_interleaved_start_id_blocked.cpp",
292
+ source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH, core_ranges=cores, compile_time_args=[block] + ta_out,
293
+ named_compile_time_args=named, runtime_args=wr_rt, config=ttnn.WriterConfigDescriptor())
294
+ ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False, fp32_dest_acc_en=True)
295
+ utd = [ttnn.UnpackToDestMode.Default] * 64
296
+ utd[26] = ttnn.UnpackToDestMode.UnpackToDestFp32
297
+ ccfg.unpack_to_dest_mode = utd
298
+ fast = fast and fuse and gamma is None and beta is None and Wt % block == 0
299
+ if fast:
300
+ # MAST3R_OPT "lnf": model-local fewer-pass add + LN compute kernel (kernels/ln_fast_compute.cpp)
301
+ ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
302
+ fp32_dest_acc_en=True)
303
+ inv_w_bits = struct.unpack("<I", struct.pack("<f", 1.0 / W))[0]
304
+ cp = ttnn.KernelDescriptor(
305
+ kernel_source=os.path.join(_KDIR, "ln_fast_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
306
+ core_ranges=cores, compile_time_args=[Wt, block, inv_w_bits, eps_bits], runtime_args=cp_rt, config=ccfg,
307
+ defines=[("LNF_PROF", "1")] if os.environ.get("MAST3R_LNF_PROF") else [])
308
+ f32, b16 = ttnn.float32, ttnn.bfloat16
309
+
310
+ def cbf(idx, n, fmt, page):
311
+ return ttnn.CBDescriptor(total_size=n * page, core_ranges=cores, format_descriptors=[
312
+ ttnn.CBFormatDescriptor(buffer_index=idx, data_format=fmt, page_size=page)])
313
+ cbs = [cbf(0, 2 * block, b16, 2048), cbf(1, 2 * block, b16, 2048), cbf(2, 2, b16, 2048), cbf(3, 2, b16, 2048),
314
+ cbf(16, 2 * block, b16, 2048), cbf(17, 2 * block, b16, 2048), cbf(23, Wt, b16, 2048),
315
+ cbf(24, 1, b16, 2048), cbf(18, 1, f32, 4096), cbf(20, 1, f32, 4096), cbf(21, 1, f32, 4096)]
316
+ prog = ttnn.ProgramDescriptor(kernels=rds + (wrs if split else [wr]) + [cp], semaphores=[], cbs=cbs)
317
+ ttnn.generic_op(list(xs) + list(residual_ys) + outs + sums, prog)
318
+ return outs, sums
319
+ cp = ttnn.KernelDescriptor(
320
+ kernel_source=os.path.join(_KDIR, "ln_add_compute.cpp") if fuse else _LN_K + "compute/layernorm.cpp",
321
+ source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
322
+ core_ranges=cores, compile_time_args=[Wt, block, int(gamma is not None), int(beta is not None), 1, 1, 0, W, 32],
323
+ named_compile_time_args=named,
324
+ defines=pre, runtime_args=cp_rt, config=ccfg)
325
+
326
+ def cb(idx, n, fmt, page):
327
+ return ttnn.CBDescriptor(total_size=n * page, core_ranges=cores,
328
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=idx, data_format=fmt, page_size=page)])
329
+ f32, b16 = ttnn.float32, ttnn.bfloat16
330
+ Wnb = -(-Wt // block) * block
331
+ cbs = [cb(0, 2 * block if fuse else Wnb, b16, 2048), cb(16, 2 * block, b16, 2048), cb(18, 2, f32, 4096),
332
+ cb(2, 2, b16, 2048), cb(3, 2, b16, 2048), cb(19, 2, f32, 4096), cb(24, Wnb, f32, 4096),
333
+ cb(20, Wnb, f32, 4096), cb(21, 8, f32, 4096)]
334
+ if gamma is not None or beta is not None:
335
+ cbs += [cb(22, 2 * block, f32, 4096)]
336
+ if gamma is not None:
337
+ cbs += [cb(5, Wnb, b16, 2048)]
338
+ if beta is not None:
339
+ cbs += [cb(6, Wnb, b16, 2048)]
340
+ if fuse:
341
+ # cb_x in bf16 (stock: fp32): the LN then sees exactly the bf16 sum ttnn.add would produce -> bit-identical
342
+ cbs += [cb(23, Wnb, b16, 2048), cb(1, 2 * block, b16, 2048), cb(17, 2 * block, b16, 2048)]
343
+ prog = ttnn.ProgramDescriptor(kernels=rds + (wrs if split else [wr]) + [cp], semaphores=[], cbs=cbs)
344
+ aff = [t for t in (gamma, beta) if t is not None]
345
+ ttnn.generic_op(list(xs) + (list(residual_ys) if fuse else []) + aff + outs + sums, prog)
346
+ return (outs, sums) if fuse else outs
347
+
348
+
349
+ def _in0_dummy(x):
350
+ """A sharded in0 (rejected obsf: tools_prof/obsf_probe.py) is read by the stock interleaved in0 sender through a TensorAccessor built for
351
+ the sharded tensor: the descriptor is created for a shape-equal L1-interleaved placeholder, then patched."""
352
+ if not x.memory_config().is_sharded():
353
+ return None
354
+ return ttnn.allocate_tensor_on_device(x.shape, x.dtype, ttnn.TILE_LAYOUT, x.device(), ttnn.L1_MEMORY_CONFIG)
355
+
356
+
357
+ def _patch_in0_ta(desc, dummy, x):
358
+ ta_d = list(ttnn.TensorAccessorArgs(dummy).get_compile_time_args())
359
+ ta_x = list(ttnn.TensorAccessorArgs(x).get_compile_time_args())
360
+ assert not list(ttnn.TensorAccessorArgs(x).get_common_runtime_args())
361
+ n = 0
362
+ for k in desc.kernels:
363
+ if str(k.kernel_source).endswith(_IN0_SENDER):
364
+ cta = list(k.compile_time_args)
365
+ assert cta[28:28 + len(ta_d)] == ta_d, (cta[28:40], ta_d)
366
+ k.compile_time_args = cta[:28] + ta_x + cta[28 + len(ta_d):]
367
+ ra = k.runtime_args
368
+ for cr in k.core_ranges.ranges():
369
+ for xx in range(cr.start.x, cr.end.x + 1):
370
+ for yy in range(cr.start.y, cr.end.y + 1):
371
+ a = list(ra[xx][yy])
372
+ if a:
373
+ assert a[0] == dummy.buffer_address(), (a[0], dummy.buffer_address())
374
+ a[0] = x.buffer_address()
375
+ ra[xx][yy] = a
376
+ n += 1
377
+ assert n == 1, n
378
+
379
+
380
+ _RESMM_SRC = os.path.join(_KDIR, "mm_gelu_compute.cpp")
381
+
382
+
383
+ def colmajor_ta_args(x):
384
+ """MAST3R_OPT "obsc": compile-time TensorAccessor args of a "virtual transposed" view of the BLOCK_SHARDED (row-major
385
+ shard grid) TILE tensor ``x`` [.., M, K] whose shards hold their tiles in COLUMN-major order (what a matmul with
386
+ out_subblock_w == 1 packs into its own shard): pages [Kt, Mt], shard [sk, sm] (row-major inside the virtual shard
387
+ == column-major inside the real one), shard (i, j) on the core of real shard (j, i). Page id of tile (r, k) = k * Mt + r.
388
+ Format (tt_metal/impl/buffers/tensor_accessor_args.cpp): [cfg, page, rank, nbanks, tshape.., sshape.., coords..]."""
389
+ a = list(ttnn.TensorAccessorArgs(x).get_compile_time_args())
390
+ cfg, page, rank, nb = a[:4]
391
+ assert rank == 2, a[:8]
392
+ t = a[4:6]
393
+ sh = a[6:8]
394
+ co = a[8:]
395
+ nbanks = nb & 0xFFFF
396
+ assert len(co) == (nbanks + 1) // 2, (len(co), nbanks)
397
+ coords = []
398
+ for i in range(nbanks):
399
+ w = co[i // 2]
400
+ coords.append(((w >> 8) & 0xFF, w & 0xFF) if i % 2 == 0 else ((w >> 24) & 0xFF, (w >> 16) & 0xFF))
401
+ nr, nc = -(-t[0] // sh[0]), -(-t[1] // sh[1]) # real shard grid (rows = M blocks, cols = K blocks)
402
+ assert nr * nc == nbanks, (nr, nc, nbanks)
403
+ vco = [coords[j * nc + i] for i in range(nc) for j in range(nr)] # virtual shard (i, j) -> real shard (j, i)
404
+ packed = []
405
+ for i in range(0, nbanks, 2):
406
+ (x1, y1) = vco[i]
407
+ if i + 1 < nbanks:
408
+ (x2, y2) = vco[i + 1]
409
+ packed.append(((x2 & 0xFF) << 24) | ((y2 & 0xFF) << 16) | ((x1 & 0xFF) << 8) | (y1 & 0xFF))
410
+ else:
411
+ packed.append(((x1 & 0xFF) << 8) | (y1 & 0xFF))
412
+ return [cfg, page, rank, nb, t[1], t[0], sh[1], sh[0]] + packed
413
+
414
+
415
+ def _patch_in0_colmajor(desc, dummy, x, Mt, Kt):
416
+ """MAST3R_OPT "obsc": the stock interleaved in0 sender reads a column-major-shard BLOCK_SHARDED in0 (colmajor_ta_args):
417
+ tile (r, k) is virtual page k * Mt + r -> strides w = Mt, h = 1, next K block = in0_block_w * Mt, next h block =
418
+ in0_block_h; start tile r0 * Kt -> r0."""
419
+ ta_d = list(ttnn.TensorAccessorArgs(dummy).get_compile_time_args())
420
+ ta_x = colmajor_ta_args(x)
421
+ n = 0
422
+ for k in desc.kernels:
423
+ if str(k.kernel_source).endswith(_IN0_SENDER):
424
+ cta = list(k.compile_time_args)
425
+ assert cta[28:28 + len(ta_d)] == ta_d, (cta[28:40], ta_d)
426
+ assert cta[0] == 1 and cta[1] == Kt, cta[:4]
427
+ ibw, ibh = cta[4], cta[5]
428
+ assert cta[2] == ibw and cta[3] == ibh * Kt, cta[:6]
429
+ cta[0], cta[1], cta[2], cta[3] = Mt, 1, ibw * Mt, ibh
430
+ k.compile_time_args = cta[:28] + ta_x + cta[28 + len(ta_d):]
431
+ ra = k.runtime_args
432
+ for cr in k.core_ranges.ranges():
433
+ for xx in range(cr.start.x, cr.end.x + 1):
434
+ for yy in range(cr.start.y, cr.end.y + 1):
435
+ a = list(ra[xx][yy])
436
+ if a:
437
+ assert a[0] == dummy.buffer_address(), (a[0], dummy.buffer_address())
438
+ assert a[1] % Kt == 0, (a[1], Kt)
439
+ a[0] = x.buffer_address()
440
+ a[1] = a[1] // Kt
441
+ ra[xx][yy] = a
442
+ n += 1
443
+ assert n == 1, n
444
+
445
+
446
+ def _add_resid(desc, resid, out, pc):
447
+ """MAST3R_OPT "resmm": fold ``out = ... + resid`` into a 2-D mcast linear's FUSE_BIAS epilogue. ``resid`` must be
448
+ BLOCK_SHARDED with exactly the output's shard spec (== the per-core output block); the compute kernel becomes the
449
+ model-local copy with MAST3R_RESID_CB and reads resid through a CB globally allocated on its shard."""
450
+ rmc, omc = resid.memory_config(), out.memory_config()
451
+ assert rmc.is_sharded() and omc.is_sharded() and rmc == omc, (rmc, omc)
452
+ assert pc.out_block_h == pc.per_core_M and pc.out_block_w == pc.per_core_N
453
+ assert pc.out_subblock_h == 1 or pc.out_subblock_w == pc.per_core_N, "row-major packing of the sharded output"
454
+ assert [int(d) for d in resid.shape] == [int(d) for d in out.shape] and resid.dtype == ttnn.bfloat16
455
+ used = set()
456
+ for c in desc.cbs:
457
+ for f in c.format_descriptors:
458
+ used.add(int(f.buffer_index))
459
+ idx = next(i for i in range(8, 32) if i not in used)
460
+ desc.cbs = list(desc.cbs) + [ttnn.cb_descriptor_from_sharded_tensor(idx, resid)]
461
+ n = 0
462
+ for k in desc.kernels:
463
+ src = str(k.kernel_source)
464
+ if src.endswith("bmm_large_block_zm_fused_bias_activation.cpp") or src == _RESMM_SRC:
465
+ assert any(d[0] == "FUSE_BIAS" for d in k.defines), "resmm needs a fused bias"
466
+ k.kernel_source = _RESMM_SRC
467
+ k.defines = list(k.defines) + [("MAST3R_RESID_CB", str(idx))]
468
+ n += 1
469
+ assert n == 1, n
470
+
471
+
472
+ def linear_descriptor(x, w, b, pc, ck, out_mc, compute_source=None, extra_defines=(), resid=None, in0_colmajor=False):
473
+ """ttnn.linear(x, w, bias=b, program_config=pc (2-D mcast), compute_kernel_config=ck) built from ttnn's own
474
+ MatmulMultiCoreReuseMcast2DProgramFactory.create_descriptor and run through generic_op. With
475
+ ``compute_source`` the compute kernel file is swapped for a model-local copy (same CBs / args / defines)."""
476
+ m = ttnn._ttnn.operations.matmul
477
+ p = m.MatmulParams()
478
+ p.program_config = pc
479
+ p.bcast_batch = True
480
+ p.output_mem_config = out_mc
481
+ p.output_dtype = ttnn.bfloat16
482
+ p.compute_kernel_config = ck
483
+ p.untilize_out = False
484
+ p.transpose_a = False
485
+ p.transpose_b = False
486
+ i = m.MatmulInputs()
487
+ dummy = _in0_dummy(x)
488
+ i.input_tensors = [dummy if dummy is not None else x, w]
489
+ i.optional_input_tensors = [b]
490
+ shp = [int(d) for d in x.shape][:-1] + [int(w.shape[-1])]
491
+ outs = [ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, x.device(), out_mc)]
492
+ i.optional_output_tensors = [outs[0]]
493
+ desc = m.MatmulMultiCoreReuseMcast2DProgramFactory.create_descriptor(p, i, outs)
494
+ if dummy is not None:
495
+ if in0_colmajor:
496
+ ps = [int(d) for d in x.padded_shape]
497
+ _patch_in0_colmajor(desc, dummy, x, math.prod(ps[:-1]) // 32, ps[-1] // 32)
498
+ else:
499
+ _patch_in0_ta(desc, dummy, x)
500
+ ttnn.deallocate(dummy)
501
+ else:
502
+ assert not in0_colmajor
503
+ if compute_source is not None:
504
+ n = 0
505
+ for k in desc.kernels:
506
+ if str(k.kernel_source).endswith("bmm_large_block_zm_fused_bias_activation.cpp"):
507
+ k.kernel_source = compute_source
508
+ if extra_defines:
509
+ k.defines = list(k.defines) + list(extra_defines)
510
+ n += 1
511
+ assert n == 1, n
512
+ if resid is not None:
513
+ _add_resid(desc, resid, outs[0], pc)
514
+ ttnn.generic_op([x, w, b] + ([resid] if resid is not None else []) + list(outs), desc)
515
+ return outs[0]
516
+
517
+
518
+ def phase_interleave_rm(ys, h0, w0, C, device, out_mc=None, pb=32):
519
+ """ROW_MAJOR variant of phase_interleave: ys are [1, 1, h0*w0, C] ROW_MAJOR bf16 interleaved (page = pixel row),
520
+ output [1, 1, 4*h0*w0, C] ROW_MAJOR (pixel-row pages). Exact."""
521
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
522
+ row = C * 2
523
+ lr = row.bit_length() - 1
524
+ assert 1 << lr == row and (2 * w0) % pb == 0
525
+ out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, 4 * h0 * w0, C]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, device,
526
+ out_mc)
527
+ hs = out_mc.is_sharded()
528
+ if hs:
529
+ # MAST3R_OPT "pilhs": write the HEIGHT_SHARDED spec directly (shard rows a multiple of the unit size pb)
530
+ spec = out_mc.shard_spec
531
+ shard_rows = int(spec.shape[0])
532
+ assert out_mc.memory_layout == ttnn.TensorMemoryLayout.HEIGHT_SHARDED and shard_rows % pb == 0 and int(spec.shape[1]) == C
533
+ assert spec.orientation == ttnn.ShardOrientation.ROW_MAJOR
534
+ coords = []
535
+ for rng in spec.grid.ranges():
536
+ for yy in range(rng.start.y, rng.end.y + 1):
537
+ for xx in range(rng.start.x, rng.end.x + 1):
538
+ coords.append((xx, yy))
539
+ coords.sort(key=lambda c: (c[1], c[0])) # row-major shard order
540
+ assert len(coords) * shard_rows >= 4 * h0 * w0
541
+ noc = []
542
+ for (xx, yy) in coords:
543
+ v = device.worker_core_from_logical_core(ttnn.CoreCoord(xx, yy))
544
+ noc += [v.x, v.y]
545
+ U = 4 * h0 * w0 // pb
546
+ grid = device.compute_with_storage_grid_size()
547
+ gx, gy = grid.x, grid.y
548
+ n = gx * gy
549
+ base, extra = divmod(U, n)
550
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
551
+ rd_rt, wr_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
552
+ start = 0
553
+ for i in range(n):
554
+ x, y = i % gx, i // gx
555
+ cnt = base + (1 if i < extra else 0)
556
+ rd_rt[x][y] = [t.buffer_address() for t in ys] + [start, cnt]
557
+ wr_rt[x][y] = [start, cnt] if hs else [out.buffer_address(), start, cnt]
558
+ start += cnt
559
+ rd = ttnn.KernelDescriptor(
560
+ kernel_source=os.path.join(_KDIR, "phase_il_rm_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
561
+ core_ranges=cores, compile_time_args=[w0, pb, lr], runtime_args=rd_rt,
562
+ defines=[("SRC_DRAM", "true" if _is_dram(ys[0]) else "false")], config=ttnn.ReaderConfigDescriptor())
563
+ if hs:
564
+ wr = ttnn.KernelDescriptor(
565
+ kernel_source=os.path.join(_KDIR, "phase_il_rm_writer_hs.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
566
+ core_ranges=cores, compile_time_args=[pb, lr, shard_rows // pb], runtime_args=wr_rt,
567
+ common_runtime_args=[out.buffer_address()] + noc, config=ttnn.WriterConfigDescriptor())
568
+ else:
569
+ wr = ttnn.KernelDescriptor(
570
+ kernel_source=os.path.join(_KDIR, "phase_il_rm_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
571
+ core_ranges=cores, compile_time_args=[pb, lr], runtime_args=wr_rt,
572
+ defines=[("OUT_DRAM", "true" if _is_dram(out) else "false")], config=ttnn.WriterConfigDescriptor())
573
+ cb = ttnn.CBDescriptor(total_size=2 * pb * row, core_ranges=cores,
574
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=0, data_format=ttnn.bfloat16,
575
+ page_size=pb * row)])
576
+ prog = ttnn.ProgramDescriptor(kernels=[rd, wr], semaphores=[], cbs=[cb])
577
+ ttnn.generic_op(list(ys) + [out], prog)
578
+ return out
579
+
580
+
581
+ def tail_interleave(t_rm, hl, wl, device, out_mc=None):
582
+ """Head-tail polyphase interleave: t_rm [1, 1, 16, hl*wl] ROW_MAJOR bf16 (row (2a+b)*4 + c, col i*wl + j) ->
583
+ [1, 1, 4*hl, 4*wl] ROW_MAJOR with page (c, i) = hi-res rows 2i, 2i+1 of channel c (NCHW [4, 2hl, 2wl]).
584
+ Model-local data-movement kernel; exact."""
585
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
586
+ trow = hl * wl * 2
587
+ lt = trow.bit_length() - 1
588
+ page = 4 * wl * 2
589
+ lp = page.bit_length() - 1
590
+ assert 1 << lt == trow and 1 << lp == page
591
+ out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, 4 * hl, 4 * wl]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, device,
592
+ out_mc)
593
+ U = 4 * hl
594
+ grid = device.compute_with_storage_grid_size()
595
+ gx, gy = grid.x, grid.y
596
+ n = gx * gy
597
+ base, extra = divmod(U, n)
598
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
599
+ rd_rt, wr_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
600
+ start = 0
601
+ for i in range(n):
602
+ x, y = i % gx, i // gx
603
+ cnt = base + (1 if i < extra else 0)
604
+ rd_rt[x][y] = [t_rm.buffer_address(), start, cnt]
605
+ wr_rt[x][y] = [out.buffer_address(), start, cnt]
606
+ start += cnt
607
+ rd = ttnn.KernelDescriptor(
608
+ kernel_source=os.path.join(_KDIR, "tail_il_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
609
+ core_ranges=cores, compile_time_args=[hl, wl, lt], runtime_args=rd_rt,
610
+ defines=[("SRC_DRAM", "true" if _is_dram(t_rm) else "false")], config=ttnn.ReaderConfigDescriptor())
611
+ wr = ttnn.KernelDescriptor(
612
+ kernel_source=os.path.join(_KDIR, "tail_il_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
613
+ core_ranges=cores, compile_time_args=[wl, lp], runtime_args=wr_rt,
614
+ defines=[("OUT_DRAM", "true" if _is_dram(out) else "false")], config=ttnn.WriterConfigDescriptor())
615
+ cbs = [ttnn.CBDescriptor(total_size=2 * 4 * wl * 2, core_ranges=cores,
616
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=0, data_format=ttnn.bfloat16,
617
+ page_size=4 * wl * 2)]),
618
+ ttnn.CBDescriptor(total_size=2 * page, core_ranges=cores,
619
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=1, data_format=ttnn.bfloat16,
620
+ page_size=page)])]
621
+ prog = ttnn.ProgramDescriptor(kernels=[rd, wr], semaphores=[], cbs=cbs)
622
+ ttnn.generic_op([t_rm, out], prog)
623
+ return out
624
+
625
+
626
+ def _mm_params(pc, ck, out_mc):
627
+ m = ttnn._ttnn.operations.matmul
628
+ p = m.MatmulParams()
629
+ p.program_config = pc
630
+ p.bcast_batch = True
631
+ p.output_mem_config = out_mc
632
+ p.output_dtype = ttnn.bfloat16
633
+ p.compute_kernel_config = ck
634
+ p.untilize_out = False
635
+ p.transpose_a = False
636
+ p.transpose_b = False
637
+ return p
638
+
639
+
640
+ def _swap_compute(desc, compute_source):
641
+ n = 0
642
+ for k in desc.kernels:
643
+ if str(k.kernel_source).endswith("bmm_large_block_zm_fused_bias_activation.cpp"):
644
+ k.kernel_source = compute_source
645
+ n += 1
646
+ assert n == 1, n
647
+
648
+
649
+ def dual_linear(jobs, ck, out_mc=None, compute_source=None, resids=None):
650
+ """Several independent ``x_i @ w_i + b_i`` (2-D mcast) in ONE program on disjoint core blocks:
651
+ ``jobs`` = [(x, w, b, pc)] where each pc is a MatmulMultiCoreReuseMultiCastProgramConfig with its own
652
+ ``allowed_worker_cores`` block (ttnn's factory places the mcast grid there). The per-job descriptors come from
653
+ ttnn's MatmulMultiCoreReuseMcast2DProgramFactory.create_descriptor and are combined with
654
+ ttnn.merge_program_descriptors; each job computes exactly what ttnn.linear with that pc computes."""
655
+ m = ttnn._ttnn.operations.matmul
656
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
657
+ descs, outs, ios = [], [], []
658
+ mcs = out_mc if isinstance(out_mc, (list, tuple)) else [out_mc] * len(jobs) # per-job (MAST3R_OPT "obsd")
659
+ for (x, w, b, pc), out_mc in zip(jobs, mcs):
660
+ p = _mm_params(pc, ck, out_mc)
661
+ i = m.MatmulInputs()
662
+ dummy = _in0_dummy(x)
663
+ i.input_tensors = [dummy if dummy is not None else x, w]
664
+ i.optional_input_tensors = [b]
665
+ shp = [int(d) for d in x.shape][:-1] + [int(w.shape[-1])]
666
+ o = ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, x.device(), out_mc)
667
+ i.optional_output_tensors = [o]
668
+ d = m.MatmulMultiCoreReuseMcast2DProgramFactory.create_descriptor(p, i, [o])
669
+ if dummy is not None:
670
+ _patch_in0_ta(d, dummy, x)
671
+ ttnn.deallocate(dummy)
672
+ if compute_source is not None:
673
+ _swap_compute(d, compute_source)
674
+ r = resids[len(outs)] if resids is not None else None
675
+ if r is not None:
676
+ _add_resid(d, r, o, pc)
677
+ descs.append(d)
678
+ outs.append(o)
679
+ ios += [x, w, b] + ([r] if r is not None else [])
680
+ prog = ttnn.merge_program_descriptors(descs)
681
+ ttnn.generic_op(ios + outs, prog)
682
+ return outs
683
+
684
+
685
+ def add3_relu(x, skip, c2, device, out_mc=None, fp32_dest=False):
686
+ """DPT refinenet plumbing: s = x + (skip + c2) (two bf16-rounded FPU adds, the order of ttnn.add(x, ttnn.add(skip,
687
+ c2))) and relu(s) in one program (model-local kernel). All [1, 1, N, C] TILE bf16 interleaved. Returns (s, relu(s))."""
688
+ out_mc = out_mc or ttnn.DRAM_MEMORY_CONFIG
689
+ shp = [int(d) for d in x.shape]
690
+ s_ = ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
691
+ r_ = ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
692
+ U = 1
693
+ for d in shp:
694
+ U *= d
695
+ U //= 1024
696
+ grid = device.compute_with_storage_grid_size()
697
+ gx, gy = grid.x, grid.y
698
+ n = gx * gy
699
+ base, extra = divmod(U, n)
700
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
701
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
702
+ start = 0
703
+ for i in range(n):
704
+ xx, yy = i % gx, i // gx
705
+ cnt = base + (1 if i < extra else 0)
706
+ rd_rt[xx][yy] = [start, cnt]
707
+ wr_rt[xx][yy] = [start, cnt]
708
+ cp_rt[xx][yy] = [cnt]
709
+ start += cnt
710
+ rd = ttnn.KernelDescriptor(
711
+ kernel_source=os.path.join(_KDIR, "add3_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
712
+ core_ranges=cores, compile_time_args=[], runtime_args=rd_rt,
713
+ common_runtime_args=[x.buffer_address(), skip.buffer_address(), c2.buffer_address()],
714
+ defines=[("IN_DRAM", "true" if _is_dram(x) else "false")], config=ttnn.ReaderConfigDescriptor())
715
+ wr = ttnn.KernelDescriptor(
716
+ kernel_source=os.path.join(_KDIR, "add3_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
717
+ core_ranges=cores, compile_time_args=[], runtime_args=wr_rt,
718
+ common_runtime_args=[s_.buffer_address(), r_.buffer_address()],
719
+ defines=[("OUT_DRAM", "true" if _is_dram(s_) else "false")], config=ttnn.WriterConfigDescriptor())
720
+ cp = ttnn.KernelDescriptor(
721
+ kernel_source=os.path.join(_KDIR, "add3_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
722
+ core_ranges=cores, compile_time_args=[], runtime_args=cp_rt,
723
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
724
+ fp32_dest_acc_en=fp32_dest))
725
+ cbs = [_cb(0, 2, cores), _cb(1, 2, cores), _cb(2, 2, cores), _cb(24, 2, cores), _cb(16, 2, cores), _cb(17, 2, cores)]
726
+ ttnn.generic_op([x, skip, c2, s_, r_], ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs))
727
+ return s_, r_
728
+
729
+
730
+ def _hs_ranges(t, device):
731
+ """Per-core tile ranges of a HEIGHT_SHARDED TILE tensor whose shard grid is the first n cores of the compute grid in
732
+ row-major order (ttnn.num_cores_to_corerangeset(n, grid, row_wise=True), ShardOrientation.ROW_MAJOR). Returns
733
+ (CoreRangeSet, [(x, y, start, count)]) or None if the spec is not of that form."""
734
+ mc = t.memory_config()
735
+ if mc.memory_layout != ttnn.TensorMemoryLayout.HEIGHT_SHARDED or t.layout != ttnn.TILE_LAYOUT:
736
+ return None
737
+ ss = mc.shard_spec
738
+ if ss is None or ss.orientation != ttnn.ShardOrientation.ROW_MAJOR:
739
+ return None
740
+ g = device.compute_with_storage_grid_size()
741
+ n = ss.grid.num_cores()
742
+ if ss.grid != ttnn.num_cores_to_corerangeset(n, g, row_wise=True):
743
+ return None
744
+ sh, sw = ss.shape
745
+ C = int(t.shape[-1])
746
+ if sw != C or sh % 32 or C % 32:
747
+ return None
748
+ shp = [int(d) for d in t.shape]
749
+ rows = 1
750
+ for d in shp[:-1]:
751
+ rows *= d
752
+ U = rows * C // 1024
753
+ per = (sh // 32) * (C // 32)
754
+ out = []
755
+ for k in range(n):
756
+ s = k * per
757
+ out.append((k % g.x, k // g.x, s, max(0, min(U, s + per) - s)))
758
+ return ss.grid, out
759
+
760
+
761
+ def add3_relu_s(x, skip, c2, device, out_mc, r_mc=None):
762
+ """add3_relu with c2 HEIGHT_SHARDED (read from each core's own shard: no S2I) and optionally relu(s) written
763
+ straight into a height-sharded tensor with c2's spec (r_mc; the next conv's input: no I2S). x / skip interleaved
764
+ (same buffer type), s interleaved (out_mc). Same compute kernel as add3_relu -> bit-identical."""
765
+ hr = _hs_ranges(c2, device)
766
+ assert hr is not None
767
+ cores, rng = hr
768
+ shp = [int(d) for d in x.shape]
769
+ s_ = ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
770
+ r_shard = r_mc is not None and r_mc.is_sharded()
771
+ if r_shard:
772
+ assert r_mc == c2.memory_config()
773
+ r_ = ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, r_mc or out_mc)
774
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
775
+ for xx, yy, st, cnt in rng:
776
+ rd_rt[xx][yy] = [st, cnt]
777
+ wr_rt[xx][yy] = [st, cnt]
778
+ cp_rt[xx][yy] = [cnt]
779
+ rd = ttnn.KernelDescriptor(
780
+ kernel_source=os.path.join(_KDIR, "add3s_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
781
+ core_ranges=cores, compile_time_args=[], runtime_args=rd_rt,
782
+ common_runtime_args=[x.buffer_address(), skip.buffer_address(), c2.buffer_address()],
783
+ defines=[("IN_DRAM", "true" if _is_dram(x) else "false"), ("C_SHARD", "1"), ("C_DRAM", "false"), ("N_IN", "3")],
784
+ config=ttnn.ReaderConfigDescriptor())
785
+ wr = ttnn.KernelDescriptor(
786
+ kernel_source=os.path.join(_KDIR, "add3s_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
787
+ core_ranges=cores, compile_time_args=[], runtime_args=wr_rt,
788
+ common_runtime_args=[s_.buffer_address(), r_.buffer_address()],
789
+ defines=[("OUT_DRAM", "true" if _is_dram(s_) else "false"), ("N_OUT", "2"), ("R_SHARD", "1" if r_shard else "0"),
790
+ ("R_DRAM", "true" if _is_dram(r_) else "false")], config=ttnn.WriterConfigDescriptor())
791
+ cp = ttnn.KernelDescriptor(
792
+ kernel_source=os.path.join(_KDIR, "add3_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
793
+ core_ranges=cores, compile_time_args=[], runtime_args=cp_rt,
794
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
795
+ fp32_dest_acc_en=False))
796
+ cbs = [_cb(0, 2, cores), _cb(1, 2, cores), _cb(2, 2, cores), _cb(24, 2, cores), _cb(16, 2, cores), _cb(17, 2, cores)]
797
+ ttnn.generic_op([x, skip, c2, s_, r_], ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs))
798
+ return s_, r_
799
+
800
+
801
+ def add2_s(a, c2, device, out_mc):
802
+ """a + c2 (ttnn.add's single FPU add, bit-identical) with c2 HEIGHT_SHARDED read from each core's own shard (no
803
+ S2I); a interleaved, output interleaved (out_mc)."""
804
+ hr = _hs_ranges(c2, device)
805
+ assert hr is not None
806
+ cores, rng = hr
807
+ shp = [int(d) for d in a.shape]
808
+ o_ = ttnn.allocate_tensor_on_device(ttnn.Shape(shp), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
809
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
810
+ for xx, yy, st, cnt in rng:
811
+ rd_rt[xx][yy] = [st, cnt]
812
+ wr_rt[xx][yy] = [st, cnt]
813
+ cp_rt[xx][yy] = [cnt]
814
+ rd = ttnn.KernelDescriptor(
815
+ kernel_source=os.path.join(_KDIR, "add3s_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
816
+ core_ranges=cores, compile_time_args=[], runtime_args=rd_rt,
817
+ common_runtime_args=[a.buffer_address(), a.buffer_address(), c2.buffer_address()],
818
+ defines=[("IN_DRAM", "true" if _is_dram(a) else "false"), ("C_SHARD", "1"), ("C_DRAM", "false"), ("N_IN", "2")],
819
+ config=ttnn.ReaderConfigDescriptor())
820
+ wr = ttnn.KernelDescriptor(
821
+ kernel_source=os.path.join(_KDIR, "add3s_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
822
+ core_ranges=cores, compile_time_args=[], runtime_args=wr_rt,
823
+ common_runtime_args=[o_.buffer_address(), o_.buffer_address()],
824
+ defines=[("OUT_DRAM", "true" if _is_dram(o_) else "false"), ("N_OUT", "1"), ("R_SHARD", "0"), ("R_DRAM", "false")],
825
+ config=ttnn.WriterConfigDescriptor())
826
+ cp = ttnn.KernelDescriptor(
827
+ kernel_source=os.path.join(_KDIR, "add2_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
828
+ core_ranges=cores, compile_time_args=[], runtime_args=cp_rt,
829
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
830
+ fp32_dest_acc_en=False))
831
+ cbs = [_cb(0, 2, cores), _cb(2, 2, cores), _cb(16, 2, cores)]
832
+ ttnn.generic_op([a, c2, o_], ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs))
833
+ return o_
834
+
835
+
836
+ def tail_fused(zs, bias_full, hl, wl, device, out_mc=None):
837
+ """Fused head tail (MAST3R_OPT "tailf"): zs = the 4 per-phase 1x1 outputs [1, 1, hl*wl, 32] TILE (phase p nonzero
838
+ only in columns 4p..4p+3), bias_full = [1, 1, 32, 32] TILE with every row = the 16 head.4 biases (+ zeros).
839
+ Returns [1, 1, 4*hl, 4*wl] ROW_MAJOR NCHW rows (= tail_interleave(untilize(slice(transpose(z0+z1+z2+z3) + bcol))),
840
+ bit-identical): one program instead of 3 adds + transpose + bias add + slice + untilize + interleave."""
841
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
842
+ assert wl % 64 == 0
843
+ tph = wl // 64 # tiles per unit (unit = half a low-res image row)
844
+ page = 4 * wl * 2
845
+ lp = page.bit_length() - 1
846
+ assert 1 << lp == page
847
+ out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, 4 * hl, 4 * wl]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, device,
848
+ out_mc)
849
+ grid = device.compute_with_storage_grid_size()
850
+ gx, gy = grid.x, grid.y
851
+ n = gx * gy
852
+ base, extra = divmod(2 * hl, n)
853
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
854
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
855
+ start = 0
856
+ for i in range(n):
857
+ x, y = i % gx, i // gx
858
+ cnt = base + (1 if i < extra else 0)
859
+ rd_rt[x][y] = [start, cnt]
860
+ wr_rt[x][y] = [start, cnt, out.buffer_address()]
861
+ cp_rt[x][y] = [cnt * tph]
862
+ start += cnt
863
+ rd = ttnn.KernelDescriptor(
864
+ kernel_source=os.path.join(_KDIR, "tailf_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
865
+ core_ranges=cores, compile_time_args=[tph], runtime_args=rd_rt,
866
+ common_runtime_args=[z.buffer_address() for z in zs] + [bias_full.buffer_address()],
867
+ defines=[("Z_DRAM", "true" if _is_dram(zs[0]) else "false"), ("B_DRAM", "true" if _is_dram(bias_full) else "false")],
868
+ config=ttnn.ReaderConfigDescriptor())
869
+ wr = ttnn.KernelDescriptor(
870
+ kernel_source=os.path.join(_KDIR, "tailf_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
871
+ core_ranges=cores, compile_time_args=[tph, wl, hl, lp], runtime_args=wr_rt,
872
+ defines=[("OUT_DRAM", "true" if _is_dram(out) else "false")], config=ttnn.WriterConfigDescriptor())
873
+ cp = ttnn.KernelDescriptor(
874
+ kernel_source=os.path.join(_KDIR, "tailf_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
875
+ core_ranges=cores, compile_time_args=[tph], runtime_args=cp_rt,
876
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
877
+ fp32_dest_acc_en=False))
878
+ cbs = [_cb(0, 2, cores), _cb(1, 2, cores), _cb(2, 2, cores), _cb(3, 2, cores), _cb(4, 1, cores),
879
+ _cb(24, 1, cores), _cb(25, 1, cores), _cb(26, 1, cores), _cb(27, 1, cores), _cb(16, 2, cores),
880
+ ttnn.CBDescriptor(total_size=2 * 8 * tph * 128, core_ranges=cores,
881
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=8, data_format=ttnn.bfloat16,
882
+ page_size=page)])]
883
+ ttnn.generic_op(list(zs) + [bias_full, out], ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs))
884
+ return out
885
+
886
+
887
+ _IN0_SENDER = "reader_bmm_tile_layout_in0_sender_padding.cpp"
888
+
889
+
890
+ def _patch_in0_heads(desc, dummy_addr, ctx, heads, nt, dt, b0):
891
+ """MAST3R_OPT "pcat": swap the 2-D mcast matmul's in0 sender reader for the model-local remapping copy and point its
892
+ in0 address (runtime arg 0) from the shape-only dummy at the SDPA output ``ctx``."""
893
+ n = 0
894
+ for k in desc.kernels:
895
+ if str(k.kernel_source).endswith(_IN0_SENDER):
896
+ k.kernel_source = os.path.join(_KDIR, "mm_in0_heads_reader.cpp")
897
+ k.defines = list(k.defines) + [("MAST3R_HEADS", str(heads)), ("MAST3R_NT", str(nt)), ("MAST3R_DT", str(dt)),
898
+ ("MAST3R_B0", str(b0))]
899
+ ra = k.runtime_args
900
+ for cr in k.core_ranges.ranges():
901
+ for x in range(cr.start.x, cr.end.x + 1):
902
+ for y in range(cr.start.y, cr.end.y + 1):
903
+ a = list(ra[x][y])
904
+ if a:
905
+ assert a[0] == dummy_addr, (a[0], dummy_addr)
906
+ a[0] = ctx.buffer_address()
907
+ ra[x][y] = a
908
+ n += 1
909
+ assert n == 1, n
910
+
911
+
912
+ def linear_from_heads(jobs, ck, out_mc=None, resids=None):
913
+ """MAST3R_OPT "pcat": x_i @ w_i + b_i where x_i = concat_heads(ctx_i)[branch b0_i] is read straight from the SDPA
914
+ output ctx_i [B, H, N, dh] (TILE interleaved) by a remapping in0 reader -- no concat-heads program. jobs =
915
+ [(ctx, b0, nb, w, b, pc)]: rows of batches b0 .. b0 + nb - 1; several jobs (disjoint allowed_worker_cores) are
916
+ merged into one program as in dual_linear. Bit-identical to heads_concat + ttnn.linear with the same pc."""
917
+ m = ttnn._ttnn.operations.matmul
918
+ out_mc = out_mc or ttnn.L1_MEMORY_CONFIG
919
+ descs, outs, ios, dummies = [], [], [], []
920
+ mcs = out_mc if isinstance(out_mc, (list, tuple)) else [out_mc] * len(jobs) # per-job (MAST3R_OPT "obsd")
921
+ for (ctx, b0, nb, w, b, pc), out_mc in zip(jobs, mcs):
922
+ Bc, H, N, dh = (int(d) for d in ctx.shape)
923
+ K = H * dh
924
+ dev = ctx.device()
925
+ dummy = ttnn.allocate_tensor_on_device(ttnn.Shape([nb, N, K]), ttnn.bfloat16, ttnn.TILE_LAYOUT, dev,
926
+ ctx.memory_config())
927
+ p = _mm_params(pc, ck, out_mc)
928
+ i = m.MatmulInputs()
929
+ i.input_tensors = [dummy, w]
930
+ i.optional_input_tensors = [b]
931
+ o = ttnn.allocate_tensor_on_device(ttnn.Shape([nb, N, int(w.shape[-1])]), ttnn.bfloat16, ttnn.TILE_LAYOUT, dev,
932
+ out_mc)
933
+ i.optional_output_tensors = [o]
934
+ d = m.MatmulMultiCoreReuseMcast2DProgramFactory.create_descriptor(p, i, [o])
935
+ _patch_in0_heads(d, dummy.buffer_address(), ctx, H, N // 32, dh // 32, b0)
936
+ r = resids[len(outs)] if resids is not None else None
937
+ if r is not None:
938
+ _add_resid(d, r, o, pc)
939
+ descs.append(d)
940
+ outs.append(o)
941
+ ios += [ctx, w, b] + ([r] if r is not None else [])
942
+ dummies.append(dummy)
943
+ for t in dummies:
944
+ ttnn.deallocate(t)
945
+ prog = descs[0] if len(descs) == 1 else ttnn.merge_program_descriptors(descs)
946
+ ttnn.generic_op(ios + outs, prog)
947
+ return outs
948
+
949
+
950
+ def ups2_tables():
951
+ """MAST3R_OPT "tups": the 12 constant interpolation tiles [12 * 32, 32] (row block 6 * wi + kind; wi 0 -> vertical
952
+ weight 0.25, 1 -> 0.75; kind 0 E_EDGE, 1 E_L, 2 E_R, 3 O_L, 4 O_R, 5 O_EDGE). Half-pixel bilinear x2 with clamping
953
+ (== ttnn.upsample): out X = 2k -> 0.25 x[k-1] + 0.75 x[k], X = 2k+1 -> 0.75 x[k] + 0.25 x[k+1]; out tile t' = 2m covers
954
+ the in-tile m (E_R) and pixel 31 of in-tile m-1 (E_L); t' = 2m+1 in-tile m (O_L) and pixel 0 of in-tile m+1 (O_R); the
955
+ image edges clamp (E_EDGE / O_EDGE). Every entry is exact in bf16."""
956
+ import torch
957
+ M = torch.zeros(6, 32, 32, dtype=torch.float64)
958
+ for i in range(32):
959
+ k = i // 2 # even t': in-tile local index of x[k]
960
+ if i % 2 == 0:
961
+ if k - 1 >= 0:
962
+ M[2, i, k - 1] += 0.25
963
+ else:
964
+ M[1, i, 31] += 0.25 # x[k-1] = pixel 31 of the previous in-tile
965
+ M[2, i, k] += 0.75
966
+ else:
967
+ M[2, i, k] += 0.75
968
+ M[2, i, k + 1] += 0.25
969
+ k = 16 + i // 2 # odd t'
970
+ if i % 2 == 0:
971
+ M[3, i, k - 1] += 0.25
972
+ M[3, i, k] += 0.75
973
+ else:
974
+ M[3, i, k] += 0.75
975
+ if k + 1 <= 31:
976
+ M[3, i, k + 1] += 0.25
977
+ else:
978
+ M[4, i, 0] += 0.25 # x[k+1] = pixel 0 of the next in-tile
979
+ M[0] = M[2]
980
+ M[0, 0, 0] += 0.25 # left edge: x[-1] clamps to x[0]
981
+ M[5] = M[3]
982
+ M[5, 31, 31] += 0.25 # right edge: x[W] clamps to x[W-1]
983
+ A = torch.cat([M * 0.25, M * 0.75], 0).reshape(12 * 32, 32)
984
+ return A
985
+
986
+
987
+ def upsample2x_tile(x, H, W, C, device, a_tiles, out_mc=None):
988
+ """MAST3R_OPT "tups": bilinear x2 upsample (half-pixel, clamped == ttnn.upsample) of x [1, 1, H*W, C] TILE bf16
989
+ interleaved (pixel rows, row-major image), straight to [1, 1, 4*H*W, C] TILE (no untilize / reshard / halo / tilize).
990
+ W must be a multiple of 32. Exact products and fp32 sums (see ups2_compute.cpp)."""
991
+ out_mc = out_mc or ttnn.DRAM_MEMORY_CONFIG
992
+ assert W % 32 == 0 and C % 32 == 0
993
+ NX, CT = W // 32, C // 32
994
+ out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, 4 * H * W, C]), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
995
+ U = 2 * H * CT
996
+ grid = device.compute_with_storage_grid_size()
997
+ gx, gy = grid.x, grid.y
998
+ n = gx * gy
999
+ base, extra = divmod(U, n)
1000
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
1001
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
1002
+ start = 0
1003
+ for i in range(n):
1004
+ xx, yy = i % gx, i // gx
1005
+ cnt = base + (1 if i < extra else 0)
1006
+ rd_rt[xx][yy] = [start, cnt]
1007
+ wr_rt[xx][yy] = [start, cnt]
1008
+ cp_rt[xx][yy] = [start, cnt]
1009
+ start += cnt
1010
+ assert start == U
1011
+ rd = ttnn.KernelDescriptor(
1012
+ kernel_source=os.path.join(_KDIR, "ups2_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1013
+ core_ranges=cores, compile_time_args=[NX, CT, H], runtime_args=rd_rt,
1014
+ common_runtime_args=[x.buffer_address(), a_tiles.buffer_address()],
1015
+ defines=[("IN_DRAM", "true" if _is_dram(x) else "false"), ("A_DRAM", "true" if _is_dram(a_tiles) else "false")],
1016
+ config=ttnn.ReaderConfigDescriptor())
1017
+ wr = ttnn.KernelDescriptor(
1018
+ kernel_source=os.path.join(_KDIR, "ups2_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1019
+ core_ranges=cores, compile_time_args=[NX, CT], runtime_args=wr_rt, common_runtime_args=[out.buffer_address()],
1020
+ defines=[("OUT_DRAM", "true" if _is_dram(out) else "false")], config=ttnn.WriterConfigDescriptor())
1021
+ cp = ttnn.KernelDescriptor(
1022
+ kernel_source=os.path.join(_KDIR, "ups2_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1023
+ core_ranges=cores, compile_time_args=[NX, CT], runtime_args=cp_rt,
1024
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
1025
+ fp32_dest_acc_en=True))
1026
+ cbs = [_cb(0, 4 * NX, cores), _cb(1, 12, cores), _cb(16, 4, cores)]
1027
+ ttnn.generic_op([x, a_tiles, out], ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs))
1028
+ return out
1029
+
1030
+
1031
+ def strip_gather(xs, h0, w0, C, device):
1032
+ """MAST3R_OPT "sgat": the ring-strip inputs (tb [2, 3, w0, C], lr [2, 3, h0, C], ROW_MAJOR, DRAM) gathered straight from
1033
+ the HEIGHT_SHARDED TILE tensor xs [1, 1, h0*w0, C] (see kernels/strip_gather.cpp). Exact."""
1034
+ mc = xs.memory_config()
1035
+ spec = mc.shard_spec
1036
+ assert mc.memory_layout == ttnn.TensorMemoryLayout.HEIGHT_SHARDED and spec.orientation == ttnn.ShardOrientation.ROW_MAJOR
1037
+ assert int(spec.shape[0]) % 32 == 0 and int(spec.shape[1]) == C and C % 32 == 0
1038
+ PER, CT = int(spec.shape[0]) // 32, C // 32
1039
+ coords = []
1040
+ for rng in spec.grid.ranges():
1041
+ for yy in range(rng.start.y, rng.end.y + 1):
1042
+ for xx in range(rng.start.x, rng.end.x + 1):
1043
+ coords.append((xx, yy))
1044
+ coords.sort(key=lambda c: (c[1], c[0]))
1045
+ noc = []
1046
+ for (xx, yy) in coords:
1047
+ v = device.worker_core_from_logical_core(ttnn.CoreCoord(xx, yy))
1048
+ noc += [v.x, v.y]
1049
+ tb = ttnn.allocate_tensor_on_device(ttnn.Shape([2, 3, w0, C]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, device,
1050
+ ttnn.DRAM_MEMORY_CONFIG)
1051
+ lr = ttnn.allocate_tensor_on_device(ttnn.Shape([2, 3, h0, C]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, device,
1052
+ ttnn.DRAM_MEMORY_CONFIG)
1053
+ P = 6 * (w0 + h0)
1054
+ grid = device.compute_with_storage_grid_size()
1055
+ gx, gy = grid.x, grid.y
1056
+ n = gx * gy
1057
+ base, extra = divmod(P, n)
1058
+ maxp = base + (1 if extra else 0)
1059
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
1060
+ rt = ttnn.RuntimeArgs()
1061
+ start = 0
1062
+ for i in range(n):
1063
+ xx, yy = i % gx, i // gx
1064
+ cnt = base + (1 if i < extra else 0)
1065
+ rt[xx][yy] = [start, cnt]
1066
+ start += cnt
1067
+ assert start == P
1068
+ k = ttnn.KernelDescriptor(
1069
+ kernel_source=os.path.join(_KDIR, "strip_gather.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1070
+ core_ranges=cores, compile_time_args=[CT, PER, w0, h0], runtime_args=rt,
1071
+ common_runtime_args=[xs.buffer_address(), tb.buffer_address(), lr.buffer_address()] + noc,
1072
+ config=ttnn.ReaderConfigDescriptor())
1073
+ page = CT * 64
1074
+ cb = ttnn.CBDescriptor(total_size=max(1, maxp) * page, core_ranges=cores,
1075
+ format_descriptors=[ttnn.CBFormatDescriptor(buffer_index=0, data_format=ttnn.bfloat16, page_size=page)])
1076
+ ttnn.generic_op([xs, tb, lr], ttnn.ProgramDescriptor(kernels=[k], semaphores=[], cbs=[cb]))
1077
+ return tb, lr
1078
+
1079
+
1080
+ def ups2h_tables():
1081
+ """MAST3R_OPT "tups", W = 16: 4 constant tiles [4 * 32, 32], index 2 * wi + h (wi 0 -> vertical weight 0.25, 1 -> 0.75;
1082
+ h = which half of the input tile holds the source image row): out px i of the 32-px output row, k = i // 2:
1083
+ even i -> 0.25 x[max(k-1, 0)] + 0.75 x[k], odd i -> 0.75 x[k] + 0.25 x[min(k+1, 15)], x at in-tile rows 16 h + k."""
1084
+ import torch
1085
+ M = torch.zeros(2, 32, 32, dtype=torch.float64)
1086
+ for h in range(2):
1087
+ for i in range(32):
1088
+ k = i // 2
1089
+ if i % 2 == 0:
1090
+ M[h, i, 16 * h + max(k - 1, 0)] += 0.25
1091
+ M[h, i, 16 * h + k] += 0.75
1092
+ else:
1093
+ M[h, i, 16 * h + k] += 0.75
1094
+ M[h, i, 16 * h + min(k + 1, 15)] += 0.25
1095
+ return torch.cat([M * 0.25, M * 0.75], 0).reshape(4 * 32, 32)
1096
+
1097
+
1098
+ def upsample2x_tile_half(x, H, C, device, a_tiles, out_mc=None):
1099
+ """MAST3R_OPT "tups" for a 16-px-wide image (W = 16, two image rows per input tile row): [1, 1, 16*H, C] TILE ->
1100
+ [1, 1, 64*H, C] TILE (32-px output rows). Same arithmetic as upsample2x_tile."""
1101
+ out_mc = out_mc or ttnn.DRAM_MEMORY_CONFIG
1102
+ assert C % 32 == 0 and H % 2 == 0
1103
+ CT = C // 32
1104
+ out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, 64 * H, C]), ttnn.bfloat16, ttnn.TILE_LAYOUT, device, out_mc)
1105
+ U = 2 * H * CT
1106
+ grid = device.compute_with_storage_grid_size()
1107
+ gx, gy = grid.x, grid.y
1108
+ n = gx * gy
1109
+ base, extra = divmod(U, n)
1110
+ cores = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))])
1111
+ rd_rt, wr_rt, cp_rt = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
1112
+ start = 0
1113
+ for i in range(n):
1114
+ xx, yy = i % gx, i // gx
1115
+ cnt = base + (1 if i < extra else 0)
1116
+ rd_rt[xx][yy] = [start, cnt]
1117
+ wr_rt[xx][yy] = [start, cnt]
1118
+ cp_rt[xx][yy] = [start, cnt]
1119
+ start += cnt
1120
+ assert start == U
1121
+ rd = ttnn.KernelDescriptor(
1122
+ kernel_source=os.path.join(_KDIR, "ups2h_reader.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1123
+ core_ranges=cores, compile_time_args=[CT, H], runtime_args=rd_rt,
1124
+ common_runtime_args=[x.buffer_address(), a_tiles.buffer_address()],
1125
+ defines=[("IN_DRAM", "true" if _is_dram(x) else "false"), ("A_DRAM", "true" if _is_dram(a_tiles) else "false")],
1126
+ config=ttnn.ReaderConfigDescriptor())
1127
+ wr = ttnn.KernelDescriptor(
1128
+ kernel_source=os.path.join(_KDIR, "ups2h_writer.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1129
+ core_ranges=cores, compile_time_args=[], runtime_args=wr_rt, common_runtime_args=[out.buffer_address()],
1130
+ defines=[("OUT_DRAM", "true" if _is_dram(out) else "false")], config=ttnn.WriterConfigDescriptor())
1131
+ cp = ttnn.KernelDescriptor(
1132
+ kernel_source=os.path.join(_KDIR, "ups2h_compute.cpp"), source_type=ttnn.KernelDescriptor.SourceType.FILE_PATH,
1133
+ core_ranges=cores, compile_time_args=[CT, H], runtime_args=cp_rt,
1134
+ config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False,
1135
+ fp32_dest_acc_en=True))
1136
+ cbs = [_cb(0, 4, cores), _cb(1, 4, cores), _cb(16, 4, cores)]
1137
+ ttnn.generic_op([x, a_tiles, out], ttnn.ProgramDescriptor(kernels=[rd, wr, cp], semaphores=[], cbs=cbs))
1138
+ return out
code/models/demos/mast3r/tt/kernels/add2_compute.cpp ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): out = a + c (CB 0 + CB 2 -> CB 16), one FPU add with the CB formats /
3
+ // rounding of ttnn.add (same as add3_compute's adds).
4
+ #include <cstdint>
5
+ #include "api/compute/common.h"
6
+ #include "api/compute/eltwise_binary.h"
7
+ #include "api/compute/compute_kernel_hw_startup.h"
8
+ #include "api/compute/cb_api.h"
9
+ #include "api/compute/pack.h"
10
+
11
+ constexpr uint32_t CB_A = 0, CB_C = 2, CB_OUT = 16;
12
+
13
+ void kernel_main() {
14
+ const uint32_t u_count = get_arg_val<uint32_t>(0);
15
+ if (u_count == 0) {
16
+ return;
17
+ }
18
+ compute_kernel_hw_startup(CB_A, CB_C, CB_OUT);
19
+ add_init(CB_A, CB_C);
20
+ for (uint32_t u = 0; u < u_count; ++u) {
21
+ cb_wait_front(CB_A, 1);
22
+ cb_wait_front(CB_C, 1);
23
+ cb_reserve_back(CB_OUT, 1);
24
+ tile_regs_acquire();
25
+ add_tiles(CB_A, CB_C, 0, 0, 0);
26
+ tile_regs_commit();
27
+ tile_regs_wait();
28
+ pack_tile(0, CB_OUT);
29
+ tile_regs_release();
30
+ cb_push_back(CB_OUT, 1);
31
+ cb_pop_front(CB_A, 1);
32
+ cb_pop_front(CB_C, 1);
33
+ }
34
+ }
code/models/demos/mast3r/tt/kernels/add3_compute.cpp ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): t = skip + c2; sum = x + t; out_relu = relu(sum). FPU adds with the
3
+ // same CB formats / rounding as two ttnn.add calls, relu on the SFPU (exact); writes sum (CB 16) and relu(sum) (CB 17).
4
+ #include <cstdint>
5
+ #include "api/compute/common.h"
6
+ #include "api/compute/eltwise_binary.h"
7
+ #include "api/compute/eltwise_unary/eltwise_unary.h"
8
+ #include "api/compute/eltwise_unary/relu.h"
9
+ #include "api/compute/compute_kernel_hw_startup.h"
10
+ #include "api/compute/cb_api.h"
11
+ #include "api/compute/pack.h"
12
+
13
+ constexpr uint32_t CB_X = 0, CB_S = 1, CB_C = 2, CB_T = 24, CB_SUM = 16, CB_RELU = 17;
14
+
15
+ void kernel_main() {
16
+ const uint32_t u_count = get_arg_val<uint32_t>(0);
17
+ if (u_count == 0) {
18
+ return;
19
+ }
20
+ compute_kernel_hw_startup(CB_S, CB_C, CB_T);
21
+ for (uint32_t u = 0; u < u_count; ++u) {
22
+ cb_wait_front(CB_S, 1);
23
+ cb_wait_front(CB_C, 1);
24
+ cb_reserve_back(CB_T, 1);
25
+ add_init(CB_S, CB_C);
26
+ tile_regs_acquire();
27
+ add_tiles(CB_S, CB_C, 0, 0, 0);
28
+ tile_regs_commit();
29
+ tile_regs_wait();
30
+ pack_tile(0, CB_T);
31
+ tile_regs_release();
32
+ cb_push_back(CB_T, 1);
33
+ cb_pop_front(CB_S, 1);
34
+ cb_pop_front(CB_C, 1);
35
+
36
+ cb_wait_front(CB_X, 1);
37
+ cb_wait_front(CB_T, 1);
38
+ cb_reserve_back(CB_SUM, 1);
39
+ add_init(CB_X, CB_T);
40
+ tile_regs_acquire();
41
+ add_tiles(CB_X, CB_T, 0, 0, 0);
42
+ tile_regs_commit();
43
+ tile_regs_wait();
44
+ pack_tile(0, CB_SUM);
45
+ tile_regs_release();
46
+ cb_push_back(CB_SUM, 1);
47
+
48
+ cb_reserve_back(CB_RELU, 1);
49
+ add_init(CB_X, CB_T);
50
+ tile_regs_acquire();
51
+ add_tiles(CB_X, CB_T, 0, 0, 0);
52
+ relu_tile_init();
53
+ relu_tile(0);
54
+ tile_regs_commit();
55
+ tile_regs_wait();
56
+ pack_tile(0, CB_RELU);
57
+ tile_regs_release();
58
+ cb_push_back(CB_RELU, 1);
59
+ cb_pop_front(CB_X, 1);
60
+ cb_pop_front(CB_T, 1);
61
+ }
62
+ }
code/models/demos/mast3r/tt/kernels/add3_reader.cpp ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): DPT refinenet "x + (skip + c2)" + relu, reader.
3
+ // Reads tile u of x (CB 0), skip (CB 1), c2 (CB 2) for u in this core's range (all TILE bf16 interleaved, same shape).
4
+ #include <stdint.h>
5
+ #include "api/dataflow/dataflow_api.h"
6
+
7
+ void kernel_main() {
8
+ constexpr uint32_t LOG2_TILE = 11;
9
+ constexpr uint32_t TB = 1u << LOG2_TILE;
10
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
11
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
12
+ const InterleavedPow2AddrGen<IN_DRAM> gx = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
13
+ const InterleavedPow2AddrGen<IN_DRAM> gs = {.bank_base_address = get_common_arg_val<uint32_t>(1), .log_base_2_of_page_size = LOG2_TILE};
14
+ const InterleavedPow2AddrGen<IN_DRAM> gc = {.bank_base_address = get_common_arg_val<uint32_t>(2), .log_base_2_of_page_size = LOG2_TILE};
15
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
16
+ cb_reserve_back(0, 1);
17
+ cb_reserve_back(1, 1);
18
+ cb_reserve_back(2, 1);
19
+ noc_async_read(gx.get_noc_addr(u), get_write_ptr(0), TB);
20
+ noc_async_read(gs.get_noc_addr(u), get_write_ptr(1), TB);
21
+ noc_async_read(gc.get_noc_addr(u), get_write_ptr(2), TB);
22
+ noc_async_read_barrier();
23
+ cb_push_back(0, 1);
24
+ cb_push_back(1, 1);
25
+ cb_push_back(2, 1);
26
+ }
27
+ }
code/models/demos/mast3r/tt/kernels/add3_writer.cpp ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): writes sum (CB 16) and relu(sum) (CB 17) tile u.
3
+ #include <stdint.h>
4
+ #include "api/dataflow/dataflow_api.h"
5
+
6
+ void kernel_main() {
7
+ constexpr uint32_t LOG2_TILE = 11;
8
+ constexpr uint32_t TB = 1u << LOG2_TILE;
9
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
10
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
11
+ const InterleavedPow2AddrGen<OUT_DRAM> g0 = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
12
+ const InterleavedPow2AddrGen<OUT_DRAM> g1 = {.bank_base_address = get_common_arg_val<uint32_t>(1), .log_base_2_of_page_size = LOG2_TILE};
13
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
14
+ cb_wait_front(16, 1);
15
+ noc_async_write(get_read_ptr(16), g0.get_noc_addr(u), TB);
16
+ cb_wait_front(17, 1);
17
+ noc_async_write(get_read_ptr(17), g1.get_noc_addr(u), TB);
18
+ noc_async_writes_flushed();
19
+ cb_pop_front(16, 1);
20
+ cb_pop_front(17, 1);
21
+ }
22
+ noc_async_write_barrier();
23
+ }
code/models/demos/mast3r/tt/kernels/add3s_reader.cpp ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): DPT refinenet "x + (skip + c2)" + relu, reader, shard-aware.
3
+ // Reads tile u of x (CB 0), skip (CB 1) from interleaved buffers and c2 (CB 2) either interleaved or (C_SHARD) from this
4
+ // core's own height shard (the core's range [u_start, u_start + u_count) is exactly its shard's tile range).
5
+ // N_IN = 3: x, skip, c2 ; N_IN = 2: x, c2 only (CB 0 and CB 2).
6
+ #include <stdint.h>
7
+ #include "api/dataflow/dataflow_api.h"
8
+
9
+ void kernel_main() {
10
+ constexpr uint32_t LOG2_TILE = 11;
11
+ constexpr uint32_t TB = 1u << LOG2_TILE;
12
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
13
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
14
+ const InterleavedPow2AddrGen<IN_DRAM> gx = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
15
+ const InterleavedPow2AddrGen<IN_DRAM> gs = {.bank_base_address = get_common_arg_val<uint32_t>(1), .log_base_2_of_page_size = LOG2_TILE};
16
+ const uint32_t c_base = get_common_arg_val<uint32_t>(2);
17
+ #if !C_SHARD
18
+ const InterleavedPow2AddrGen<C_DRAM> gc = {.bank_base_address = c_base, .log_base_2_of_page_size = LOG2_TILE};
19
+ #endif
20
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
21
+ cb_reserve_back(0, 1);
22
+ noc_async_read(gx.get_noc_addr(u), get_write_ptr(0), TB);
23
+ #if N_IN == 3
24
+ cb_reserve_back(1, 1);
25
+ noc_async_read(gs.get_noc_addr(u), get_write_ptr(1), TB);
26
+ #endif
27
+ cb_reserve_back(2, 1);
28
+ #if C_SHARD
29
+ noc_async_read(get_noc_addr(c_base + ((u - u_start) << LOG2_TILE)), get_write_ptr(2), TB);
30
+ #else
31
+ noc_async_read(gc.get_noc_addr(u), get_write_ptr(2), TB);
32
+ #endif
33
+ noc_async_read_barrier();
34
+ cb_push_back(0, 1);
35
+ #if N_IN == 3
36
+ cb_push_back(1, 1);
37
+ #endif
38
+ cb_push_back(2, 1);
39
+ }
40
+ }
code/models/demos/mast3r/tt/kernels/add3s_writer.cpp ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): writer for add3s / add2s. N_OUT = 2: sum (CB 16, interleaved) and
3
+ // relu (CB 17; R_SHARD: into this core's own height shard of the output); N_OUT = 1: CB 16 only.
4
+ #include <stdint.h>
5
+ #include "api/dataflow/dataflow_api.h"
6
+
7
+ void kernel_main() {
8
+ constexpr uint32_t LOG2_TILE = 11;
9
+ constexpr uint32_t TB = 1u << LOG2_TILE;
10
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
11
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
12
+ const InterleavedPow2AddrGen<OUT_DRAM> g0 = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
13
+ #if N_OUT == 2
14
+ const uint32_t r_base = get_common_arg_val<uint32_t>(1);
15
+ #if !R_SHARD
16
+ const InterleavedPow2AddrGen<R_DRAM> g1 = {.bank_base_address = r_base, .log_base_2_of_page_size = LOG2_TILE};
17
+ #endif
18
+ #endif
19
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
20
+ cb_wait_front(16, 1);
21
+ noc_async_write(get_read_ptr(16), g0.get_noc_addr(u), TB);
22
+ #if N_OUT == 2
23
+ cb_wait_front(17, 1);
24
+ #if R_SHARD
25
+ noc_async_write(get_read_ptr(17), get_noc_addr(r_base + ((u - u_start) << LOG2_TILE)), TB);
26
+ #else
27
+ noc_async_write(get_read_ptr(17), g1.get_noc_addr(u), TB);
28
+ #endif
29
+ #endif
30
+ noc_async_writes_flushed();
31
+ cb_pop_front(16, 1);
32
+ #if N_OUT == 2
33
+ cb_pop_front(17, 1);
34
+ #endif
35
+ }
36
+ noc_async_write_barrier();
37
+ }
code/models/demos/mast3r/tt/kernels/fattn_reader.cpp ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // mast3r-p150 model-local SDPA reader (MAST3R_OPT "fattn"), used with ttnn's stock streaming SDPA compute kernel
4
+ // (sdpa.cpp, single K chunk = the whole key sequence) and its stock writer through ttnn.generic_op.
5
+ //
6
+ // Differences to the stock reader_interleaved.cpp:
7
+ // * the whole K^T / V of a (batch, head) is loaded ONCE per core and stays resident: cb_k / cb_v hold exactly one
8
+ // chunk (= the full sequence), so after the compute pops it, reserve_back returns the same L1 slot and the reader
9
+ // re-publishes it (push_back) without any NoC traffic while the head does not change;
10
+ // * no KV chain forwarding (no inter-core coupling), no masks / paging / sinks;
11
+ // * flat B*H*q_chunk scheduling with per-core [g_start, g_start + g_count) ranges chosen by the host (balanced to
12
+ // the tile row, not to the stock q-chunk count).
13
+ // The CB tile layouts are the stock ones: Q row-major [Sq_chunk_t x DHt], K transposed [DHt x Skt], V row-major
14
+ // [Skt x DHt].
15
+
16
+ #include <cstdint>
17
+ #include "api/dataflow/dataflow_api.h"
18
+
19
+ void kernel_main() {
20
+ constexpr uint32_t q_num_chunks = get_compile_time_arg_val(0);
21
+ constexpr uint32_t Sq_chunk_t = get_compile_time_arg_val(1);
22
+ constexpr uint32_t Sqt = get_compile_time_arg_val(2);
23
+ constexpr uint32_t Skt = get_compile_time_arg_val(3);
24
+ constexpr uint32_t DHt = get_compile_time_arg_val(4);
25
+ constexpr uint32_t cb_q = get_compile_time_arg_val(5);
26
+ constexpr uint32_t cb_k = get_compile_time_arg_val(6);
27
+ constexpr uint32_t cb_v = get_compile_time_arg_val(7);
28
+ constexpr uint32_t kv_heads_div = get_compile_time_arg_val(8); // q head index -> kv head index divisor (1)
29
+ constexpr auto q_args = TensorAccessorArgs<9>();
30
+ constexpr auto k_args = TensorAccessorArgs<q_args.next_compile_time_args_offset()>();
31
+ constexpr auto v_args = TensorAccessorArgs<k_args.next_compile_time_args_offset()>();
32
+
33
+ const uint32_t q_addr = get_arg_val<uint32_t>(0);
34
+ const uint32_t k_addr = get_arg_val<uint32_t>(1);
35
+ const uint32_t v_addr = get_arg_val<uint32_t>(2);
36
+ const uint32_t g_start = get_arg_val<uint32_t>(3);
37
+ const uint32_t g_count = get_arg_val<uint32_t>(4);
38
+
39
+ constexpr uint32_t tile_bytes = get_tile_size(cb_q);
40
+ const auto q_rd = TensorAccessor(q_args, q_addr, tile_bytes);
41
+ const auto k_rd = TensorAccessor(k_args, k_addr, tile_bytes);
42
+ const auto v_rd = TensorAccessor(v_args, v_addr, tile_bytes);
43
+
44
+ constexpr uint32_t q_tiles = Sq_chunk_t * DHt;
45
+ constexpr uint32_t kv_tiles = Skt * DHt;
46
+
47
+ // NoC read throttle: with ~120 cores reading L1-interleaved tiles at once, more than ~2 outstanding tile reads per
48
+ // core congests the NoC (tools_prof/r7_read_bw.py: 2 -> 955 GB/s aggregate, unthrottled -> 270 GB/s)
49
+ #ifndef BAR
50
+ #define BAR 2
51
+ #endif
52
+ uint32_t nb = 0;
53
+ uint32_t prev_kv = 0xFFFFFFFF;
54
+ for (uint32_t i = 0; i < g_count; ++i) {
55
+ const uint32_t g = g_start + i;
56
+ const uint32_t bh = g / q_num_chunks; // flat (batch, q head)
57
+ const uint32_t qc = g - bh * q_num_chunks;
58
+ const uint32_t kvh = bh / kv_heads_div;
59
+ #ifdef FATTN_DIAG_NOKV
60
+ const bool new_kv = false;
61
+ #elif defined(FATTN_DIAG_FIRSTKV)
62
+ const bool new_kv = i == 0;
63
+ #else
64
+ const bool new_kv = kvh != prev_kv;
65
+ #endif
66
+
67
+ // K^T first (the compute waits on K before Q)
68
+ cb_reserve_back(cb_k, kv_tiles);
69
+ if (new_kv) {
70
+ const uint32_t base = get_write_ptr(cb_k);
71
+ uint32_t tid = kvh * Skt * DHt;
72
+ for (uint32_t r = 0; r < Skt; ++r) {
73
+ for (uint32_t c = 0; c < DHt; ++c) {
74
+ noc_async_read(k_rd.get_noc_addr(tid++), base + (c * Skt + r) * tile_bytes, tile_bytes);
75
+ if ((++nb & (BAR - 1)) == 0) noc_async_read_barrier();
76
+ }
77
+ }
78
+ }
79
+ // Q chunk
80
+ cb_reserve_back(cb_q, q_tiles);
81
+ {
82
+ uint32_t wp = get_write_ptr(cb_q);
83
+ uint32_t tid = (bh * Sqt + qc * Sq_chunk_t) * DHt;
84
+ for (uint32_t t = 0; t < q_tiles; ++t) {
85
+ noc_async_read(q_rd.get_noc_addr(tid++), wp, tile_bytes);
86
+ if ((++nb & (BAR - 1)) == 0) noc_async_read_barrier();
87
+ wp += tile_bytes;
88
+ }
89
+ }
90
+ noc_async_read_barrier();
91
+ cb_push_back(cb_k, kv_tiles);
92
+ cb_push_back(cb_q, q_tiles);
93
+
94
+ cb_reserve_back(cb_v, kv_tiles);
95
+ if (new_kv) {
96
+ uint32_t wp = get_write_ptr(cb_v);
97
+ uint32_t tid = kvh * Skt * DHt;
98
+ for (uint32_t t = 0; t < kv_tiles; ++t) {
99
+ noc_async_read(v_rd.get_noc_addr(tid++), wp, tile_bytes);
100
+ if ((++nb & (BAR - 1)) == 0) noc_async_read_barrier();
101
+ wp += tile_bytes;
102
+ }
103
+ noc_async_read_barrier();
104
+ }
105
+ cb_push_back(cb_v, kv_tiles);
106
+ prev_kv = kvh;
107
+ }
108
+ }
code/models/demos/mast3r/tt/kernels/fsdpa/compute_common.hpp ADDED
@@ -0,0 +1,2484 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // mast3r-p150 model-local copy of ttnn sdpa compute (tt-metal 8b98410e730); FSDPA_PROF=1 enables the stock profiling zones.
2
+ // SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
3
+ //
4
+ // SPDX-License-Identifier: Apache-2.0
5
+
6
+ #pragma once
7
+
8
+ #include <cstdint>
9
+
10
+ #define REDUCE_OP (PoolType::MAX)
11
+ #define REDUCE_DIM (ReduceDim::REDUCE_ROW)
12
+
13
+ #include "api/debug/assert.h"
14
+ #include "api/compute/compute_kernel_api.h"
15
+ #include "api/compute/binary_max_min.h"
16
+ #include "api/compute/eltwise_binary.h"
17
+ #include "api/compute/eltwise_unary/exp.h"
18
+ #include "api/compute/eltwise_unary/recip.h"
19
+ #include "api/compute/eltwise_unary/softplus.h"
20
+ #include "api/compute/eltwise_unary/negative.h"
21
+ #include "api/compute/eltwise_unary/binop_with_scalar.h"
22
+ #include "api/compute/bcast.h"
23
+ #include "api/compute/tile_move_copy.h"
24
+ #include "api/compute/matmul.h"
25
+ #include "api/compute/reduce.h"
26
+ #include "api/compute/reduce_custom.h"
27
+ #include "api/dataflow/circular_buffer.h"
28
+ #include "cpp/ttnn/operations/transformer/sdpa/device/kernels/q_chunk_remapping.hpp"
29
+ #include "cpp/ttnn/operations/transformer/sdpa/device/kernels/dataflow/chunked_prefill_utils.hpp"
30
+ #include "cpp/ttnn/kernel_lib/dest_helpers.hpp"
31
+ #if defined(TRISC_MATH) || defined(TRISC_PACK)
32
+ #include "experimental/llk_sfpu/ckernel_sfpu_sdpa.h"
33
+ #endif
34
+
35
+ ALWI void sdpa_reduce_copy_tile_to_dst_init_short(uint32_t cbid, uint32_t transpose = 0) {
36
+ UNPACK((llk_unpack_A_init<BroadcastType::NONE, false, EltwiseBinaryReuseDestType::NONE, UnpackToDestEn>(
37
+ transpose, true /*transpose within 16x16 face*/, cbid)));
38
+
39
+ MATH((llk_math_eltwise_unary_datacopy_init<
40
+ DataCopyType::A2D,
41
+ DST_ACCUM_MODE,
42
+ BroadcastType::NONE,
43
+ false, // is_int_fpu_en
44
+ PackMode::Default>(cbid)));
45
+ }
46
+
47
+ /**
48
+ * in0_cb = max(in0_cb, in1_cb)
49
+ */
50
+ template <uint32_t num_tiles>
51
+ void max_block_inplace(uint32_t in0, uint32_t in1) {
52
+ CircularBuffer cb_in0(in0);
53
+ CircularBuffer cb_in1(in1);
54
+ // inputs come in full, outputs go out full
55
+ copy_tile_to_dst_init_short(in0);
56
+ copy_tile_to_dst_init_short(in1);
57
+ binary_max_tile_init();
58
+ constexpr uint32_t dst_reg_0 = 0;
59
+ constexpr uint32_t dst_reg_1 = 1;
60
+ cb_in0.wait_front(num_tiles);
61
+ cb_in1.wait_front(num_tiles);
62
+ for (uint32_t i = 0; i < num_tiles; ++i) {
63
+ tile_regs_acquire();
64
+ copy_tile(in0, i, dst_reg_0);
65
+ copy_tile(in1, i, dst_reg_1);
66
+ binary_max_tile(dst_reg_0, dst_reg_1, dst_reg_0, VectorMode::C);
67
+ tile_regs_commit();
68
+ tile_regs_wait();
69
+ pack_tile(dst_reg_0, in0);
70
+ tile_regs_release();
71
+ }
72
+ cb_in0.pop_front(num_tiles);
73
+ cb_in0.reserve_back(num_tiles);
74
+ cb_in0.push_back(num_tiles);
75
+ }
76
+
77
+ /**
78
+ * out_cb = eltwise_max(in0, in1)
79
+ */
80
+ template <VectorMode vector_mode = VectorMode::RC>
81
+ void max_block(uint32_t in0, uint32_t in1, uint32_t out_cb, uint32_t num_tiles) {
82
+ CircularBuffer cb_in0(in0);
83
+ CircularBuffer cb_in1(in1);
84
+ CircularBuffer cb_out(out_cb);
85
+ // inputs come in full, outputs go out full
86
+ copy_tile_to_dst_init_short(in0);
87
+ binary_max_tile_init();
88
+
89
+ constexpr uint32_t dst_reg_0 = 0;
90
+ constexpr uint32_t dst_reg_1 = 1;
91
+ cb_in0.wait_front(num_tiles);
92
+ cb_in1.wait_front(num_tiles);
93
+ cb_out.reserve_back(num_tiles);
94
+ for (uint32_t i = 0; i < num_tiles; ++i) {
95
+ tile_regs_acquire();
96
+ copy_tile(in0, i, dst_reg_0);
97
+ copy_tile(in1, i, dst_reg_1);
98
+ binary_max_tile(dst_reg_0, dst_reg_1, dst_reg_0, vector_mode);
99
+ tile_regs_commit();
100
+ tile_regs_wait();
101
+ pack_tile(dst_reg_0, out_cb, i);
102
+ tile_regs_release();
103
+ }
104
+ cb_out.push_back(num_tiles);
105
+ }
106
+
107
+ /**
108
+ * out_cb = reduce[MAX,SUM](in0_cb * scale_cb)
109
+ */
110
+ template <
111
+ PoolType pool_type,
112
+ ReduceDim reduce_dim,
113
+ uint32_t in0_cb,
114
+ uint32_t scale_cb,
115
+ uint32_t rows,
116
+ uint32_t cols,
117
+ VectorMode vector_mode = VectorMode::C>
118
+ void reduce_c(uint32_t out_cb, uint32_t prev_cb, bool do_eltwise_max = false) {
119
+ CircularBuffer cb_in0(in0_cb);
120
+ CircularBuffer cb_scale(scale_cb);
121
+ CircularBuffer cb_out(out_cb);
122
+ CircularBuffer cb_prev(prev_cb);
123
+ // Precondition: in0_cb has rows*cols produced. in0_cb has tiles in row-major order
124
+ // Precondition: scale_cb has 1 produced
125
+ // Precondition: out_cb has rows free
126
+ // Postcondition: in0_cb has rows*cols produced
127
+ // Precondition: scale_cb has 1 produced
128
+ // Postcondition: out_cb has rows produced
129
+ // If do_eltwise_max == true, prev_cb has rows produced.
130
+
131
+ constexpr uint32_t num_tiles = rows * cols;
132
+
133
+ #if defined REDUCE_GRANULARITY
134
+ constexpr uint32_t dst_tiles = (rows < REDUCE_GRANULARITY) ? rows : REDUCE_GRANULARITY;
135
+ constexpr uint32_t granularity = (rows >= REDUCE_GRANULARITY) ? (rows / REDUCE_GRANULARITY) : 1;
136
+ #else
137
+ constexpr uint32_t dst_tiles = 1;
138
+ constexpr uint32_t granularity = rows;
139
+ #endif
140
+
141
+ cb_scale.wait_front(1);
142
+ cb_out.reserve_back(rows);
143
+
144
+ const uint32_t num_tiles_to_wait = dst_tiles * cols;
145
+ uint32_t in0_wait_tiles = num_tiles_to_wait;
146
+
147
+ uint32_t row_start_idx = 0;
148
+ for (uint32_t g = 0; g < granularity; g++) {
149
+ cb_in0.wait_front(in0_wait_tiles);
150
+ tile_regs_acquire();
151
+
152
+ if (do_eltwise_max) {
153
+ cb_prev.wait_front(g * dst_tiles);
154
+ /**
155
+ * Copy previous max values into DST register.
156
+ * Note that this special invocation of copy_tile is necessary to produce
157
+ * tiles in DST with transposed faces, as `reduce_block_max_row` expects.
158
+ */
159
+ reconfig_data_format_srca(prev_cb);
160
+ sdpa_reduce_copy_tile_to_dst_init_short(prev_cb);
161
+ for (uint32_t i = 0; i < dst_tiles; i++) {
162
+ const uint32_t cur_max_dst_idx = i;
163
+ copy_tile(prev_cb, (row_start_idx + i), cur_max_dst_idx);
164
+ }
165
+ reconfig_data_format_srca(in0_cb);
166
+ }
167
+
168
+ /**
169
+ * For `dst_tiles` number of rows, compute the max into the even indices of the DST register.
170
+ */
171
+ reduce_block_max_row_init<cols>(out_cb);
172
+ for (uint32_t i = 0; i < dst_tiles; i++) {
173
+ const uint32_t reduce_dst_idx = i;
174
+ reduce_block_max_row<cols>(in0_cb, scale_cb, (row_start_idx + i) * cols, reduce_dst_idx);
175
+ }
176
+ reduce_block_max_row_uninit(in0_cb);
177
+
178
+ tile_regs_commit();
179
+ tile_regs_wait();
180
+ pack_reconfig_data_format(out_cb);
181
+ for (uint32_t i = 0; i < dst_tiles; i++) {
182
+ const uint32_t cur_max_dst_idx = i;
183
+ pack_tile<true>(cur_max_dst_idx, out_cb, (row_start_idx + i));
184
+ }
185
+ tile_regs_release();
186
+
187
+ row_start_idx += dst_tiles;
188
+ in0_wait_tiles += num_tiles_to_wait;
189
+ }
190
+
191
+ cb_out.push_back(rows);
192
+ }
193
+
194
+ /**
195
+ * out_cb = reduce[MAX,SUM](in0_cb * scale_cb)
196
+ *
197
+ * In this version cols does not have to be a compile-time constant.
198
+ */
199
+ template <
200
+ PoolType pool_type,
201
+ ReduceDim reduce_dim,
202
+ uint32_t in0_cb,
203
+ uint32_t scale_cb,
204
+ uint32_t rows,
205
+ VectorMode vector_mode = VectorMode::C>
206
+ void reduce_c(uint32_t out_cb, uint32_t prev_cb, uint32_t cols, bool do_eltwise_max = false) {
207
+ CircularBuffer cb_in0(in0_cb);
208
+ CircularBuffer cb_scale(scale_cb);
209
+ CircularBuffer cb_out(out_cb);
210
+ // Precondition: in0_cb has rows*cols produced. in0_cb has tiles in row-major order
211
+ // Precondition: scale_cb has 1 produced
212
+ // Precondition: out_cb has rows free
213
+ // Postcondition: in0_cb has rows*cols produced
214
+ // Precondition: scale_cb has 1 produced
215
+ // Postcondition: out_cb has rows produced
216
+
217
+ uint32_t num_tiles = rows * cols;
218
+ cb_scale.wait_front(1);
219
+ cb_in0.wait_front(num_tiles);
220
+ cb_out.reserve_back(rows);
221
+
222
+ pack_reconfig_data_format(out_cb);
223
+
224
+ binary_max_tile_init();
225
+ constexpr uint32_t reduce_dst_idx = 0;
226
+ constexpr uint32_t prev_max_dst_idx = 1;
227
+
228
+ for (uint32_t i = 0; i < rows; i++) {
229
+ reconfig_data_format_srca(in0_cb);
230
+ tile_regs_acquire();
231
+ reduce_init<pool_type, reduce_dim>(in0_cb, scale_cb, out_cb);
232
+ for (uint32_t j = 0; j < cols; j++) {
233
+ reduce_tile<pool_type, reduce_dim>(in0_cb, scale_cb, i * cols + j, 0, reduce_dst_idx);
234
+ }
235
+ reduce_uninit();
236
+ if (do_eltwise_max) {
237
+ reconfig_data_format_srca(prev_cb);
238
+ copy_tile_to_dst_init_short(prev_cb);
239
+ copy_tile(prev_cb, i, prev_max_dst_idx);
240
+ binary_max_tile(reduce_dst_idx, prev_max_dst_idx, reduce_dst_idx, vector_mode);
241
+ }
242
+
243
+ tile_regs_commit();
244
+ tile_regs_wait();
245
+ pack_tile(reduce_dst_idx, out_cb);
246
+ tile_regs_release();
247
+ }
248
+
249
+ cb_out.push_back(rows);
250
+ }
251
+
252
+ #ifdef TRISC_MATH
253
+ template <bool legacy_compat = true>
254
+ void recip_tile_first_column(uint32_t idst) {
255
+ SFPU_UNARY_CALL(DST_SYNC_MODE, DST_ACCUM_MODE, calculate_recip_first_column, (legacy_compat), idst, VectorMode::C);
256
+ }
257
+ #endif
258
+
259
+ /**
260
+ * in_cb = 1 / in_cb
261
+ */
262
+ void recip_block_inplace(uint32_t in_cb, uint32_t num_tiles) {
263
+ CircularBuffer cb_in(in_cb);
264
+ // Precondition: in_cb has num_tiles produced
265
+ // Postcondition: in_cb has num_tiles produced
266
+ reconfig_data_format_srca(in_cb);
267
+ copy_tile_to_dst_init_short(in_cb);
268
+ recip_tile_init();
269
+ pack_reconfig_data_format(in_cb);
270
+
271
+ cb_in.wait_front(num_tiles);
272
+ for (uint32_t i = 0; i < num_tiles; ++i) {
273
+ tile_regs_acquire();
274
+ copy_tile(in_cb, i, 0);
275
+ MATH((recip_tile_first_column(0)));
276
+ tile_regs_commit();
277
+ tile_regs_wait();
278
+ pack_tile(0, in_cb);
279
+ tile_regs_release();
280
+ }
281
+ cb_in.pop_front(num_tiles);
282
+ cb_in.reserve_back(num_tiles);
283
+ cb_in.push_back(num_tiles);
284
+ }
285
+
286
+ /**
287
+ * in0_cb = exp((in0_cb - in1_cb) * scale_fp32)
288
+ */
289
+ template <
290
+ uint32_t in0_cb,
291
+ uint32_t rows,
292
+ uint32_t scale_fp32,
293
+ bool write_result_inplace = true,
294
+ bool do_reduce = true,
295
+ VectorMode vector_mode = VectorMode::RC>
296
+ void sub_exp_block_bcast_cols_inplace(uint32_t in1_cb, uint32_t reduce_cb, uint32_t cols) {
297
+ CircularBuffer cb_in0(in0_cb);
298
+ CircularBuffer cb_in1(in1_cb);
299
+ CircularBuffer cb_reduce(reduce_cb);
300
+ // Precondition: in0_cb has rows*cols produced
301
+ // Precondition: in1_cb has rows produced
302
+ // Postcondition: in0_cb has rows*cols produced
303
+ // Postcondition: in1_cb has rows produced
304
+ // llk_unpack_AB_init (inside sub_bcast_cols_init) validates the live
305
+ // unpacker configuration. Reconfigure first: qk_im can be FP32 while the
306
+ // row maximum is BF16, notably after applying a windowed BF16 mask.
307
+ reconfig_data_format(in0_cb, in1_cb);
308
+ sub_bcast_cols_init(in0_cb, in1_cb);
309
+
310
+ // The exponential function uses InputClamping::None for better performance. This version
311
+ // produces incorrect outputs for inputs <~ -88, but those outputs are guaranteed to be negative.
312
+ // Enable packer ReLU to zero any negative values produced by the exponential approximation.
313
+ exp_tile_init<true /* approx */, scale_fp32, InputClamping::None>();
314
+ PACK((llk_pack_relu_config(ReluConfig::zero())));
315
+
316
+ cb_in0.wait_front(rows * cols);
317
+ cb_in1.wait_front(rows);
318
+ if constexpr (do_reduce) {
319
+ cb_reduce.reserve_back(rows);
320
+ }
321
+
322
+ #ifdef SUB_EXP_GRANULARITY
323
+ uint32_t dst_tiles = (cols < SUB_EXP_GRANULARITY) ? cols : SUB_EXP_GRANULARITY;
324
+ uint32_t granularity = (cols >= SUB_EXP_GRANULARITY) ? (cols / SUB_EXP_GRANULARITY) : 1;
325
+ #else
326
+ uint32_t dst_tiles = cols;
327
+ uint32_t granularity = 1;
328
+ #endif
329
+
330
+ for (uint32_t i = 0; i < rows; ++i) {
331
+ for (uint32_t u = 0; u < granularity; u++) {
332
+ tile_regs_acquire();
333
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
334
+ sub_tiles_bcast_cols(in0_cb, in1_cb, j, i, j);
335
+ constexpr int iterations = (vector_mode == VectorMode::RC) ? 32 /*ITER*/ : 8 /*ITER*/;
336
+ constexpr VectorMode vector_mode_exp = (vector_mode == VectorMode::RC) ? VectorMode::None : vector_mode;
337
+ exp_tile<true /* approx */, false /* scale_en */, InputClamping::None, iterations>(j, vector_mode_exp);
338
+ }
339
+ tile_regs_commit();
340
+
341
+ if constexpr (write_result_inplace) {
342
+ cb_in0.pop_front(dst_tiles);
343
+ cb_in0.reserve_back(dst_tiles);
344
+ }
345
+
346
+ tile_regs_wait();
347
+
348
+ if constexpr (write_result_inplace) {
349
+ pack_reconfig_data_format(in0_cb);
350
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
351
+ pack_tile(j, in0_cb);
352
+ }
353
+ // Granular write output to enable following matmul unpack to start early.
354
+ cb_in0.push_back(dst_tiles);
355
+ }
356
+
357
+ if constexpr (do_reduce) {
358
+ pack_reconfig_data_format(reduce_cb);
359
+ // While we have results in DST, take advantage of L1 accumulation
360
+ // to reduce row x cols tiles to rows x 1 tiles.
361
+ if (u > 0) {
362
+ // If on the same row, keep accumulating
363
+ PACK((llk_pack_reconfig_l1_acc(1)));
364
+ }
365
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
366
+ pack_tile<true>(j, reduce_cb, i);
367
+ if (u == 0 && j == 0) {
368
+ // If this was the first tile of a row, start accumulating
369
+ PACK((llk_pack_reconfig_l1_acc(1)));
370
+ }
371
+ }
372
+ }
373
+ tile_regs_release();
374
+ if constexpr (do_reduce) {
375
+ PACK((llk_pack_reconfig_l1_acc(0)));
376
+ }
377
+ }
378
+ }
379
+ if constexpr (do_reduce) {
380
+ cb_reduce.push_back(rows);
381
+ }
382
+
383
+ PACK((llk_pack_relu_config(ReluConfig::none())));
384
+ }
385
+
386
+ /**
387
+ * out_cb = in0_cb * in1_cb
388
+ * @tparam rows - Number of rows of tiles
389
+ * @tparam cols - Number of columns of tiles
390
+ * @tparam immediate_pop - If true, uses tile-by-tile processing with immediate CB pop after each tile.
391
+ * If false, uses batched processing with deferred CB pop, processing multiple tiles in
392
+ * parallel.
393
+ * @tparam pack_accumulate - If true, enables L1 accumulation to accumulate results onto existing tiles
394
+ * in out_cb. Only supported when immediate_pop=false.
395
+ */
396
+ template <uint32_t rows, uint32_t cols, bool immediate_pop, bool pack_accumulate>
397
+ void mul_block_bcast_cols(uint32_t in0_cb, uint32_t in1_cb, uint32_t out_cb) {
398
+ CircularBuffer cb_in0(in0_cb);
399
+ CircularBuffer cb_in1(in1_cb);
400
+ CircularBuffer cb_out(out_cb);
401
+ // Precondition: in0_cb has rows*cols produced
402
+ // Precondition: in1_cb has rows produced
403
+ // Precondition: out_cb has rows*cols produced
404
+ // Postcondition: in0_cb empty
405
+ // Postcondition: in1_cb empty
406
+ // Postcondition: out_cb has rows*cols produced
407
+
408
+ constexpr uint32_t num_tiles = rows * cols;
409
+
410
+ reconfig_data_format(in0_cb, in1_cb);
411
+ pack_reconfig_data_format(out_cb);
412
+ mul_bcast_cols_init(in0_cb, in1_cb);
413
+ cb_in0.wait_front(num_tiles);
414
+ cb_in1.wait_front(rows);
415
+
416
+ if constexpr (immediate_pop) {
417
+ static_assert(!pack_accumulate, "Unsupported parameter configuration");
418
+ for (uint32_t i = 0; i < rows; ++i) {
419
+ for (uint32_t j = 0; j < cols; ++j) {
420
+ tile_regs_acquire();
421
+ mul_tiles_bcast_cols(in0_cb, in1_cb, 0, i, 0);
422
+ tile_regs_commit();
423
+ cb_in0.pop_front(1);
424
+ cb_out.reserve_back(1);
425
+ tile_regs_wait();
426
+ pack_tile(0, out_cb);
427
+ tile_regs_release();
428
+ cb_out.push_back(1);
429
+ }
430
+ }
431
+ cb_in1.pop_front(rows);
432
+ } else {
433
+ #ifdef DHT_GRANULARITY
434
+ constexpr uint32_t dst_tiles = (cols < DHT_GRANULARITY) ? cols : DHT_GRANULARITY;
435
+ constexpr uint32_t granularity = (cols >= DHT_GRANULARITY) ? (cols / DHT_GRANULARITY) : 1;
436
+ #else
437
+ constexpr uint32_t dst_tiles = 1;
438
+ constexpr uint32_t granularity = cols;
439
+ #endif
440
+ PACK((llk_pack_reconfig_l1_acc(pack_accumulate)));
441
+ if (!pack_accumulate) {
442
+ cb_out.reserve_back(num_tiles);
443
+ }
444
+ uint32_t in0_index = 0;
445
+ for (uint32_t i = 0; i < rows; ++i) {
446
+ for (uint32_t u = 0; u < granularity; ++u) {
447
+ tile_regs_acquire();
448
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
449
+ mul_tiles_bcast_cols(in0_cb, in1_cb, in0_index, i, j);
450
+ in0_index++;
451
+ }
452
+ tile_regs_commit();
453
+ tile_regs_wait();
454
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
455
+ pack_tile(j, out_cb);
456
+ }
457
+ tile_regs_release();
458
+ }
459
+ }
460
+ cb_in1.pop_front(rows);
461
+ cb_in0.pop_front(num_tiles);
462
+ if (pack_accumulate) {
463
+ PACK((llk_pack_reconfig_l1_acc(false)));
464
+ cb_out.pop_front(num_tiles);
465
+ cb_out.reserve_back(num_tiles);
466
+ cb_out.push_back(num_tiles);
467
+ } else {
468
+ cb_out.push_back(num_tiles);
469
+ }
470
+ }
471
+ }
472
+
473
+ /**
474
+ * in0_cb *= in1_cb
475
+ */
476
+ template <uint32_t rows, uint32_t cols>
477
+ void mul_block_bcast_cols_inplace(uint32_t in0_cb, uint32_t in1_cb) {
478
+ CircularBuffer cb_in0(in0_cb);
479
+ CircularBuffer cb_in1(in1_cb);
480
+ // Precondition: in0_cb has rows*cols produced
481
+ // Precondition: in1_cb has rows produced
482
+ // Postcondition: in0_cb has rows*cols produced
483
+ // Postcondition: in1_cb has rows consumed
484
+
485
+ constexpr uint32_t num_tiles = rows * cols;
486
+
487
+ #ifdef DHT_GRANULARITY
488
+ constexpr uint32_t dst_tiles = (cols < DHT_GRANULARITY) ? cols : DHT_GRANULARITY;
489
+ constexpr uint32_t granularity = (cols >= DHT_GRANULARITY) ? (cols / DHT_GRANULARITY) : 1;
490
+ #else
491
+ constexpr uint32_t dst_tiles = 1;
492
+ constexpr uint32_t granularity = cols;
493
+ #endif
494
+
495
+ reconfig_data_format(in0_cb, in1_cb);
496
+ mul_bcast_cols_init(in0_cb, in1_cb);
497
+ pack_reconfig_data_format(in0_cb);
498
+ cb_in0.wait_front(num_tiles);
499
+ cb_in1.wait_front(rows);
500
+ for (uint32_t i = 0; i < rows; ++i) {
501
+ for (uint32_t u = 0; u < granularity; ++u) {
502
+ tile_regs_acquire();
503
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
504
+ mul_tiles_bcast_cols(in0_cb, in1_cb, j, i, j);
505
+ }
506
+ tile_regs_commit();
507
+ cb_in0.pop_front(dst_tiles);
508
+ cb_in0.reserve_back(dst_tiles);
509
+ tile_regs_wait();
510
+ for (uint32_t j = 0; j < dst_tiles; ++j) {
511
+ pack_tile(j, in0_cb);
512
+ }
513
+ cb_in0.push_back(dst_tiles);
514
+ tile_regs_release();
515
+ }
516
+ }
517
+ cb_in1.pop_front(rows);
518
+ reconfig_data_format_srcb(in0_cb);
519
+ }
520
+
521
+ template <uint32_t in1_scalar_cb, uint32_t num_tiles>
522
+ void mul_block_bcast_scalar_inplace(uint32_t in0_cb) {
523
+ CircularBuffer cb_in0(in0_cb);
524
+ CircularBuffer cb_in1_scalar(in1_scalar_cb);
525
+ // Precondition: in0_cb has num_tiles produced
526
+ // Precondition: in1_scalar_cb has 1 produced
527
+ // Postcondition: in0_cb has num_tiles produced
528
+ // Postcondition: in1_scalar_cb has 1 produced
529
+
530
+ #ifdef STATS_GRANULARITY
531
+ constexpr uint32_t dst_tiles = STATS_GRANULARITY;
532
+ constexpr uint32_t granularity = num_tiles / STATS_GRANULARITY;
533
+ #else
534
+ constexpr uint32_t dst_tiles = 1;
535
+ constexpr uint32_t granularity = num_tiles;
536
+ #endif
537
+
538
+ reconfig_data_format(in0_cb, in1_scalar_cb);
539
+ mul_bcast_scalar_init(in0_cb, in1_scalar_cb);
540
+ cb_in0.wait_front(num_tiles);
541
+ cb_in1_scalar.wait_front(1);
542
+ uint32_t in0_index = 0;
543
+ for (uint32_t g = 0; g < granularity; ++g) {
544
+ tile_regs_acquire();
545
+ for (uint32_t i = 0; i < dst_tiles; ++i) {
546
+ mul_tiles_bcast_scalar(in0_cb, in1_scalar_cb, in0_index, 0, i);
547
+ in0_index++;
548
+ }
549
+ tile_regs_commit();
550
+ tile_regs_wait();
551
+ for (uint32_t i = 0; i < dst_tiles; ++i) {
552
+ pack_tile(i, in0_cb);
553
+ }
554
+ tile_regs_release();
555
+ }
556
+ cb_in0.pop_front(num_tiles);
557
+ cb_in0.reserve_back(num_tiles);
558
+ cb_in0.push_back(num_tiles);
559
+ }
560
+
561
+ /**
562
+ * in0_cb += in1_cb
563
+ */
564
+ template <bool pop_in1 = true>
565
+ void add_block_inplace(uint32_t in0_cb, uint32_t in1_cb, uint32_t num_tiles) {
566
+ CircularBuffer cb_in0(in0_cb);
567
+ CircularBuffer cb_in1(in1_cb);
568
+ // Precondition: in0_cb and in1_cb have num_tiles produced
569
+ // Postcondition: in0_cb has num_tiles produced
570
+ // Postcondition: in1_cb has num_tiles consumed
571
+
572
+ reconfig_data_format(in0_cb, in1_cb);
573
+ pack_reconfig_data_format(in0_cb);
574
+ add_init(in0_cb, in1_cb);
575
+ cb_in0.wait_front(num_tiles);
576
+ cb_in1.wait_front(num_tiles);
577
+ for (uint32_t i = 0; i < num_tiles; i++) {
578
+ tile_regs_acquire();
579
+ add_tiles(in0_cb, in1_cb, i, i, 0);
580
+ tile_regs_commit();
581
+ tile_regs_wait();
582
+ pack_tile(0, in0_cb);
583
+ tile_regs_release();
584
+ }
585
+
586
+ cb_in0.pop_front(num_tiles);
587
+ if (pop_in1) {
588
+ cb_in1.pop_front(num_tiles);
589
+ }
590
+ cb_in0.reserve_back(num_tiles);
591
+ cb_in0.push_back(num_tiles);
592
+ }
593
+
594
+ void mul_tiles_bcast_cols_inplace(uint32_t in0_cb, uint32_t in1_cb, uint32_t num_tiles) {
595
+ CircularBuffer cb_in0(in0_cb);
596
+ CircularBuffer cb_in1(in1_cb);
597
+ /**
598
+ * Given in0_cb and in1_cb, multiply each tile of in0_cb by the corresponding tile of in1_cb
599
+ * and bcast cols of in1_cb.
600
+ */
601
+ // Precondition: in0_cb and in1_cb have num_tiles produced
602
+ // Postcondition: in0_cb has num_tiles produced
603
+ // Postcondition: in1_cb has num_tiles produced
604
+
605
+ reconfig_data_format(in0_cb, in1_cb);
606
+ mul_bcast_cols_init(in0_cb, in1_cb);
607
+ pack_reconfig_data_format(in0_cb);
608
+ cb_in0.wait_front(num_tiles);
609
+ cb_in1.wait_front(num_tiles);
610
+ for (uint32_t i = 0; i < num_tiles; i++) {
611
+ tile_regs_acquire();
612
+ mul_tiles_bcast_cols(in0_cb, in1_cb, 0, i, 0);
613
+ tile_regs_commit();
614
+ cb_in0.pop_front(1);
615
+ cb_in0.reserve_back(1);
616
+ tile_regs_wait();
617
+ pack_tile(0, in0_cb);
618
+ tile_regs_release();
619
+ cb_in0.push_back(1);
620
+ }
621
+ }
622
+
623
+ /**
624
+ * in0_cb *= in1_cb
625
+ */
626
+ void mul_block_inplace(uint32_t in0_cb, uint32_t in1_cb, uint32_t num_tiles) {
627
+ CircularBuffer cb_in0(in0_cb);
628
+ CircularBuffer cb_in1(in1_cb);
629
+ // Precondition: in0_cb and in1_cb have num_tiles produced
630
+ // Postcondition: in0_cb has num_tiles produced
631
+ // Postcondition: in1_cb has num_tiles produced
632
+
633
+ mul_init(in0_cb, in1_cb);
634
+ cb_in0.wait_front(num_tiles);
635
+ cb_in1.wait_front(num_tiles);
636
+ for (uint32_t i = 0; i < num_tiles; i++) {
637
+ invalidate_l1_cache();
638
+ tile_regs_acquire();
639
+ mul_tiles(in0_cb, in1_cb, 0, i, 0);
640
+ tile_regs_commit();
641
+ cb_in0.pop_front(1);
642
+ cb_in0.reserve_back(1);
643
+ tile_regs_wait();
644
+ pack_tile(0, in0_cb);
645
+ tile_regs_release();
646
+ cb_in0.push_back(1);
647
+ }
648
+ }
649
+
650
+ #if defined(TRISC_MATH) || defined(TRISC_PACK)
651
+
652
+ template <bool SDPA_EXP_APPROX_MODE, uint16_t scale_bf16>
653
+ void exp_tile_first_column(uint32_t idst) {
654
+ SFPU_UNARY_CALL(
655
+ DST_SYNC_MODE,
656
+ DST_ACCUM_MODE,
657
+ calculate_exponential_first_column,
658
+ (SDPA_EXP_APPROX_MODE, scale_bf16),
659
+ idst,
660
+ VectorMode::C);
661
+ }
662
+ #endif // defined(TRISC_MATH) || defined(TRISC_PACK)
663
+
664
+ /**
665
+ * out_cb = exp((in0_cb - in1_cb) * scale_fp32)
666
+ */
667
+ template <uint32_t scale_fp32>
668
+ void sub_exp_block(uint32_t in0_cb, uint32_t in1_cb, uint32_t out_cb, uint32_t num_tiles) {
669
+ CircularBuffer cb_in0(in0_cb);
670
+ CircularBuffer cb_in1(in1_cb);
671
+ CircularBuffer cb_out(out_cb);
672
+ // Precondition: in0_cb and in1_cb have num_tiles produced
673
+ // Postcondition: out_cb has num_tiles produced
674
+ // Postcondition: in0_cb and in1_cb has num_tiles produced
675
+
676
+ sub_init(in0_cb, in1_cb);
677
+ exp_tile_init<EXP_APPROX_MODE>();
678
+ cb_in0.wait_front(num_tiles);
679
+ cb_in1.wait_front(num_tiles);
680
+ cb_out.reserve_back(num_tiles);
681
+
682
+ // Convert scale_fp32 to bf16 scale
683
+ constexpr uint16_t scale_bf16 = scale_fp32 >> 16;
684
+
685
+ for (uint32_t i = 0; i < num_tiles; i++) {
686
+ invalidate_l1_cache();
687
+ tile_regs_acquire();
688
+ sub_tiles(in0_cb, in1_cb, i, i, 0);
689
+ MATH((exp_tile_first_column<EXP_APPROX_MODE, scale_bf16>(0)));
690
+ tile_regs_commit();
691
+ tile_regs_wait();
692
+ pack_tile(0, out_cb);
693
+ tile_regs_release();
694
+ cb_out.push_back(1);
695
+ }
696
+ }
697
+
698
+ #ifdef TRISC_MATH
699
+ template <VectorMode vector_mode = VectorMode::C>
700
+ void fused_max_sub_exp_add_tile(uint32_t idst, int scale_bf16) {
701
+ SFPU_UNARY_CALL_NO_TEMPLATE_ARGS(
702
+ DST_SYNC_MODE, DST_ACCUM_MODE, calculate_fused_max_sub_exp_add_tile, idst, vector_mode, scale_bf16);
703
+ }
704
+ #endif
705
+
706
+ template <uint32_t scale_fp32, VectorMode vector_mode = VectorMode::C>
707
+ void correction_block(
708
+ uint32_t cb_worker_max,
709
+ uint32_t cb_worker_sum,
710
+ uint32_t cb_cur_max,
711
+ uint32_t cb_prev_max,
712
+ uint32_t cb_cur_sum,
713
+ uint32_t cb_prev_sum,
714
+ uint32_t cb_exp_max_diff,
715
+ uint32_t cb_exp_max_diff_2,
716
+ uint32_t num_head_tiles) {
717
+ CircularBuffer cb_worker_max_obj(cb_worker_max);
718
+ CircularBuffer cb_worker_sum_obj(cb_worker_sum);
719
+ CircularBuffer cb_cur_max_obj(cb_cur_max);
720
+ CircularBuffer cb_prev_max_obj(cb_prev_max);
721
+ CircularBuffer cb_cur_sum_obj(cb_cur_sum);
722
+ CircularBuffer cb_prev_sum_obj(cb_prev_sum);
723
+ CircularBuffer cb_exp_max_diff_obj(cb_exp_max_diff);
724
+ CircularBuffer cb_exp_max_diff_2_obj(cb_exp_max_diff_2);
725
+ cb_worker_max_obj.wait_front(num_head_tiles);
726
+ cb_worker_sum_obj.wait_front(num_head_tiles);
727
+ cb_prev_max_obj.wait_front(num_head_tiles);
728
+ cb_prev_sum_obj.wait_front(num_head_tiles);
729
+
730
+ cb_cur_max_obj.reserve_back(num_head_tiles);
731
+ cb_cur_sum_obj.reserve_back(num_head_tiles);
732
+ cb_exp_max_diff_obj.reserve_back(num_head_tiles);
733
+ cb_exp_max_diff_2_obj.reserve_back(num_head_tiles);
734
+
735
+ constexpr uint32_t dst_reg_0 = 0; // dst_reg_0 is used for prev_max
736
+ constexpr uint32_t dst_reg_1 = 1; // dst_reg_1 is used for worker_max
737
+ constexpr uint32_t dst_reg_2 = 2; // dst_reg_2 is used for cur_max
738
+ constexpr uint32_t dst_reg_3 = 3; // dst_reg_3 is used for prev_sum, returns cur_sum
739
+ constexpr uint32_t dst_reg_4 = 4; // dst_reg_4 is used for worker_sum
740
+
741
+ // convert scale from fp32 to bf16
742
+ constexpr uint16_t scale_bf16 = scale_fp32 >> 16;
743
+
744
+ for (uint32_t i = 0; i < num_head_tiles; i++) {
745
+ tile_regs_acquire();
746
+ copy_tile_to_dst_init_short(cb_worker_max);
747
+ exp_tile_init<EXP_APPROX_MODE>();
748
+ copy_tile(cb_prev_max, i, dst_reg_0);
749
+ copy_tile(cb_worker_max, i, dst_reg_1);
750
+ copy_tile(cb_prev_sum, i, dst_reg_3);
751
+ copy_tile(cb_worker_sum, i, dst_reg_4);
752
+ MATH((fused_max_sub_exp_add_tile<vector_mode>(0, scale_bf16)));
753
+ tile_regs_commit();
754
+ tile_regs_wait();
755
+ pack_tile(dst_reg_0, cb_exp_max_diff);
756
+ pack_tile(dst_reg_1, cb_exp_max_diff_2);
757
+ pack_tile(dst_reg_2, cb_cur_max);
758
+ pack_tile(dst_reg_3, cb_cur_sum);
759
+ tile_regs_release();
760
+ cb_cur_max_obj.push_back(1);
761
+ cb_cur_sum_obj.push_back(1);
762
+ cb_exp_max_diff_obj.push_back(1);
763
+ cb_exp_max_diff_2_obj.push_back(1);
764
+ }
765
+ cb_prev_sum_obj.pop_front(num_head_tiles);
766
+ cb_worker_sum_obj.pop_front(num_head_tiles);
767
+ }
768
+
769
+ /**
770
+ * in_cb -> out_cb
771
+ */
772
+ template <bool pop_in_cb>
773
+ void move_block(uint32_t in_cb, uint32_t out_cb, uint32_t num_tiles) {
774
+ CircularBuffer cb_in(in_cb);
775
+ CircularBuffer cb_out(out_cb);
776
+ // Precondition: in_cb has num_tiles produced
777
+ // Precondition: out_cb has num_tiles free
778
+ // Postcondition: in_cb has num_tiles consumed
779
+ // Postcondition: out_cb has num_tiles produced
780
+
781
+ copy_tile_to_dst_init_short(in_cb);
782
+
783
+ cb_in.wait_front(num_tiles);
784
+ cb_out.reserve_back(num_tiles);
785
+
786
+ #pragma GCC unroll 0
787
+ for (uint32_t i = 0; i < num_tiles; i++) {
788
+ tile_regs_acquire();
789
+ copy_tile(in_cb, i, 0 /*dst*/);
790
+ tile_regs_commit();
791
+ tile_regs_wait();
792
+ pack_tile(0, out_cb);
793
+ tile_regs_release();
794
+ cb_out.push_back(1);
795
+ }
796
+ if (pop_in_cb) {
797
+ cb_in.pop_front(num_tiles);
798
+ }
799
+ }
800
+
801
+ void copy_block(uint32_t in_cb, uint32_t out_cb, uint32_t num_tiles) {
802
+ CircularBuffer cb_in(in_cb);
803
+ CircularBuffer cb_out(out_cb);
804
+ // Precondition: in_cb has num_tiles produced
805
+ // Precondition: out_cb has num_tiles free
806
+ // Postcondition: in_cb has num_tiles consumed
807
+ // Postcondition: out_cb has num_tiles produced
808
+ copy_tile_to_dst_init_short(in_cb);
809
+ cb_in.wait_front(num_tiles);
810
+ cb_out.reserve_back(num_tiles);
811
+ #pragma GCC unroll 0
812
+ for (uint32_t i = 0; i < num_tiles; i++) {
813
+ tile_regs_acquire();
814
+ copy_tile(in_cb, i, 0 /*dst*/);
815
+ tile_regs_commit();
816
+ tile_regs_wait();
817
+ pack_tile(0, out_cb);
818
+ tile_regs_release();
819
+ cb_out.push_back(1);
820
+ }
821
+ cb_in.pop_front(num_tiles);
822
+ }
823
+
824
+ void log_block(uint32_t in_cb, uint32_t out_cb, uint32_t num_tiles) {
825
+ reconfig_data_format_srca(in_cb);
826
+ pack_reconfig_data_format(out_cb);
827
+ CircularBuffer cb_in(in_cb);
828
+ CircularBuffer cb_out(out_cb);
829
+ copy_tile_to_dst_init_short(in_cb);
830
+ log_tile_init();
831
+ cb_in.wait_front(num_tiles);
832
+ cb_out.reserve_back(num_tiles);
833
+
834
+ for (uint32_t i = 0; i < num_tiles; i++) {
835
+ tile_regs_acquire();
836
+ copy_tile(in_cb, i, 0 /*dst*/);
837
+ log_tile(0);
838
+ tile_regs_commit();
839
+ tile_regs_wait();
840
+ pack_tile(0, out_cb);
841
+ tile_regs_release();
842
+ cb_out.push_back(1);
843
+ }
844
+ }
845
+
846
+ void sigmoid_sub(uint32_t in0_cb, uint32_t in1_cb, uint32_t out_cb, uint32_t num_tiles) {
847
+ CircularBuffer cb_in0(in0_cb);
848
+ CircularBuffer cb_in1(in1_cb);
849
+ CircularBuffer cb_out(out_cb);
850
+ // out_cb = sigmoid(in0_cb - in1_cb)
851
+ /**
852
+ * sigmoid(x) is accurately implemented as 1 / (1 + exp(-x))
853
+ * This function manually implements the composite, accurate sigmoid.
854
+ *
855
+ * Each input tile has only the first column containing valid data, so VectorMode::C is a useful optimization.
856
+ */
857
+ cb_in0.wait_front(num_tiles);
858
+ cb_in1.wait_front(num_tiles);
859
+ cb_out.reserve_back(num_tiles);
860
+ sub_init(in0_cb, in1_cb);
861
+ exp_tile_init<false>();
862
+ // recip_tile_first_column<false>() calls the scalar sfpu_reciprocal_iter path, so initialize exactly
863
+ // that SFPU state here. Blackhole needs vConstFloatPrgm0 = 2.0 for Newton-Raphson; Wormhole
864
+ // needs vConstFloatPrgm0/1/2 loaded with reciprocal polynomial coefficients.
865
+ // This init programs persistent SFPU constants, not per-tile data. It intentionally comes after
866
+ // exp_tile_init<false>() because the exp call below is the custom exp_tile_first_column<false>(),
867
+ // whose polynomial implementation loads the constants it consumes into LREGs in the tile body; it
868
+ // does not depend on vConstFloatPrgm* state from exp_tile_init. Conversely, that exp body also
869
+ // does not clobber the reciprocal constants, so one reciprocal init before the tile loop is enough.
870
+ MATH((ckernel::sfpu::sfpu_reciprocal_init<false>()));
871
+
872
+ for (uint32_t i = 0; i < num_tiles; i++) {
873
+ tile_regs_acquire();
874
+ sub_tiles(in0_cb, in1_cb, i, i, 0);
875
+ // exp_tile<false, true /*SCALE_EN*/>(0, (int)VectorMode::C, (uint16_t)0xBF80 /*bf16(-1.0) scale*/);
876
+ MATH((exp_tile_first_column<false /*APPROX_MODE*/, (uint16_t)0xBF80 /*bf16(-1.0) scale*/>(0)));
877
+ // add_unary_tile(0 /*dst_index*/, 0x3F800000); // Call the macro directly to get access to VectorMode argument
878
+ MATH(SFPU_UNARY_CALL(
879
+ DST_SYNC_MODE,
880
+ DST_ACCUM_MODE,
881
+ calculate_binop_with_scalar,
882
+ (APPROX, ADD_UNARY, 8 /* ITERATIONS */),
883
+ 0 /*dst_index*/,
884
+ VectorMode::C,
885
+ 0x3F800000 /*scalar*/));
886
+ // recip_tile<false>(0, (int)VectorMode::C);
887
+ MATH((recip_tile_first_column<false>(0 /*dst_index*/)));
888
+ tile_regs_commit();
889
+ tile_regs_wait();
890
+ pack_tile(0, out_cb);
891
+ tile_regs_release();
892
+ }
893
+ cb_out.push_back(num_tiles);
894
+ }
895
+
896
+ #ifdef TRISC_MATH
897
+ void softplus_tile_first_column(uint32_t idst, uint beta, uint beta_reciprocal, uint threshold) {
898
+ SFPU_UNARY_CALL_NO_TEMPLATE_ARGS(
899
+ DST_SYNC_MODE,
900
+ DST_ACCUM_MODE,
901
+ calculate_softplus_first_column,
902
+ idst,
903
+ VectorMode::C,
904
+ beta,
905
+ beta_reciprocal,
906
+ threshold);
907
+ }
908
+ #endif
909
+
910
+ void logsigmoid_sub(uint32_t in0_cb, uint32_t in1_cb, uint32_t out_cb, uint32_t num_tiles) {
911
+ CircularBuffer cb_in0(in0_cb);
912
+ CircularBuffer cb_in1(in1_cb);
913
+ CircularBuffer cb_out(out_cb);
914
+ // out_cb = logsigmoid(in0_cb - in1_cb)
915
+ // Implemented as softplus for numerical stability. logsigmoid(x) = -softplus(-x)
916
+ cb_in0.wait_front(num_tiles);
917
+ cb_in1.wait_front(num_tiles);
918
+ cb_out.reserve_back(num_tiles);
919
+ sub_init(in0_cb, in1_cb);
920
+ softplus_tile_init();
921
+ constexpr uint32_t const_1_fp32 = 0x3F800000;
922
+ constexpr uint32_t const_20_fp32 = 0x41A00000;
923
+
924
+ for (uint32_t i = 0; i < num_tiles; i++) {
925
+ tile_regs_acquire();
926
+ // Negate input to softplus by swapping inputs to sub
927
+ sub_tiles(in1_cb, in0_cb, i, i, 0);
928
+ // softplus_tile(0, 0x3F800000, 0x3F800000, 0x41A00000); // beta, beta_reciprocal, threshold
929
+ // MATH((llk_math_eltwise_unary_sfpu_softplus<APPROX>(
930
+ // 0,
931
+ // const_1_fp32 /*beta*/,
932
+ // const_1_fp32 /*beta_reciprocal*/,
933
+ // const_20_fp32 /*threshold*/,
934
+ // (int)VectorMode::C)));
935
+
936
+ MATH((softplus_tile_first_column(0, const_1_fp32, const_1_fp32, const_20_fp32)));
937
+ // Negate the output of softplus
938
+ negative_tile(0);
939
+ tile_regs_commit();
940
+ tile_regs_wait();
941
+ pack_tile(0, out_cb);
942
+ tile_regs_release();
943
+ }
944
+ cb_out.push_back(num_tiles);
945
+ }
946
+
947
+ /**
948
+ * out_cb = in0_cb - in1_cb
949
+ * Compile with size optimization to prevent binary size exceeding the limit.
950
+ */
951
+ __attribute__((optimize("Os"))) void sub_block(uint32_t in0_cb, uint32_t in1_cb, uint32_t out_cb, uint32_t num_tiles) {
952
+ CircularBuffer cb_in0(in0_cb);
953
+ CircularBuffer cb_in1(in1_cb);
954
+ CircularBuffer cb_out(out_cb);
955
+ cb_in0.wait_front(num_tiles);
956
+ cb_in1.wait_front(num_tiles);
957
+ cb_out.reserve_back(num_tiles);
958
+ sub_init(in0_cb, in1_cb);
959
+
960
+ for (uint32_t i = 0; i < num_tiles; i++) {
961
+ tile_regs_acquire();
962
+ sub_tiles(in0_cb, in1_cb, i, i, 0);
963
+ tile_regs_commit();
964
+ tile_regs_wait();
965
+ pack_tile(0, out_cb);
966
+ tile_regs_release();
967
+ }
968
+ cb_out.push_back(num_tiles);
969
+ }
970
+
971
+ /**
972
+ * out_cb = in0_cb @ in1_cb
973
+ */
974
+ ALWI void matmul_blocks(
975
+ const uint32_t& in0_cb,
976
+ const uint32_t& in1_cb,
977
+ const uint32_t& out_cb,
978
+ const uint32_t& M,
979
+ const uint32_t& N,
980
+ const uint32_t& K,
981
+ const uint32_t& num_blocks,
982
+ const uint32_t& in0_num_subblocks,
983
+ const uint32_t& in1_num_subblocks,
984
+ const uint32_t& in0_block_w,
985
+ const uint32_t& subblock_h,
986
+ const uint32_t& subblock_w,
987
+ const bool& transpose,
988
+ const bool& add_mask = false,
989
+ const uint32_t& mask_cb = 0,
990
+ const uint32_t& zero_cb = 0) {
991
+ // precondition: in0_cb has M*K produced
992
+ // precondition: in1_cb has K*N produced
993
+ // postcondition: in0_cb is full, in1_cb is empty
994
+ // postcondition: out_cb has M*N produced
995
+
996
+ CircularBuffer cb_in0(in0_cb);
997
+ CircularBuffer cb_in1(in1_cb);
998
+ CircularBuffer cb_out(out_cb);
999
+ CircularBuffer cb_mask(mask_cb);
1000
+ CircularBuffer cb_zero(zero_cb);
1001
+
1002
+ matmul_block_init(
1003
+ in0_cb, in1_cb, transpose /*transpose*/, subblock_w /*ct_dim*/, subblock_h /*rt_dim*/, in0_block_w /*kt_dim*/);
1004
+
1005
+ const uint32_t output_num_tiles = M * N;
1006
+ const uint32_t out_subblock_num_tiles = subblock_h * subblock_w;
1007
+ const uint32_t in0_subblock_all_cols_num_tiles = subblock_h * N;
1008
+
1009
+ uint32_t in0_index_offset = 0;
1010
+
1011
+ const uint32_t in0_subblock_num_tiles = subblock_h * in0_block_w;
1012
+ uint32_t in0_wait_tiles = in0_subblock_num_tiles;
1013
+
1014
+ reconfig_data_format(in1_cb, in0_cb);
1015
+ cb_in1.wait_front(K * N);
1016
+ cb_out.reserve_back(output_num_tiles);
1017
+
1018
+ for (uint32_t in0_subblock = 0; in0_subblock < in0_num_subblocks; ++in0_subblock) {
1019
+ cb_in0.wait_front(in0_wait_tiles);
1020
+ uint32_t in1_index_offset = 0;
1021
+ for (uint32_t in1_subblock = 0; in1_subblock < in1_num_subblocks; ++in1_subblock) {
1022
+ tile_regs_acquire();
1023
+
1024
+ uint32_t dst_index = 0;
1025
+ uint32_t in0_index = in0_index_offset;
1026
+ uint32_t in1_index = in1_index_offset;
1027
+
1028
+ for (uint32_t inner_dim = 0; inner_dim < in0_block_w; inner_dim++) {
1029
+ matmul_block(
1030
+ in0_cb, in1_cb, in0_index, in1_index, dst_index, transpose, subblock_w, subblock_h, in0_block_w);
1031
+ in0_index++;
1032
+ in1_index += N;
1033
+ }
1034
+ if (add_mask) {
1035
+ cb_mask.wait_front(out_subblock_num_tiles);
1036
+ cb_zero.wait_front(1);
1037
+ reconfig_data_format(zero_cb, mask_cb);
1038
+ add_init(zero_cb, mask_cb, true);
1039
+ for (uint32_t i = 0; i < out_subblock_num_tiles; i++) {
1040
+ add_tiles(zero_cb, mask_cb, 0, i, i);
1041
+ }
1042
+ reconfig_data_format(in1_cb, in0_cb);
1043
+ matmul_block_init(in0_cb, in1_cb, transpose, subblock_w, subblock_h, in0_block_w);
1044
+ }
1045
+ tile_regs_commit();
1046
+ tile_regs_wait();
1047
+ uint32_t dst_idx = 0;
1048
+ uint32_t out_col_offset = in1_subblock * subblock_w;
1049
+ for (uint32_t r = 0; r < subblock_h; r++) {
1050
+ uint32_t out_row_offset = r * N;
1051
+ for (uint32_t c = 0; c < subblock_w; c++) {
1052
+ pack_tile<true>(dst_idx, out_cb, out_row_offset + out_col_offset + c);
1053
+ dst_idx++;
1054
+ }
1055
+ }
1056
+ tile_regs_release();
1057
+ in1_index_offset += subblock_w;
1058
+ }
1059
+ in0_index_offset += subblock_h * in0_block_w;
1060
+ in0_wait_tiles += in0_subblock_num_tiles;
1061
+ // Somewhat granularize the push of in0 subblocks
1062
+ cb_out.push_back(in0_subblock_all_cols_num_tiles);
1063
+ }
1064
+ cb_in1.pop_front(K * N);
1065
+ }
1066
+
1067
+ template <uint32_t M>
1068
+ void matmul_reduce(uint32_t in1_cb, const uint32_t& out_cb) {
1069
+ CircularBuffer cb_in1(in1_cb);
1070
+ CircularBuffer cb_out(out_cb);
1071
+ // precondition: in0_cb has M*K produced
1072
+ // precondition: in1_cb has K*N produced
1073
+ // postcondition: in0_cb is full, in1_cb is empty
1074
+ // postcondition: out_cb has M*N produced
1075
+
1076
+ constexpr uint32_t N = 1; // Result of reduce is 1 column
1077
+ constexpr uint32_t in0_block_w = N;
1078
+ constexpr uint32_t subblock_w = N;
1079
+ // Reuse the Sq_chunk_t granularity chosen for sub_exp_block
1080
+ #ifdef STATS_GRANULARITY
1081
+ constexpr uint32_t subblock_h = STATS_GRANULARITY;
1082
+ constexpr uint32_t in0_num_subblocks = M / STATS_GRANULARITY;
1083
+ #else
1084
+ constexpr uint32_t subblock_h = 1;
1085
+ constexpr uint32_t in0_num_subblocks = M;
1086
+ #endif
1087
+
1088
+ /**
1089
+ * Use matmul on Mx1 input to reduce rows within tile to produce Mx1 output.
1090
+ */
1091
+
1092
+ // matmul_block_init validates the live reverse-order unpacker setup
1093
+ // (in1_cb -> SrcA, out_cb -> SrcB), so establish it before init.
1094
+ reconfig_data_format(in1_cb, out_cb);
1095
+ matmul_block_init(
1096
+ out_cb, in1_cb, 0 /*transpose*/, subblock_w /*ct_dim*/, subblock_h /*rt_dim*/, in0_block_w /*kt_dim*/);
1097
+
1098
+ constexpr uint32_t output_num_tiles = M * N;
1099
+ constexpr uint32_t out_subblock_num_tiles = subblock_h * subblock_w;
1100
+
1101
+ pack_reconfig_data_format(out_cb);
1102
+ cb_in1.wait_front(N);
1103
+ cb_out.wait_front(M);
1104
+
1105
+ for (uint32_t in0_subblock = 0; in0_subblock < in0_num_subblocks; ++in0_subblock) {
1106
+ tile_regs_acquire();
1107
+
1108
+ uint32_t dst_index = 0;
1109
+ uint32_t in0_index = 0;
1110
+ uint32_t in1_index = 0;
1111
+
1112
+ matmul_block(out_cb, in1_cb, in0_index, in1_index, dst_index, 0, subblock_w, subblock_h, in0_block_w);
1113
+
1114
+ tile_regs_commit();
1115
+ cb_out.pop_front(subblock_h);
1116
+
1117
+ tile_regs_wait();
1118
+ for (uint32_t i = 0; i < subblock_h; i++) {
1119
+ pack_tile(i, out_cb);
1120
+ }
1121
+ tile_regs_release();
1122
+ cb_out.push_back(subblock_h);
1123
+ }
1124
+ }
1125
+
1126
+ /**
1127
+ * Batch-stamp a single tile onto a range of positions in out_cb using L1 accumulate.
1128
+ * Caller must have already called copy_tile_to_dst_init_short and llk_pack_reconfig_l1_acc(1).
1129
+ *
1130
+ * @tparam dst_batch Max tiles per DST cycle (DST register capacity, typically 8 for fp16b half-sync).
1131
+ */
1132
+ template <uint32_t dst_batch>
1133
+ void stamp_tile_range_l1_acc(
1134
+ uint32_t src_cb, uint32_t src_tile_idx, uint32_t out_cb, uint32_t out_offset, uint32_t count) {
1135
+ for (uint32_t base = 0; base < count; base += dst_batch) {
1136
+ uint32_t batch = (count - base < dst_batch) ? (count - base) : dst_batch;
1137
+ tile_regs_acquire();
1138
+ for (uint32_t i = 0; i < batch; i++) {
1139
+ copy_tile(src_cb, src_tile_idx, i);
1140
+ }
1141
+ tile_regs_commit();
1142
+ tile_regs_wait();
1143
+ for (uint32_t i = 0; i < batch; i++) {
1144
+ pack_tile<true>(i, out_cb, out_offset + base + i);
1145
+ }
1146
+ tile_regs_release();
1147
+ }
1148
+ }
1149
+
1150
+ template <uint32_t dst_batch>
1151
+ void apply_padded_mask_lightweight_runtime(
1152
+ uint32_t neginf_cb,
1153
+ uint32_t neginf_tile_idx,
1154
+ uint32_t out_cb,
1155
+ uint32_t num_padded,
1156
+ uint32_t num_cols,
1157
+ uint32_t num_rows,
1158
+ uint32_t row_base = 0) { // first out_cb tile-row of this query band; nonzero when heads span >1 DEST band
1159
+ uint32_t start = num_cols - num_padded;
1160
+
1161
+ reconfig_data_format_srca(neginf_cb);
1162
+ pack_reconfig_data_format(out_cb);
1163
+ copy_tile_to_dst_init_short(neginf_cb);
1164
+ PACK((llk_pack_reconfig_l1_acc(1)));
1165
+
1166
+ for (uint32_t row = 0; row < num_rows; row++) {
1167
+ stamp_tile_range_l1_acc<dst_batch>(
1168
+ neginf_cb, neginf_tile_idx, out_cb, (row_base + row) * num_cols + start, num_padded);
1169
+ }
1170
+
1171
+ PACK((llk_pack_reconfig_l1_acc(0)));
1172
+ }
1173
+
1174
+ /**
1175
+ * Lightweight partial mask: L1-accumulate a partial mask tile (0 for valid, -inf for padded columns)
1176
+ * onto the boundary tile position in out_cb. The partial tile is permanently fronted in the CB.
1177
+ *
1178
+ * @param mask_cb CB holding mask tiles, permanently fronted
1179
+ * @param partial_tile_idx Index of the partial tile within the CB
1180
+ * @param out_cb QK intermediate CB (already wait-fronted)
1181
+ * @param boundary_col Column index within the chunk where the boundary tile is
1182
+ * @param num_cols Total K tiles per row (Sk_chunk_t)
1183
+ * @param num_rows Q tiles per chunk (Sq_chunk_t)
1184
+ */
1185
+ void apply_partial_mask_lightweight(
1186
+ uint32_t mask_cb,
1187
+ uint32_t partial_tile_idx,
1188
+ uint32_t out_cb,
1189
+ uint32_t boundary_col,
1190
+ uint32_t num_cols,
1191
+ uint32_t num_rows,
1192
+ uint32_t row_base = 0) { // first out_cb tile-row of this query band; nonzero when heads span >1 DEST band
1193
+ reconfig_data_format_srca(mask_cb);
1194
+ pack_reconfig_data_format(out_cb);
1195
+ copy_tile_to_dst_init_short(mask_cb);
1196
+ PACK((llk_pack_reconfig_l1_acc(1)));
1197
+
1198
+ for (uint32_t row = 0; row < num_rows; row++) {
1199
+ tile_regs_acquire();
1200
+ copy_tile(mask_cb, partial_tile_idx, 0);
1201
+ tile_regs_commit();
1202
+ tile_regs_wait();
1203
+ pack_tile<true>(0, out_cb, (row_base + row) * num_cols + boundary_col);
1204
+ tile_regs_release();
1205
+ }
1206
+
1207
+ PACK((llk_pack_reconfig_l1_acc(0)));
1208
+ }
1209
+
1210
+ /**
1211
+ * Lightweight causal mask: stamps neginf and diagonal tiles onto QKT using L1 accumulate.
1212
+ *
1213
+ * For Q tile-row i (0..num_rows-1) processing K chunk starting at k_start_tile:
1214
+ * diag_col = q_start_tile + i - k_start_tile
1215
+ * - diag_col < 0: entire row above diagonal -> stamp neginf on all num_cols tiles
1216
+ * - 0 <= diag_col < num_cols: diagonal tile at col diag_col, neginf at cols diag_col+1..num_cols-1
1217
+ * - diag_col >= num_cols: entire row below diagonal -> no mask needed
1218
+ */
1219
+ template <uint32_t dst_batch>
1220
+ void apply_causal_mask_lightweight(
1221
+ uint32_t mask_cb,
1222
+ uint32_t neginf_idx,
1223
+ uint32_t diag_idx,
1224
+ uint32_t out_cb,
1225
+ uint32_t q_start_tile,
1226
+ uint32_t k_start_tile,
1227
+ uint32_t num_rows,
1228
+ uint32_t num_cols,
1229
+ uint32_t straddle_col = 0,
1230
+ uint32_t straddle_jump = 0) {
1231
+ reconfig_data_format_srca(mask_cb);
1232
+ pack_reconfig_data_format(out_cb);
1233
+ copy_tile_to_dst_init_short(mask_cb);
1234
+ PACK((llk_pack_reconfig_l1_acc(1)));
1235
+
1236
+ for (uint32_t row = 0; row < num_rows; row++) {
1237
+ uint32_t row_offset = row * num_cols;
1238
+ const int32_t q_pos = static_cast<int32_t>(q_start_tile + row);
1239
+
1240
+ if (straddle_col == 0) {
1241
+ // Fast path: K coords contiguous across cols.
1242
+ int32_t diag_col = q_pos - static_cast<int32_t>(k_start_tile);
1243
+ if (diag_col < 0) {
1244
+ // Entire row above diagonal -> stamp all neginf
1245
+ stamp_tile_range_l1_acc<dst_batch>(mask_cb, neginf_idx, out_cb, row_offset, num_cols);
1246
+ } else if (static_cast<uint32_t>(diag_col) < num_cols) {
1247
+ // Stamp the diagonal tile
1248
+ tile_regs_acquire();
1249
+ copy_tile(mask_cb, diag_idx, 0);
1250
+ tile_regs_commit();
1251
+ tile_regs_wait();
1252
+ pack_tile<true>(0, out_cb, row_offset + static_cast<uint32_t>(diag_col));
1253
+ tile_regs_release();
1254
+
1255
+ // Stamp neginf tiles to the right of diagonal
1256
+ uint32_t neginf_start = static_cast<uint32_t>(diag_col) + 1;
1257
+ if (neginf_start < num_cols) {
1258
+ stamp_tile_range_l1_acc<dst_batch>(
1259
+ mask_cb, neginf_idx, out_cb, row_offset + neginf_start, num_cols - neginf_start);
1260
+ }
1261
+ }
1262
+ // else: diag_col >= num_cols -> entire row below diagonal, no mask needed
1263
+ } else {
1264
+ // Chunked-prefill straddle: K coord jumps by straddle_jump at col >= straddle_col
1265
+ // (the K-chunk crosses a slab boundary). Evaluate per-col.
1266
+ for (uint32_t col = 0; col < num_cols; col++) {
1267
+ int32_t k_pos = static_cast<int32_t>(k_start_tile) + static_cast<int32_t>(col);
1268
+ if (col >= straddle_col) {
1269
+ k_pos += static_cast<int32_t>(straddle_jump);
1270
+ }
1271
+ if (k_pos > q_pos) {
1272
+ stamp_tile_range_l1_acc<dst_batch>(mask_cb, neginf_idx, out_cb, row_offset + col, 1);
1273
+ } else if (k_pos == q_pos) {
1274
+ tile_regs_acquire();
1275
+ copy_tile(mask_cb, diag_idx, 0);
1276
+ tile_regs_commit();
1277
+ tile_regs_wait();
1278
+ pack_tile<true>(0, out_cb, row_offset + col);
1279
+ tile_regs_release();
1280
+ }
1281
+ }
1282
+ }
1283
+ }
1284
+
1285
+ PACK((llk_pack_reconfig_l1_acc(0)));
1286
+ }
1287
+
1288
+ /**
1289
+ * Context for lightweight mask application.
1290
+ * All mask tiles reside in a single CB. This struct stores the pre-resolved mask metadata used when
1291
+ * lightweight masking is enabled; enablement itself is controlled by the `lightweight_mask_enabled`
1292
+ * template parameter(s), not by default-constructing this context.
1293
+ */
1294
+ struct LightweightMaskContext {
1295
+ bool is_causal = false; // Causal masking active for this context instance
1296
+ uint32_t neginf_tile_idx = 0; // Index of -inf tile in the mask CB
1297
+ uint32_t causal_diag_tile_idx = 0; // Index of causal diagonal tile in the mask CB
1298
+ uint32_t primary_diag_tile_idx = 0; // Causal diagonal, or sliding-window trailing-primary tile
1299
+ uint32_t sliding_leading_prev_tile_idx = 0; // Index of previous sliding-window leading tile
1300
+ uint32_t sliding_leading_tile_idx = 0; // Index of current sliding-window leading tile in the mask CB
1301
+ uint32_t sliding_trailing_next_tile_idx = 0; // Index of next sliding-window trailing tile
1302
+ uint32_t global_n_padded_tiles = 0; // Fully padded K tile columns for global_n chunk
1303
+ uint32_t local_n_padded_tiles = 0; // Fully padded K tile columns for local_n chunk
1304
+ uint32_t joint_n_padded_tiles = 0; // Fully padded K tile columns for joint_l chunk
1305
+ uint32_t global_n_partial_col = 0; // Column within tile where global_n padding starts (0 = no partial)
1306
+ uint32_t joint_l_partial_col = 0; // Column within tile where joint_l padding starts (0 = no partial)
1307
+ uint32_t global_n_partial_tile_idx = 0; // Index of global_n partial tile in the mask CB
1308
+ uint32_t joint_l_partial_tile_idx = 0; // Index of joint_l partial tile in the mask CB
1309
+ uint32_t straddle_num_padded_tiles = 0; // Trailing -inf tiles on straddle chunk (0 = inactive)
1310
+ uint32_t straddle_mask_chunk_id = 0; // K chunk index where straddle mask applies
1311
+
1312
+ /**
1313
+ * Resolve which mask type applies for a given K chunk and return pre-resolved params.
1314
+ * Called once per K chunk; the resolved values are then used per subblock.
1315
+ */
1316
+ void resolve_for_chunk(
1317
+ uint32_t Sk_chunk_t,
1318
+ uint32_t k_chunk,
1319
+ uint32_t num_local_k_chunks,
1320
+ bool ring_iter_needs_global_n_mask,
1321
+ bool ring_iter_needs_joint_n_mask,
1322
+ bool local_n_needs_masking,
1323
+ uint32_t global_n_mask_chunk_id,
1324
+ uint32_t local_n_mask_chunk_id,
1325
+ uint32_t joint_n_mask_chunk_id,
1326
+ uint32_t& out_num_padded,
1327
+ uint32_t& out_boundary_col,
1328
+ uint32_t& out_partial_tile_idx,
1329
+ bool& out_has_partial) const {
1330
+ out_num_padded = 0;
1331
+ out_has_partial = false;
1332
+ out_boundary_col = 0;
1333
+ out_partial_tile_idx = 0;
1334
+
1335
+ if (ring_iter_needs_global_n_mask && k_chunk == global_n_mask_chunk_id) {
1336
+ out_num_padded = global_n_padded_tiles;
1337
+ out_has_partial = (global_n_partial_col > 0);
1338
+ out_boundary_col = Sk_chunk_t - out_num_padded - (out_has_partial ? 1 : 0);
1339
+ out_partial_tile_idx = global_n_partial_tile_idx;
1340
+ } else if (local_n_needs_masking && k_chunk == local_n_mask_chunk_id) {
1341
+ out_num_padded = local_n_padded_tiles;
1342
+ } else if (ring_iter_needs_joint_n_mask && (k_chunk - num_local_k_chunks) == joint_n_mask_chunk_id) {
1343
+ out_num_padded = joint_n_padded_tiles;
1344
+ out_has_partial = (joint_l_partial_col > 0);
1345
+ out_boundary_col = Sk_chunk_t - out_num_padded - (out_has_partial ? 1 : 0);
1346
+ out_partial_tile_idx = joint_l_partial_tile_idx;
1347
+ } else if (straddle_num_padded_tiles > 0 && k_chunk == straddle_mask_chunk_id) {
1348
+ out_num_padded = straddle_num_padded_tiles;
1349
+ }
1350
+ }
1351
+
1352
+ /**
1353
+ * Apply the lightweight mask for a given K chunk (non-streaming path).
1354
+ * When is_causal: applies causal stamp, then also applies padding stamp if this K chunk has padding.
1355
+ * When !is_causal: applies padding stamp only.
1356
+ * L1-accumulate stamps are additive, so applying both is safe (-inf + -inf = -inf, -inf + 0 = -inf).
1357
+ *
1358
+ * @tparam dst_size DST register capacity (8 for fp16b half-sync, 4 for fp32_dest_acc).
1359
+ */
1360
+ template <uint32_t dst_size>
1361
+ void apply(
1362
+ uint32_t cb_mask_in,
1363
+ uint32_t cb_qk_im,
1364
+ uint32_t Sk_chunk_t,
1365
+ uint32_t Sq_chunk_t,
1366
+ uint32_t k_chunk,
1367
+ uint32_t num_local_k_chunks,
1368
+ bool ring_iter_needs_global_n_mask,
1369
+ bool ring_iter_needs_joint_n_mask,
1370
+ bool local_n_needs_masking,
1371
+ uint32_t global_n_mask_chunk_id,
1372
+ uint32_t local_n_mask_chunk_id,
1373
+ uint32_t joint_n_mask_chunk_id,
1374
+ uint32_t q_start_tile,
1375
+ uint32_t k_start_tile,
1376
+ uint32_t straddle_col = 0,
1377
+ uint32_t straddle_jump = 0) const {
1378
+ if (is_causal) {
1379
+ apply_causal_mask_lightweight<dst_size>(
1380
+ cb_mask_in,
1381
+ neginf_tile_idx,
1382
+ causal_diag_tile_idx,
1383
+ cb_qk_im,
1384
+ q_start_tile,
1385
+ k_start_tile,
1386
+ Sq_chunk_t,
1387
+ Sk_chunk_t,
1388
+ straddle_col,
1389
+ straddle_jump);
1390
+ }
1391
+
1392
+ // Apply padding stamp (also when is_causal — the causal stamp doesn't handle K padding).
1393
+ uint32_t num_padded, boundary_col, partial_tile_idx;
1394
+ bool has_partial;
1395
+ resolve_for_chunk(
1396
+ Sk_chunk_t,
1397
+ k_chunk,
1398
+ num_local_k_chunks,
1399
+ ring_iter_needs_global_n_mask,
1400
+ ring_iter_needs_joint_n_mask,
1401
+ local_n_needs_masking,
1402
+ global_n_mask_chunk_id,
1403
+ local_n_mask_chunk_id,
1404
+ joint_n_mask_chunk_id,
1405
+ num_padded,
1406
+ boundary_col,
1407
+ partial_tile_idx,
1408
+ has_partial);
1409
+
1410
+ if (has_partial) {
1411
+ apply_partial_mask_lightweight(
1412
+ cb_mask_in, partial_tile_idx, cb_qk_im, boundary_col, Sk_chunk_t, Sq_chunk_t);
1413
+ }
1414
+ if (num_padded > 0) {
1415
+ apply_padded_mask_lightweight_runtime<dst_size>(
1416
+ cb_mask_in, neginf_tile_idx, cb_qk_im, num_padded, Sk_chunk_t, Sq_chunk_t);
1417
+ }
1418
+ }
1419
+ };
1420
+
1421
+ enum SDPAType {
1422
+ STANDARD = 0,
1423
+ JOINT = 1,
1424
+ RING = 2,
1425
+ };
1426
+
1427
+ /******************************************************************************
1428
+ * SDPA INNER LOOP *
1429
+ ******************************************************************************/
1430
+ /**
1431
+ * Use the specialized wrapper functions below instead of calling this directly.
1432
+ *
1433
+ * Template Parameters:
1434
+ * @tparam sdpa_type - SDPA variant: STANDARD, JOINT, or RING
1435
+ * @tparam cb_qk_im - QK intermediate buffer
1436
+ * @tparam cb_identity_scale_in - Identity scale buffer
1437
+ * @tparam cb_attention_sink - Attention sink buffer
1438
+ * @tparam cb_scale_in - Scale buffer
1439
+ * @tparam Sq_chunk_t - Query chunk size in tiles
1440
+ * @tparam Sk_chunk_t - Key chunk size in tiles
1441
+ * @tparam DHt - Head dimension in tiles
1442
+ * @tparam vDHt - Value head dimension in tiles
1443
+ * @tparam use_attention_sink - Whether to use attention sink
1444
+ * @tparam is_causal - Whether to use causal masking
1445
+ * @tparam use_provided_mask - Whether to use user-provided mask
1446
+ * @tparam use_padded_mask - Whether to use padding mask
1447
+ * @tparam use_joint_mask - Whether to use joint mask
1448
+ * @tparam is_chunked - Whether query is chunked
1449
+ * @tparam scale_fp32 - FP32 scale factor
1450
+ * @tparam sliding_window_size - Sliding window attention size
1451
+ * @tparam lightweight_mask_enabled - Enables the lightweight mask path (compile-time gated)
1452
+ *
1453
+ * Runtime Parameters:
1454
+ * @param Skt - Sequence length in tiles
1455
+ * @param qk_in0_block_w - QK matmul block width
1456
+ * @param qk_subblock_w - QK matmul subblock width
1457
+ * @param qk_subblock_h - QK matmul subblock height
1458
+ * @param qk_in0_num_subblocks - QK input0 subblocks
1459
+ * @param qk_in1_num_subblocks - QK input1 subblocks
1460
+ * @param qk_num_blocks - QK number of blocks
1461
+ * @param out_in0_block_w - Output matmul block width
1462
+ * @param out_subblock_w - Output matmul subblock width
1463
+ * @param out_subblock_h - Output matmul subblock height
1464
+ * @param out_in0_num_subblocks - Output input0 subblocks
1465
+ * @param out_in1_num_subblocks - Output input1 subblocks
1466
+ * @param out_num_blocks - Output number of blocks
1467
+ * @param iter_q_start - Query iteration start
1468
+ * @param iter_q_end - Query iteration end
1469
+ * @param q_num_chunks - Total query chunks
1470
+ * @param local_q_start - Local query start
1471
+ * @param chunked_q_chunk_offset - Chunked query offset
1472
+ * @param iter_k_chunk_start - Key chunk iteration start
1473
+ * @param iter_k_chunk_end - Key chunk iteration end
1474
+ * @param q_chunk_tiles - Query chunk tiles
1475
+ * @param k_chunk_tiles - Key chunk tiles
1476
+ * @param qk_chunk_tiles - QK chunk tiles
1477
+ * @param out_chunk_tiles - Output chunk tiles
1478
+ * @param mask_chunk_0 - First mask chunk index
1479
+ * @param mask_chunk_1 - Second mask chunk index
1480
+ * @param ring_iter - Ring iteration index
1481
+ * @param ring_id - Ring ID
1482
+ * @param num_local_k_chunks - Number of K chunks stored locally on this device (used in Ring SDPA)
1483
+ * @param local_padded_Nt - Padded sequence length in tiles for local K/V chunks on this device
1484
+ * @param logical_nt - Logical (unpadded) sequence length in tiles for K/V
1485
+ * @param ring_iter_needs_global_n_mask - Whether current ring iteration requires global N masking
1486
+ * @param ring_iter_needs_joint_n_mask - Whether current ring iteration requires joint N masking
1487
+ * @param local_n_needs_masking - Whether local N dimension requires masking
1488
+ * @param global_n_mask_chunk_id - K chunk index where global N mask should be applied
1489
+ * @param local_n_mask_chunk_id - K chunk index where local N mask should be applied
1490
+ * @param joint_n_mask_chunk_id - K chunk index where joint N mask should be applied (relative to joint chunks)
1491
+ * @param cb_q_in - Query input buffer
1492
+ * @param cb_k_in - Key input buffer
1493
+ * @param cb_v_in - Value input buffer
1494
+ * @param cb_mask_in - Mask input buffer
1495
+ * @param cb_col_identity - Column identity buffer
1496
+ * @param cb_out_im_A - Output intermediate buffer A
1497
+ * @param cb_out_im_B - Output intermediate buffer B
1498
+ * @param cb_max_A - Max buffer A
1499
+ * @param cb_max_B - Max buffer B
1500
+ * @param cb_sum_A - Sum buffer A
1501
+ * @param cb_sum_B - Sum buffer B
1502
+ * @param cb_exp_max_diff - Exp max diff buffer
1503
+ * @param cb_lse_in - LSE input buffer
1504
+ * @param cb_lse_out - LSE output buffer
1505
+ * @param cb_prev_out - Previous output buffer
1506
+ * @param cb_out - Output buffer
1507
+ */
1508
+ template <
1509
+ SDPAType sdpa_type,
1510
+ uint32_t cb_qk_im,
1511
+ uint32_t cb_identity_scale_in,
1512
+ uint32_t cb_attention_sink,
1513
+ uint32_t cb_scale_in,
1514
+ uint32_t Sq_chunk_t,
1515
+ uint32_t Sk_chunk_t,
1516
+ uint32_t NH,
1517
+ uint32_t DHt,
1518
+ uint32_t vDHt,
1519
+ bool use_attention_sink,
1520
+ bool use_provided_mask,
1521
+ bool use_padded_mask,
1522
+ bool use_joint_mask,
1523
+ bool is_chunked,
1524
+ uint32_t scale_fp32,
1525
+ uint32_t sliding_window_size,
1526
+ bool lightweight_mask_enabled = false,
1527
+ bool chunked_enabled = false,
1528
+ uint32_t chunked_q_local_padded_Nt = 0,
1529
+ uint32_t chunked_chunk_size_t = 0>
1530
+ void sdpa_inner_loop(
1531
+ const uint32_t Skt,
1532
+ const uint32_t qk_in0_block_w,
1533
+ const uint32_t qk_subblock_w,
1534
+ const uint32_t qk_subblock_h,
1535
+ const uint32_t qk_in0_num_subblocks,
1536
+ const uint32_t qk_in1_num_subblocks,
1537
+ const uint32_t qk_num_blocks,
1538
+ const uint32_t out_in0_block_w,
1539
+ const uint32_t out_subblock_w,
1540
+ const uint32_t out_subblock_h,
1541
+ const uint32_t out_in0_num_subblocks,
1542
+ const uint32_t out_in1_num_subblocks,
1543
+ const uint32_t out_num_blocks,
1544
+ const uint32_t iter_q_start,
1545
+ const uint32_t iter_q_end,
1546
+ const uint32_t q_num_chunks,
1547
+ const uint32_t local_q_start,
1548
+ const uint32_t chunked_q_chunk_offset,
1549
+ const uint32_t iter_k_chunk_start,
1550
+ const uint32_t iter_k_chunk_end,
1551
+ const uint32_t q_chunk_tiles,
1552
+ const uint32_t k_chunk_tiles,
1553
+ const uint32_t v_chunk_tiles,
1554
+ const uint32_t qk_chunk_tiles,
1555
+ const uint32_t out_chunk_tiles,
1556
+ const uint32_t mask_chunk_0,
1557
+ const uint32_t mask_chunk_1,
1558
+ const uint32_t ring_iter,
1559
+ const uint32_t ring_id,
1560
+ const uint32_t num_local_k_chunks,
1561
+ const uint32_t local_padded_Nt,
1562
+ const uint32_t logical_nt,
1563
+ const bool ring_iter_needs_global_n_mask,
1564
+ const bool ring_iter_needs_joint_n_mask,
1565
+ const bool local_n_needs_masking,
1566
+ const uint32_t global_n_mask_chunk_id,
1567
+ const uint32_t local_n_mask_chunk_id,
1568
+ const uint32_t joint_n_mask_chunk_id,
1569
+ const uint32_t cb_q_in,
1570
+ const uint32_t cb_k_in,
1571
+ const uint32_t cb_v_in,
1572
+ const uint32_t cb_mask_in,
1573
+ const uint32_t cb_col_identity,
1574
+ const uint32_t cb_out_im_A,
1575
+ const uint32_t cb_out_im_B,
1576
+ const uint32_t cb_max_A,
1577
+ const uint32_t cb_max_B,
1578
+ const uint32_t cb_sum_A,
1579
+ const uint32_t cb_sum_B,
1580
+ const uint32_t cb_exp_max_diff,
1581
+ const uint32_t cb_lse_in,
1582
+ const uint32_t cb_lse_out,
1583
+ const uint32_t cb_prev_out,
1584
+ const uint32_t cb_out,
1585
+ const LightweightMaskContext& lw_mask = {},
1586
+ const bool is_causal = false,
1587
+ const bool is_balanced = false,
1588
+ const bool use_zigzag_balancing = false,
1589
+ const bool is_last_ring_iter = true,
1590
+ const ChunkedContext& chunked = {}) {
1591
+ // Parameter-stable CB locals. Aliases (cb_sum_A/B, cb_max_A/B, cb_out_im_A/B) are
1592
+ // std::swap-mutated below, so they use inline CircularBuffer(alias).method() at the call
1593
+ // sites instead. cb_out is constructed conditionally near its consumer (cur.out target).
1594
+ CircularBuffer cb_q_in_obj(cb_q_in);
1595
+ CircularBuffer cb_k_in_obj(cb_k_in);
1596
+ CircularBuffer cb_v_in_obj(cb_v_in);
1597
+ CircularBuffer cb_qk_im_obj(cb_qk_im);
1598
+ CircularBuffer cb_attention_sink_obj(cb_attention_sink);
1599
+ CircularBuffer cb_lse_in_obj(cb_lse_in);
1600
+ CircularBuffer cb_prev_out_obj(cb_prev_out);
1601
+ constexpr uint32_t dst_size = compute_kernel_lib::DEST_AUTO_LIMIT;
1602
+ uint32_t KV_chunks_processed_in_iter = 0;
1603
+ const uint32_t q_per_core = iter_q_end - iter_q_start;
1604
+
1605
+ for (uint32_t q_iter = iter_q_start; q_iter < iter_q_end; ++q_iter) {
1606
+ uint32_t q_start_tile = 0; // First tile of Q chunk (tile units, both STANDARD and RING)
1607
+ uint32_t q_high_tile = 0; // STANDARD: upper tile bound for K iteration
1608
+ uint32_t causal_k_limit = 0; // RING: K-chunk index beyond which all K is above the diagonal
1609
+ if constexpr (sdpa_type == STANDARD) {
1610
+ const uint32_t linear_q_chunk = local_q_start + (q_iter - iter_q_start);
1611
+ // Mod is a no-op when the input is per-head ([0, q_num_chunks)) and extracts the
1612
+ // per-head q_chunk when it's a flat global index (global Q scheduling spans heads).
1613
+ uint32_t q_chunk = remap_q_index(linear_q_chunk, q_num_chunks, use_zigzag_balancing) % q_num_chunks;
1614
+ // Get Q chunk
1615
+ if constexpr (is_chunked) {
1616
+ q_chunk = chunked_q_chunk_offset + q_chunk;
1617
+ }
1618
+ q_start_tile = q_chunk * Sq_chunk_t;
1619
+ if (is_causal) {
1620
+ // Clamp to total K-tile extent. Mirrors reader_interleaved's clamp; without
1621
+ // both, the reader and compute disagree on K-chunk count when Q-chunk extends
1622
+ // past total K (Sq_chunk_t > Skt) → CB deadlock.
1623
+ const uint32_t q_high_unclamped = q_start_tile + Sq_chunk_t;
1624
+ q_high_tile = q_high_unclamped < Skt ? q_high_unclamped : Skt;
1625
+ } else {
1626
+ q_high_tile = Skt;
1627
+ }
1628
+ } else if (sdpa_type == RING) {
1629
+ uint32_t q_chunk = remap_q_index(q_iter, q_num_chunks, use_zigzag_balancing) % q_num_chunks;
1630
+
1631
+ if constexpr (chunked_enabled) {
1632
+ // Absolute Q tile row. Diag stamp masks K past Q's range; logical_n skip handles K past the cache.
1633
+ q_start_tile =
1634
+ chunked.q_start_idx_t + chunked.ring_index * chunked_q_local_padded_Nt + q_chunk * Sq_chunk_t;
1635
+ } else if (is_causal) {
1636
+ q_start_tile = q_chunk * Sq_chunk_t;
1637
+ causal_k_limit = (q_start_tile + Sq_chunk_t + Sk_chunk_t - 1) / Sk_chunk_t;
1638
+ }
1639
+ if (is_balanced && (q_chunk < q_num_chunks / 2)) {
1640
+ continue;
1641
+ }
1642
+ } // If ring attention
1643
+
1644
+ // Set up ping pong buffers
1645
+ uint32_t alias_prev_sum = cb_sum_A;
1646
+ uint32_t alias_cur_sum = cb_sum_B;
1647
+ uint32_t alias_prev_max = cb_max_A;
1648
+ uint32_t alias_cur_max = cb_max_B;
1649
+ uint32_t alias_mm2_prev_out = cb_out_im_A;
1650
+ uint32_t alias_mm2_cur_out = cb_out_im_B;
1651
+
1652
+ uint32_t k_chunk_end;
1653
+ if constexpr (sdpa_type == STANDARD) {
1654
+ // loop while k_low < q_high => (k_chunk * Sk_chunk_t) < q_high_tile.
1655
+ k_chunk_end = (q_high_tile + Sk_chunk_t - 1) / Sk_chunk_t;
1656
+ } else { // RING or JOINT.
1657
+ k_chunk_end = iter_k_chunk_end;
1658
+ }
1659
+
1660
+ uint32_t processed_k_chunks = 0;
1661
+
1662
+ for (uint32_t k_chunk = iter_k_chunk_start; k_chunk < k_chunk_end; ++k_chunk) {
1663
+ uint32_t kv_global_start_tile = 0; // RING only: abs K-tile index of this k_chunk's start
1664
+ if constexpr (sdpa_type == RING) {
1665
+ const bool kv_chunk_is_joint = k_chunk >= num_local_k_chunks;
1666
+ // Global index into the padded KV tensor. Chunked: non-monotonic local→abs map.
1667
+ if constexpr (chunked_enabled) {
1668
+ kv_global_start_tile =
1669
+ kv_global_tile_for_local<true, 0, chunked_chunk_size_t, chunked_q_local_padded_Nt>(
1670
+ ring_id, k_chunk * Sk_chunk_t);
1671
+ } else {
1672
+ kv_global_start_tile = local_padded_Nt * ring_id + k_chunk * Sk_chunk_t;
1673
+ }
1674
+ if (!kv_chunk_is_joint && (kv_global_start_tile >= logical_nt)) {
1675
+ // This is a KV chunk on spatial input beyond the logical N, and not joint KV. Skip it.
1676
+ continue;
1677
+ }
1678
+ }
1679
+
1680
+ KV_chunks_processed_in_iter++;
1681
+
1682
+ // Chunked-prefill: never take this skip (local-frame causal_k_limit doesn't apply —
1683
+ // the diag stamp uses absolute coords every k_chunk instead).
1684
+ if (sdpa_type == RING && !chunked_enabled && k_chunk >= causal_k_limit && is_causal) {
1685
+ cb_k_in_obj.wait_front(k_chunk_tiles);
1686
+ cb_v_in_obj.wait_front(v_chunk_tiles);
1687
+ cb_k_in_obj.pop_front(k_chunk_tiles);
1688
+ cb_v_in_obj.pop_front(v_chunk_tiles);
1689
+
1690
+ continue;
1691
+ }
1692
+
1693
+ /**
1694
+ * QK = Q_CHUNK @ K_CHUNK
1695
+ *
1696
+ * matmul_blocks internally waits on both inputs
1697
+ */
1698
+ reconfig_data_format(cb_k_in, cb_q_in);
1699
+ pack_reconfig_data_format(cb_qk_im);
1700
+ matmul_blocks(
1701
+ cb_q_in,
1702
+ cb_k_in,
1703
+ cb_qk_im,
1704
+ Sq_chunk_t,
1705
+ Sk_chunk_t,
1706
+ DHt,
1707
+ qk_num_blocks,
1708
+ qk_in0_num_subblocks,
1709
+ qk_in1_num_subblocks,
1710
+ qk_in0_block_w,
1711
+ qk_subblock_h,
1712
+ qk_subblock_w,
1713
+ true /*transpose*/);
1714
+
1715
+ /**
1716
+ * Note
1717
+ * Typically, scores is multiplied by a scalar here. We employed an optimization
1718
+ * where we fuse the scaling into exp both in exp(x - max) and exp(prev_max - cur_max).
1719
+ * This gives us scaling for free on the performance-critical exp(x - max) computation.
1720
+ */
1721
+
1722
+ bool apply_mask = false;
1723
+ bool needs_padding_mask =
1724
+ (sdpa_type == RING) &&
1725
+ ((ring_iter_needs_global_n_mask && k_chunk == global_n_mask_chunk_id) ||
1726
+ (local_n_needs_masking && k_chunk == local_n_mask_chunk_id) ||
1727
+ (ring_iter_needs_joint_n_mask && (k_chunk - num_local_k_chunks) == joint_n_mask_chunk_id) ||
1728
+ (lw_mask.straddle_num_padded_tiles > 0 && k_chunk == lw_mask.straddle_mask_chunk_id));
1729
+ if (sdpa_type == RING && !is_causal) {
1730
+ apply_mask = needs_padding_mask;
1731
+ } else if (is_causal || sliding_window_size > 0) {
1732
+ // Chunked-prefill (RING): use abs K so the causal-overlap test matches the diag stamp's frame.
1733
+ uint32_t k_low_idx;
1734
+ uint32_t k_high_idx;
1735
+ if constexpr (sdpa_type == RING && chunked_enabled) {
1736
+ k_low_idx = kv_global_start_tile;
1737
+ k_high_idx = k_low_idx + Sk_chunk_t;
1738
+ // Straddle: chunk crosses a per-chunk slab boundary in the local cache. The
1739
+ // tail tiles land in the next slab, so the actual high global K is shifted by
1740
+ // chunk_size_t - q_local_padded_Nt.
1741
+ if (chunked_q_local_padded_Nt > 0) {
1742
+ const uint32_t local_start = k_chunk * Sk_chunk_t;
1743
+ const uint32_t slab_end_local =
1744
+ (local_start / chunked_q_local_padded_Nt + 1) * chunked_q_local_padded_Nt;
1745
+ if (local_start + Sk_chunk_t > slab_end_local) {
1746
+ k_high_idx += (chunked_chunk_size_t - chunked_q_local_padded_Nt);
1747
+ }
1748
+ }
1749
+ } else {
1750
+ k_low_idx = k_chunk * Sk_chunk_t;
1751
+ k_high_idx = k_low_idx + Sk_chunk_t;
1752
+ }
1753
+ // Apply mask if causal overlap, sliding window, or this K chunk has padding
1754
+ apply_mask = (q_start_tile < k_high_idx) || (sliding_window_size > 0) || needs_padding_mask;
1755
+ } else if constexpr (use_provided_mask) {
1756
+ apply_mask = true;
1757
+ } else if constexpr (use_padded_mask) {
1758
+ // Apply mask only on the last K chunk
1759
+ apply_mask = (k_chunk == iter_k_chunk_end - 1);
1760
+ } else if constexpr (use_joint_mask) {
1761
+ // Apply mask for specific chunk combinations
1762
+ apply_mask = (k_chunk == mask_chunk_0) || (k_chunk == mask_chunk_1);
1763
+ }
1764
+
1765
+ if (apply_mask) {
1766
+ /* QK += MASK */
1767
+ reconfig_data_format(cb_qk_im, cb_mask_in);
1768
+ if constexpr (lightweight_mask_enabled) {
1769
+ // Re-enter reserved state on cb_qk_im so the lightweight mask can be stamped in-place.
1770
+ // matmul_blocks above already pushed the QK tiles, so tiles_received has been bumped;
1771
+ // without the pop+push cycle below, reduce_c's wait-front would return immediately
1772
+ // and unpack could start before the mask stamps land in L1.
1773
+ // Safe because pop-front only moves rd_ptr (L1 data is untouched), and on a
1774
+ // single-buffered CB of size qk_chunk_tiles the re-reserved wr_ptr wraps back to the
1775
+ // same physical region holding the QK scores.
1776
+ // Warning: this won't work if cb_qk_im is double-buffered -- the stamps would land
1777
+ // in the other buffer, leaving the QK scores unmasked.
1778
+ cb_qk_im_obj.wait_front(Sk_chunk_t * Sq_chunk_t);
1779
+ cb_qk_im_obj.pop_front(Sk_chunk_t * Sq_chunk_t);
1780
+ cb_qk_im_obj.reserve_back(Sk_chunk_t * Sq_chunk_t);
1781
+ // Chunked-prefill: feed abs K (matches abs q_start_tile so diag stamp lines up).
1782
+ uint32_t k_start_tile_for_mask;
1783
+ uint32_t lw_straddle_col = 0;
1784
+ uint32_t lw_straddle_jump = 0;
1785
+ if constexpr (sdpa_type == RING && chunked_enabled) {
1786
+ k_start_tile_for_mask = kv_global_start_tile;
1787
+ // Chunked-prefill straddle: a K-chunk can begin in one per-chunk K
1788
+ // region and end in the next when k_chunk_size does not divide
1789
+ // q_local_padded_Nt. Global K is non-contiguous across the K-chunk in
1790
+ // that case (jumps by chunk_size_t - q_local_padded_Nt between regions),
1791
+ // so we signal the column boundary (straddle_col) and the jump
1792
+ // (straddle_jump) to the diag stamp — see same comment block in
1793
+ // compute_streaming.hpp:sdpa_inner_loop_step for the full picture.
1794
+ if (chunked_q_local_padded_Nt > 0) {
1795
+ const uint32_t local_start = k_chunk * Sk_chunk_t;
1796
+ const uint32_t slab_end_local =
1797
+ (local_start / chunked_q_local_padded_Nt + 1) * chunked_q_local_padded_Nt;
1798
+ if (local_start + Sk_chunk_t > slab_end_local) {
1799
+ lw_straddle_col = slab_end_local - local_start;
1800
+ lw_straddle_jump = chunked_chunk_size_t - chunked_q_local_padded_Nt;
1801
+ }
1802
+ }
1803
+ } else {
1804
+ k_start_tile_for_mask = k_chunk * Sk_chunk_t;
1805
+ }
1806
+ lw_mask.template apply<dst_size>(
1807
+ cb_mask_in,
1808
+ cb_qk_im,
1809
+ Sk_chunk_t,
1810
+ Sq_chunk_t,
1811
+ k_chunk,
1812
+ num_local_k_chunks,
1813
+ ring_iter_needs_global_n_mask,
1814
+ ring_iter_needs_joint_n_mask,
1815
+ local_n_needs_masking,
1816
+ global_n_mask_chunk_id,
1817
+ local_n_mask_chunk_id,
1818
+ joint_n_mask_chunk_id,
1819
+ q_start_tile,
1820
+ k_start_tile_for_mask,
1821
+ lw_straddle_col,
1822
+ lw_straddle_jump);
1823
+ cb_qk_im_obj.push_back(Sk_chunk_t * Sq_chunk_t);
1824
+ } else {
1825
+ add_block_inplace(cb_qk_im, cb_mask_in, qk_chunk_tiles);
1826
+ }
1827
+ }
1828
+
1829
+ /**
1830
+ * reduce_c can perform both reduce_max and eltwise max with previous result.
1831
+ * if do_eltwise_max:
1832
+ * cur_max = eltwise_max(prev_max, max(qk, dim=-1))
1833
+ * else:
1834
+ * cur_max = max(qk, dim=-1)
1835
+ *
1836
+ * Use the reduce_c overload with cols as a runtime arg which uses standard
1837
+ * reduce_tile + binary_max_tile. The overload with cols as a template arg
1838
+ * is bf16-only but cb_qk_im could be fp32.
1839
+ */
1840
+ reconfig_data_format(cb_qk_im, cb_identity_scale_in);
1841
+ reduce_c<PoolType::MAX, ReduceDim::REDUCE_ROW, cb_qk_im, cb_identity_scale_in, Sq_chunk_t>(
1842
+ alias_cur_max, alias_prev_max, Sk_chunk_t, processed_k_chunks > 0);
1843
+
1844
+ /**
1845
+ * sub_exp fuses a few operations.
1846
+ * In-place it performs `QK = exp((QK - cur_max) * scale)`
1847
+ *
1848
+ * It also partially performs reduce_sum on the output using L1 accumulation.
1849
+ * `cur_sum = sum_tiles(exp((QK - cur_max) * scale), dim=-1)`
1850
+ *
1851
+ * Partial reduce_sum is used to push the final row_reduction within a tile
1852
+ * outside of the loop over K chunks.
1853
+ */
1854
+ sub_exp_block_bcast_cols_inplace<cb_qk_im, Sq_chunk_t, scale_fp32, true>(
1855
+ alias_cur_max, alias_cur_sum, Sk_chunk_t);
1856
+
1857
+ // Reconfigure unpackers: srcA (context 0) = cb_v_in, srcB (context 1) = cb_qk_im (operands are swapped in
1858
+ // matmul)
1859
+ reconfig_data_format(cb_v_in, cb_qk_im);
1860
+ pack_reconfig_data_format(alias_mm2_cur_out);
1861
+
1862
+ /* OUT_IM = QK @ V_CHUNK */
1863
+ matmul_blocks(
1864
+ cb_qk_im,
1865
+ cb_v_in,
1866
+ alias_mm2_cur_out,
1867
+ Sq_chunk_t,
1868
+ vDHt,
1869
+ Sk_chunk_t,
1870
+ out_num_blocks,
1871
+ out_in0_num_subblocks,
1872
+ out_in1_num_subblocks,
1873
+ out_in0_block_w,
1874
+ out_subblock_h,
1875
+ out_subblock_w,
1876
+ false /*transpose*/);
1877
+
1878
+ cb_qk_im_obj.pop_front(qk_chunk_tiles);
1879
+ reconfig_data_format(alias_prev_max, alias_cur_max);
1880
+
1881
+ /* OUT_ACC += OUT_IM */
1882
+ if (processed_k_chunks > 0) {
1883
+ /**
1884
+ * cb_exp_max_diff = torch.exp((cb_prev_max - cb_cur_max) * scale)
1885
+ * Scale is fused into exp again since max is the max of unscaled scores.
1886
+ */
1887
+ sub_exp_block<scale_fp32>(alias_prev_max, alias_cur_max, cb_exp_max_diff, Sq_chunk_t);
1888
+ CircularBuffer(alias_prev_max).pop_front(Sq_chunk_t);
1889
+
1890
+ /**
1891
+ * cb_prev_sum *= cb_exp_max_diff
1892
+ * This is a bcast_cols since max_diff is a column vector and prev_sum is a partial
1893
+ * reduction, containing the sum of tiles in dim=-1 of QK.
1894
+ */
1895
+ mul_tiles_bcast_cols_inplace(alias_prev_sum, cb_exp_max_diff, Sq_chunk_t);
1896
+
1897
+ /* cb_cur_sum += cb_prev_sum */
1898
+ add_block_inplace(alias_cur_sum, alias_prev_sum, Sq_chunk_t);
1899
+
1900
+ /**
1901
+ * alias_mm2_cur_out += alias_mm2_prev_out * cb_exp_max_diff
1902
+ * This uses L1 accumulation to accumulate onto mm2_cur_out.
1903
+ */
1904
+ mul_block_bcast_cols<Sq_chunk_t, vDHt, false, true>(
1905
+ alias_mm2_prev_out, cb_exp_max_diff, alias_mm2_cur_out);
1906
+ }
1907
+
1908
+ // Swap CB handles to prepare for next iteration
1909
+ std::swap(alias_prev_sum, alias_cur_sum);
1910
+ std::swap(alias_mm2_prev_out, alias_mm2_cur_out);
1911
+ std::swap(alias_prev_max, alias_cur_max);
1912
+
1913
+ processed_k_chunks++;
1914
+ }
1915
+
1916
+ /**
1917
+ * Performs final row-reduction on the partial sum.
1918
+ */
1919
+ matmul_reduce<Sq_chunk_t>(cb_col_identity, alias_prev_sum);
1920
+
1921
+ /**
1922
+ * Process attention sink as a virtual K chunk.
1923
+ * The attention sink provides additional logits that are included in the softmax
1924
+ * denominator but don't contribute to the output (no S @ V computation).
1925
+ * This effectively allows some attention probability to be "absorbed" by the sink,
1926
+ * reducing attention weights on actual tokens.
1927
+ *
1928
+ * Shape of attention_sink: [Sq_chunk_t, 1] tiles
1929
+ * Each head has one sink logit value that is broadcast to all query positions in the chunk.
1930
+ * The reader kernel replicates the per-head value across all Sq_chunk_t positions.
1931
+ */
1932
+ if constexpr (use_attention_sink) {
1933
+ // Treat attention_sink as scores (already scaled)
1934
+ // Shape: [Sq_chunk_t, 1] tiles - same per-head sink value broadcast to all query positions
1935
+
1936
+ // 1. Update running max: cur_max = max(prev_max, attention_sink)
1937
+ // This compares the previous max with the sink logit
1938
+ reconfig_data_format(cb_attention_sink, cb_identity_scale_in);
1939
+
1940
+ reduce_c<PoolType::MAX, ReduceDim::REDUCE_ROW, cb_attention_sink, cb_identity_scale_in, Sq_chunk_t, 1>(
1941
+ alias_cur_max, alias_prev_max, true);
1942
+
1943
+ // 2. Compute exp((prev_max - cur_max) * scale) to rescale previous statistics
1944
+ sub_exp_block<scale_fp32>(alias_prev_max, alias_cur_max, cb_exp_max_diff, Sq_chunk_t);
1945
+ CircularBuffer(alias_prev_max).pop_front(Sq_chunk_t);
1946
+
1947
+ // 3. Rescale previous sum: prev_sum *= exp(prev_max - cur_max)
1948
+ mul_tiles_bcast_cols_inplace(alias_prev_sum, cb_exp_max_diff, Sq_chunk_t);
1949
+ // 4. Compute exp((attention_sink - cur_max) * scale) and accumulate in cur_sum
1950
+ // This adds the attention sink's contribution to the softmax denominator
1951
+ sub_exp_block_bcast_cols_inplace<cb_attention_sink, Sq_chunk_t, scale_fp32, false>(
1952
+ alias_cur_max, alias_cur_sum, 1);
1953
+
1954
+ // 5. Add rescaled previous sum to current sum: cur_sum += prev_sum
1955
+ add_block_inplace(alias_cur_sum, alias_prev_sum, Sq_chunk_t);
1956
+
1957
+ // 6. Update running statistics for final normalization
1958
+ std::swap(alias_prev_sum, alias_cur_sum);
1959
+ std::swap(alias_prev_max, alias_cur_max);
1960
+
1961
+ // 7. Rescale accumulated output: mm2_prev_out *= exp(prev_max - cur_max)
1962
+ // Note: We do NOT compute attention_sink @ V, so output only has real token contributions
1963
+ // But we need to rescale it due to the updated max
1964
+ mul_block_bcast_cols<Sq_chunk_t, vDHt, false, false>(
1965
+ alias_mm2_prev_out, cb_exp_max_diff, alias_mm2_cur_out);
1966
+ std::swap(alias_mm2_prev_out, alias_mm2_cur_out);
1967
+ }
1968
+
1969
+ if constexpr (sdpa_type == RING) {
1970
+ log_block(alias_prev_sum, alias_cur_max, Sq_chunk_t);
1971
+
1972
+ // Scale prev_max by scale_fp32
1973
+ mul_block_bcast_scalar_inplace<cb_scale_in, Sq_chunk_t>(alias_prev_max);
1974
+ add_block_inplace(alias_prev_max, alias_cur_max, Sq_chunk_t);
1975
+
1976
+ /* cb_cur_sum = 1.0 / cb_cur_sum */
1977
+ recip_block_inplace(alias_prev_sum, Sq_chunk_t);
1978
+ /* cb_out_accumulate_im *= cb_cur_sum */
1979
+ mul_block_bcast_cols_inplace<Sq_chunk_t, vDHt>(alias_mm2_prev_out, alias_prev_sum);
1980
+
1981
+ if (ring_iter > 0) {
1982
+ // Update output according to previous and current LSE
1983
+ /**
1984
+ * sig = torch.sigmoid(cur_lse - prev_lse)
1985
+ * out = prev_out - sig * (prev_out - cur_out)
1986
+ * lse = prev_lse - torch.logsigmoid(prev_lse - cur_lse)
1987
+ */
1988
+ cb_lse_in_obj.wait_front(Sq_chunk_t);
1989
+ cb_prev_out_obj.wait_front(out_chunk_tiles);
1990
+
1991
+ uint32_t alias_cur_lse = alias_prev_max; // full
1992
+ uint32_t alias_sig = alias_cur_max; // empty
1993
+ uint32_t alias_cur_out = alias_mm2_prev_out; // full
1994
+ uint32_t alias_sub = alias_mm2_cur_out; // empty
1995
+
1996
+ // alias_sig = sigmoid(alias_cur_lse - cb_lse_in)
1997
+ sigmoid_sub(alias_cur_lse, cb_lse_in, alias_sig, Sq_chunk_t);
1998
+
1999
+ // alias_sub = cb_prev_out - alias_cur_out
2000
+ reconfig_data_format(cb_prev_out, alias_cur_out);
2001
+ sub_block(cb_prev_out, alias_cur_out, alias_sub, out_chunk_tiles);
2002
+ // alias_sub *= alias_sig
2003
+ reconfig_data_format(alias_sub, alias_sig);
2004
+ mul_block_bcast_cols_inplace<Sq_chunk_t, vDHt>(alias_sub, alias_sig);
2005
+ // cb_out = cb_prev_out - alias_sub
2006
+ reconfig_data_format(cb_prev_out, alias_sub);
2007
+ pack_reconfig_data_format(cb_out);
2008
+ sub_block(cb_prev_out, alias_sub, cb_out, out_chunk_tiles);
2009
+ cb_prev_out_obj.pop_front(out_chunk_tiles);
2010
+ CircularBuffer(alias_cur_out).pop_front(out_chunk_tiles);
2011
+ CircularBuffer(alias_sub).pop_front(out_chunk_tiles);
2012
+
2013
+ // alias_sig = sigmoid(cb_lse_in - alias_cur_lse)
2014
+ // alias_cur_lse = log(alias_sig)
2015
+ // cb_lse_out = cb_lse_in - alias_cur_lse
2016
+ pack_reconfig_data_format(alias_sig);
2017
+ reconfig_data_format(cb_lse_in, alias_cur_lse);
2018
+ logsigmoid_sub(cb_lse_in, alias_cur_lse, alias_sig, Sq_chunk_t);
2019
+ sub_block(cb_lse_in, alias_sig, cb_lse_out, Sq_chunk_t);
2020
+ CircularBuffer(alias_sig).pop_front(Sq_chunk_t);
2021
+ CircularBuffer(alias_cur_lse).pop_front(Sq_chunk_t);
2022
+ cb_lse_in_obj.pop_front(Sq_chunk_t);
2023
+ } else {
2024
+ pack_reconfig_data_format(cb_out);
2025
+ copy_block(alias_mm2_prev_out, cb_out, out_chunk_tiles);
2026
+
2027
+ pack_reconfig_data_format(cb_lse_out);
2028
+ copy_block(alias_prev_max, cb_lse_out, Sq_chunk_t);
2029
+ }
2030
+ } else {
2031
+ /* cb_cur_sum = 1.0 / cb_cur_sum */
2032
+ recip_block_inplace(alias_prev_sum, Sq_chunk_t);
2033
+
2034
+ /* cb_out_accumulate_im *= cb_cur_sum */
2035
+ pack_reconfig_data_format(cb_out);
2036
+ mul_block_bcast_cols<Sq_chunk_t, vDHt, false, false>(alias_mm2_prev_out, alias_prev_sum, cb_out);
2037
+
2038
+ // free up cb_prev_max after K chunks
2039
+ CircularBuffer(alias_prev_max).pop_front(Sq_chunk_t);
2040
+ }
2041
+
2042
+ // When q_per_core == 1, Q is identical across ring iterations so we keep it
2043
+ // fronted in the CB and only pop on the last iteration to avoid redundant DRAM re-reads.
2044
+ if (q_per_core > 1 || is_last_ring_iter) {
2045
+ cb_q_in_obj.pop_front(q_chunk_tiles);
2046
+ }
2047
+
2048
+ // Under global Q scheduling the reader pushes one cb_attention_sink slot per Q iter,
2049
+ // so we must drain matching slots inside this loop. Pre-PR the push/pop pair was
2050
+ // 1:1 outside the loop; the new cadence is 1:1 per iter.
2051
+ if constexpr (use_attention_sink) {
2052
+ cb_attention_sink_obj.pop_front(Sq_chunk_t);
2053
+ }
2054
+ }
2055
+
2056
+ if constexpr (sdpa_type == RING) {
2057
+ if (KV_chunks_processed_in_iter % 2 == 0) {
2058
+ cb_k_in_obj.wait_front(k_chunk_tiles);
2059
+ cb_v_in_obj.wait_front(v_chunk_tiles);
2060
+ cb_k_in_obj.pop_front(k_chunk_tiles);
2061
+ cb_v_in_obj.pop_front(v_chunk_tiles);
2062
+ }
2063
+ }
2064
+ }
2065
+
2066
+ /******************************************************************************
2067
+ * SDPA WRAPPER FUNCTIONS *
2068
+ ******************************************************************************/
2069
+
2070
+ /**
2071
+ * Standard SDPA with optional causal masking, attention sink, and sliding window.
2072
+ */
2073
+ template <
2074
+ uint32_t cb_qk_im,
2075
+ uint32_t cb_identity_scale_in,
2076
+ uint32_t cb_attention_sink,
2077
+ uint32_t Sq_chunk_t,
2078
+ uint32_t Sk_chunk_t,
2079
+ uint32_t DHt,
2080
+ uint32_t vDHt,
2081
+ bool use_attention_sink,
2082
+ bool is_causal,
2083
+ bool use_provided_mask,
2084
+ bool use_padded_mask,
2085
+ bool is_chunked,
2086
+ uint32_t scale_fp32,
2087
+ uint32_t sliding_window_size,
2088
+ bool lightweight_mask_enabled = false>
2089
+ void sdpa_standard(
2090
+ const uint32_t Skt,
2091
+ const uint32_t qk_in0_block_w,
2092
+ const uint32_t qk_subblock_w,
2093
+ const uint32_t qk_subblock_h,
2094
+ const uint32_t qk_in0_num_subblocks,
2095
+ const uint32_t qk_in1_num_subblocks,
2096
+ const uint32_t qk_num_blocks,
2097
+ const uint32_t out_in0_block_w,
2098
+ const uint32_t out_subblock_w,
2099
+ const uint32_t out_subblock_h,
2100
+ const uint32_t out_in0_num_subblocks,
2101
+ const uint32_t out_in1_num_subblocks,
2102
+ const uint32_t out_num_blocks,
2103
+ const uint32_t iter_q_start,
2104
+ const uint32_t iter_q_end,
2105
+ const uint32_t q_num_chunks,
2106
+ const uint32_t local_q_start,
2107
+ const uint32_t chunked_q_chunk_offset,
2108
+ const uint32_t k_num_chunks,
2109
+ const uint32_t q_chunk_tiles,
2110
+ const uint32_t k_chunk_tiles,
2111
+ const uint32_t v_chunk_tiles,
2112
+ const uint32_t qk_chunk_tiles,
2113
+ const uint32_t out_chunk_tiles,
2114
+ const uint32_t cb_q_in,
2115
+ const uint32_t cb_k_in,
2116
+ const uint32_t cb_v_in,
2117
+ const uint32_t cb_mask_in,
2118
+ const uint32_t cb_col_identity,
2119
+ const uint32_t cb_out_im_A,
2120
+ const uint32_t cb_out_im_B,
2121
+ const uint32_t cb_max_A,
2122
+ const uint32_t cb_max_B,
2123
+ const uint32_t cb_sum_A,
2124
+ const uint32_t cb_sum_B,
2125
+ const uint32_t cb_exp_max_diff,
2126
+ const uint32_t cb_out,
2127
+ const LightweightMaskContext& lw_mask = {},
2128
+ const bool use_zigzag_balancing = false) {
2129
+ sdpa_inner_loop<
2130
+ STANDARD,
2131
+ cb_qk_im,
2132
+ cb_identity_scale_in,
2133
+ cb_attention_sink,
2134
+ 0, // cb_scale_in (not used)
2135
+ Sq_chunk_t,
2136
+ Sk_chunk_t,
2137
+ 0, // NH (not used)
2138
+ DHt,
2139
+ vDHt,
2140
+ use_attention_sink,
2141
+ use_provided_mask,
2142
+ use_padded_mask,
2143
+ false, // use_joint_mask (not used)
2144
+ is_chunked,
2145
+ scale_fp32,
2146
+ sliding_window_size,
2147
+ lightweight_mask_enabled>(
2148
+ Skt,
2149
+ qk_in0_block_w,
2150
+ qk_subblock_w,
2151
+ qk_subblock_h,
2152
+ qk_in0_num_subblocks,
2153
+ qk_in1_num_subblocks,
2154
+ qk_num_blocks,
2155
+ out_in0_block_w,
2156
+ out_subblock_w,
2157
+ out_subblock_h,
2158
+ out_in0_num_subblocks,
2159
+ out_in1_num_subblocks,
2160
+ out_num_blocks,
2161
+ iter_q_start,
2162
+ iter_q_end,
2163
+ q_num_chunks,
2164
+ local_q_start,
2165
+ chunked_q_chunk_offset,
2166
+ 0, // iter_k_chunk_start
2167
+ k_num_chunks, // iter_k_chunk_end
2168
+ q_chunk_tiles,
2169
+ k_chunk_tiles,
2170
+ v_chunk_tiles,
2171
+ qk_chunk_tiles,
2172
+ out_chunk_tiles,
2173
+ 0, // mask_chunk_0 (not used)
2174
+ 0, // mask_chunk_1 (not used)
2175
+ 0, // ring_iter (not used)
2176
+ 0, // ring_id (not used)
2177
+ 0, // num_local_k_chunks (not used)
2178
+ 0, // local_padded_Nt (not used)
2179
+ 0, // logical_nt (not used)
2180
+ false, // ring_iter_needs_global_n_mask (not used)
2181
+ false, // ring_iter_needs_joint_n_mask (not used)
2182
+ false, // local_n_needs_masking (not used)
2183
+ 0, // global_n_mask_chunk_id (not used)
2184
+ 0, // local_n_mask_chunk_id (not used)
2185
+ 0, // joint_n_mask_chunk_id (not used)
2186
+ cb_q_in,
2187
+ cb_k_in,
2188
+ cb_v_in,
2189
+ cb_mask_in,
2190
+ cb_col_identity,
2191
+ cb_out_im_A,
2192
+ cb_out_im_B,
2193
+ cb_max_A,
2194
+ cb_max_B,
2195
+ cb_sum_A,
2196
+ cb_sum_B,
2197
+ cb_exp_max_diff,
2198
+ 0, // cb_lse_in (not used)
2199
+ 0, // cb_lse_out (not used)
2200
+ 0, // cb_prev_out (not used)
2201
+ cb_out,
2202
+ lw_mask,
2203
+ is_causal,
2204
+ false, // is_balanced (not used)
2205
+ use_zigzag_balancing);
2206
+ }
2207
+
2208
+ /**
2209
+ * Joint SDPA for multi-modal attention.
2210
+ */
2211
+ template <
2212
+ uint32_t cb_qk_im,
2213
+ uint32_t cb_identity_scale_in,
2214
+ uint32_t Sq_chunk_t,
2215
+ uint32_t Sk_chunk_t,
2216
+ uint32_t DHt,
2217
+ bool use_joint_mask,
2218
+ uint32_t scale_fp32>
2219
+ void sdpa_joint(
2220
+ const uint32_t Skt,
2221
+ const uint32_t qk_in0_block_w,
2222
+ const uint32_t qk_subblock_w,
2223
+ const uint32_t qk_subblock_h,
2224
+ const uint32_t qk_in0_num_subblocks,
2225
+ const uint32_t qk_in1_num_subblocks,
2226
+ const uint32_t qk_num_blocks,
2227
+ const uint32_t out_in0_block_w,
2228
+ const uint32_t out_subblock_w,
2229
+ const uint32_t out_subblock_h,
2230
+ const uint32_t out_in0_num_subblocks,
2231
+ const uint32_t out_in1_num_subblocks,
2232
+ const uint32_t out_num_blocks,
2233
+ const uint32_t local_q_start,
2234
+ const uint32_t local_q_end,
2235
+ const uint32_t k_num_chunks,
2236
+ const uint32_t q_chunk_tiles,
2237
+ const uint32_t k_chunk_tiles,
2238
+ const uint32_t qk_chunk_tiles,
2239
+ const uint32_t out_chunk_tiles,
2240
+ const uint32_t mask_chunk_0,
2241
+ const uint32_t mask_chunk_1,
2242
+ const uint32_t cb_q_in,
2243
+ const uint32_t cb_k_in,
2244
+ const uint32_t cb_v_in,
2245
+ const uint32_t cb_mask_in,
2246
+ const uint32_t cb_col_identity,
2247
+ const uint32_t cb_out_im_A,
2248
+ const uint32_t cb_out_im_B,
2249
+ const uint32_t cb_max_A,
2250
+ const uint32_t cb_max_B,
2251
+ const uint32_t cb_sum_A,
2252
+ const uint32_t cb_sum_B,
2253
+ const uint32_t cb_exp_max_diff,
2254
+ const uint32_t cb_out) {
2255
+ sdpa_inner_loop<
2256
+ JOINT,
2257
+ cb_qk_im,
2258
+ cb_identity_scale_in,
2259
+ 0, // cb_attention_sink (not used)
2260
+ 0, // cb_scale_in (not used)
2261
+ Sq_chunk_t,
2262
+ Sk_chunk_t,
2263
+ 0, // NH (not used)
2264
+ DHt,
2265
+ DHt, // vDHt = DHt
2266
+ false, // use_attention_sink (not used)
2267
+ false, // use_provided_mask (not used)
2268
+ false, // use_padded_mask (not used)
2269
+ use_joint_mask,
2270
+ false, // is_chunked (not used)
2271
+ scale_fp32,
2272
+ 0>( // sliding_window_size (not used)
2273
+ Skt,
2274
+ qk_in0_block_w,
2275
+ qk_subblock_w,
2276
+ qk_subblock_h,
2277
+ qk_in0_num_subblocks,
2278
+ qk_in1_num_subblocks,
2279
+ qk_num_blocks,
2280
+ out_in0_block_w,
2281
+ out_subblock_w,
2282
+ out_subblock_h,
2283
+ out_in0_num_subblocks,
2284
+ out_in1_num_subblocks,
2285
+ out_num_blocks,
2286
+ local_q_start, // iter_q_start
2287
+ local_q_end, // iter_q_end
2288
+ 0, // q_num_chunks (not used)
2289
+ local_q_start,
2290
+ 0, // chunked_q_chunk_offset (not used)
2291
+ 0, // iter_k_chunk_start
2292
+ k_num_chunks, // iter_k_chunk_end
2293
+ q_chunk_tiles,
2294
+ k_chunk_tiles,
2295
+ k_chunk_tiles,
2296
+ qk_chunk_tiles,
2297
+ out_chunk_tiles,
2298
+ mask_chunk_0,
2299
+ mask_chunk_1,
2300
+ 0, // ring_iter (not used)
2301
+ 0, // ring_id (not used)
2302
+ 0, // num_local_k_chunks (not used)
2303
+ 0, // local_padded_Nt (not used)
2304
+ 0, // logical_nt (not used)
2305
+ false, // ring_iter_needs_global_n_mask (not used)
2306
+ false, // ring_iter_needs_joint_n_mask (not used)
2307
+ false, // local_n_needs_masking (not used)
2308
+ 0, // global_n_mask_chunk_id (not used)
2309
+ 0, // local_n_mask_chunk_id (not used)
2310
+ 0, // joint_n_mask_chunk_id (not used)
2311
+ cb_q_in,
2312
+ cb_k_in,
2313
+ cb_v_in,
2314
+ cb_mask_in,
2315
+ cb_col_identity,
2316
+ cb_out_im_A,
2317
+ cb_out_im_B,
2318
+ cb_max_A,
2319
+ cb_max_B,
2320
+ cb_sum_A,
2321
+ cb_sum_B,
2322
+ cb_exp_max_diff,
2323
+ 0, // cb_lse_in (not used)
2324
+ 0, // cb_lse_out (not used)
2325
+ 0, // cb_prev_out (not used)
2326
+ cb_out);
2327
+ }
2328
+
2329
+ /**
2330
+ * Ring SDPA for distributed multi-device attention.
2331
+ */
2332
+ template <
2333
+ uint32_t cb_qk_im,
2334
+ uint32_t cb_identity_scale_in,
2335
+ uint32_t cb_scale_in,
2336
+ uint32_t Sq_chunk_t,
2337
+ uint32_t Sk_chunk_t,
2338
+ uint32_t NH,
2339
+ uint32_t DHt,
2340
+ uint32_t vDHt,
2341
+ uint32_t scale_fp32,
2342
+ bool lightweight_mask_enabled = false,
2343
+ bool chunked_enabled = false,
2344
+ uint32_t chunked_q_local_padded_Nt = 0,
2345
+ uint32_t chunked_chunk_size_t = 0>
2346
+ void sdpa_ring(
2347
+ const uint32_t qk_in0_block_w,
2348
+ const uint32_t qk_subblock_w,
2349
+ const uint32_t qk_subblock_h,
2350
+ const uint32_t qk_in0_num_subblocks,
2351
+ const uint32_t qk_in1_num_subblocks,
2352
+ const uint32_t qk_num_blocks,
2353
+ const uint32_t out_in0_block_w,
2354
+ const uint32_t out_subblock_w,
2355
+ const uint32_t out_subblock_h,
2356
+ const uint32_t out_in0_num_subblocks,
2357
+ const uint32_t out_in1_num_subblocks,
2358
+ const uint32_t out_num_blocks,
2359
+ const uint32_t global_q_start,
2360
+ const uint32_t global_q_end,
2361
+ const uint32_t q_num_chunks,
2362
+ const uint32_t iter_k_chunk_start,
2363
+ const uint32_t iter_k_chunk_end,
2364
+ const uint32_t q_chunk_tiles,
2365
+ const uint32_t k_chunk_tiles,
2366
+ const uint32_t v_chunk_tiles,
2367
+ const uint32_t qk_chunk_tiles,
2368
+ const uint32_t out_chunk_tiles,
2369
+ const uint32_t ring_iter,
2370
+ const uint32_t ring_id,
2371
+ const uint32_t num_local_k_chunks,
2372
+ const uint32_t local_padded_Nt,
2373
+ const uint32_t logical_nt,
2374
+ const bool ring_iter_needs_global_n_mask,
2375
+ const bool ring_iter_needs_joint_n_mask,
2376
+ const bool local_n_needs_masking,
2377
+ const uint32_t global_n_mask_chunk_id,
2378
+ const uint32_t local_n_mask_chunk_id,
2379
+ const uint32_t joint_n_mask_chunk_id,
2380
+ const uint32_t cb_q_in,
2381
+ const uint32_t cb_k_in,
2382
+ const uint32_t cb_v_in,
2383
+ const uint32_t cb_mask_in,
2384
+ const uint32_t cb_col_identity,
2385
+ const uint32_t cb_out_im_A,
2386
+ const uint32_t cb_out_im_B,
2387
+ const uint32_t cb_max_A,
2388
+ const uint32_t cb_max_B,
2389
+ const uint32_t cb_sum_A,
2390
+ const uint32_t cb_sum_B,
2391
+ const uint32_t cb_exp_max_diff,
2392
+ const uint32_t cb_lse_in,
2393
+ const uint32_t cb_lse_out,
2394
+ const uint32_t cb_prev_out,
2395
+ const uint32_t cb_out,
2396
+ const LightweightMaskContext& lw_mask,
2397
+ const bool is_causal_ring_iter,
2398
+ const bool skip_first_half_q,
2399
+ const bool is_last_ring_iter,
2400
+ const bool use_zigzag_balancing = false,
2401
+ const ChunkedContext& chunked = {}) {
2402
+ sdpa_inner_loop<
2403
+ RING,
2404
+ cb_qk_im,
2405
+ cb_identity_scale_in,
2406
+ 0, // cb_attention_sink (not used)
2407
+ cb_scale_in,
2408
+ Sq_chunk_t,
2409
+ Sk_chunk_t,
2410
+ NH,
2411
+ DHt,
2412
+ vDHt,
2413
+ false, // use_attention_sink (not used)
2414
+ false, // use_provided_mask (not used)
2415
+ false, // use_padded_mask (not used)
2416
+ false, // use_joint_mask (not used)
2417
+ false, // is_chunked (not used)
2418
+ scale_fp32,
2419
+ 0, // sliding_window_size (not used)
2420
+ lightweight_mask_enabled,
2421
+ chunked_enabled,
2422
+ chunked_q_local_padded_Nt,
2423
+ chunked_chunk_size_t>(
2424
+ 0, // Skt (not used)
2425
+ qk_in0_block_w,
2426
+ qk_subblock_w,
2427
+ qk_subblock_h,
2428
+ qk_in0_num_subblocks,
2429
+ qk_in1_num_subblocks,
2430
+ qk_num_blocks,
2431
+ out_in0_block_w,
2432
+ out_subblock_w,
2433
+ out_subblock_h,
2434
+ out_in0_num_subblocks,
2435
+ out_in1_num_subblocks,
2436
+ out_num_blocks,
2437
+ global_q_start, // iter_q_start
2438
+ global_q_end, // iter_q_end
2439
+ q_num_chunks, // q_num_chunks (total per-head chunks: local + joint)
2440
+ 0, // local_q_start (not used)
2441
+ 0, // chunked_q_chunk_offset (not used)
2442
+ 0,
2443
+ iter_k_chunk_end,
2444
+ q_chunk_tiles,
2445
+ k_chunk_tiles,
2446
+ v_chunk_tiles,
2447
+ qk_chunk_tiles,
2448
+ out_chunk_tiles,
2449
+ 0, // mask_chunk_0 (not used)
2450
+ 0, // mask_chunk_1 (not used)
2451
+ ring_iter,
2452
+ ring_id,
2453
+ num_local_k_chunks,
2454
+ local_padded_Nt,
2455
+ logical_nt,
2456
+ ring_iter_needs_global_n_mask,
2457
+ ring_iter_needs_joint_n_mask,
2458
+ local_n_needs_masking,
2459
+ global_n_mask_chunk_id,
2460
+ local_n_mask_chunk_id,
2461
+ joint_n_mask_chunk_id,
2462
+ cb_q_in,
2463
+ cb_k_in,
2464
+ cb_v_in,
2465
+ cb_mask_in,
2466
+ cb_col_identity,
2467
+ cb_out_im_A,
2468
+ cb_out_im_B,
2469
+ cb_max_A,
2470
+ cb_max_B,
2471
+ cb_sum_A,
2472
+ cb_sum_B,
2473
+ cb_exp_max_diff,
2474
+ cb_lse_in,
2475
+ cb_lse_out,
2476
+ cb_prev_out,
2477
+ cb_out,
2478
+ lw_mask,
2479
+ is_causal_ring_iter,
2480
+ skip_first_half_q,
2481
+ use_zigzag_balancing,
2482
+ is_last_ring_iter,
2483
+ chunked);
2484
+ }
code/models/demos/mast3r/tt/kernels/fsdpa/compute_streaming.hpp ADDED
The diff for this file is too large to render. See raw diff
 
code/models/demos/mast3r/tt/kernels/fsdpa/sdpa.cpp ADDED
@@ -0,0 +1,277 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // mast3r-p150 model-local copy of ttnn sdpa compute (tt-metal 8b98410e730); FSDPA_PROF=1 enables the stock profiling zones.
2
+ // SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
3
+ //
4
+ // SPDX-License-Identifier: Apache-2.0
5
+
6
+ #include <cstdint>
7
+
8
+ #define REDUCE_OP (PoolType::MAX)
9
+ #define REDUCE_DIM (ReduceDim::REDUCE_ROW)
10
+
11
+ #include "api/compute/compute_kernel_api.h"
12
+ #include "api/compute/compute_kernel_hw_startup.h"
13
+ #include "compute_common.hpp"
14
+ #include "compute_streaming.hpp"
15
+
16
+ void kernel_main() {
17
+ constexpr uint32_t B = get_compile_time_arg_val(0);
18
+ constexpr uint32_t NQH = get_compile_time_arg_val(1);
19
+ constexpr uint32_t NKH = get_compile_time_arg_val(2);
20
+ constexpr uint32_t Skt = get_compile_time_arg_val(3);
21
+ constexpr uint32_t DHt = get_compile_time_arg_val(4);
22
+ constexpr uint32_t vDHt = get_compile_time_arg_val(5);
23
+ constexpr uint32_t Sq_chunk_t = get_compile_time_arg_val(6);
24
+ constexpr uint32_t q_num_chunks = get_compile_time_arg_val(7);
25
+ constexpr uint32_t Sk_chunk_t = get_compile_time_arg_val(8);
26
+ constexpr uint32_t k_num_chunks = get_compile_time_arg_val(9);
27
+
28
+ constexpr uint32_t qk_in0_block_w = get_compile_time_arg_val(10);
29
+ constexpr uint32_t qk_subblock_w = get_compile_time_arg_val(11);
30
+ constexpr uint32_t qk_subblock_h = get_compile_time_arg_val(12);
31
+ constexpr uint32_t qk_in0_num_subblocks = get_compile_time_arg_val(13);
32
+ constexpr uint32_t qk_in1_num_subblocks = get_compile_time_arg_val(14);
33
+ constexpr uint32_t qk_num_blocks = get_compile_time_arg_val(15);
34
+ constexpr uint32_t out_in0_block_w = get_compile_time_arg_val(16);
35
+ constexpr uint32_t out_subblock_w = get_compile_time_arg_val(17);
36
+ constexpr uint32_t out_subblock_h = get_compile_time_arg_val(18);
37
+ constexpr uint32_t out_in0_num_subblocks = get_compile_time_arg_val(19);
38
+ constexpr uint32_t out_in1_num_subblocks = get_compile_time_arg_val(20);
39
+ constexpr uint32_t out_num_blocks = get_compile_time_arg_val(21);
40
+
41
+ constexpr uint32_t num_cores = get_compile_time_arg_val(22);
42
+
43
+ constexpr bool is_causal = get_compile_time_arg_val(23) == 1;
44
+ constexpr bool use_provided_mask = get_compile_time_arg_val(24) == 1;
45
+ constexpr bool use_padded_mask = get_compile_time_arg_val(25) == 1;
46
+ constexpr bool is_chunked = get_compile_time_arg_val(26) == 1;
47
+ constexpr uint32_t scale_fp32 = get_compile_time_arg_val(27);
48
+ constexpr uint32_t sliding_window_size = get_compile_time_arg_val(28);
49
+ constexpr bool use_attention_sink = get_compile_time_arg_val(29) == 1;
50
+ constexpr bool use_streaming_compute = get_compile_time_arg_val(30) == 1;
51
+ constexpr uint32_t valid_Skt = get_compile_time_arg_val(31);
52
+ constexpr uint32_t k_partial_col = get_compile_time_arg_val(32);
53
+ // Zigzag remap flag drives the external remap_q_index call on the flat B*NQH*q_num_chunks range.
54
+ constexpr bool use_zigzag_balancing = get_compile_time_arg_val(33) == 1;
55
+
56
+ const uint32_t core_id = get_arg_val<uint32_t>(0);
57
+ const uint32_t num_phases = get_arg_val<uint32_t>(1);
58
+ const uint32_t use_chunk_start_idx_tensor = get_arg_val<uint32_t>(2);
59
+ uint32_t chunked_q_chunk_offset_phase_1 = get_arg_val<uint32_t>(3);
60
+ uint32_t chunked_q_chunk_offset_phase_2 = 0;
61
+ if (num_phases == 2) {
62
+ chunked_q_chunk_offset_phase_2 = get_arg_val<uint32_t>(4);
63
+ }
64
+
65
+ // Global Q scheduling args follow phase_2 slot.
66
+ const uint32_t global_q_start = get_arg_val<uint32_t>(5);
67
+ const uint32_t global_q_count = get_arg_val<uint32_t>(6);
68
+
69
+ constexpr uint32_t q_chunk_tiles = Sq_chunk_t * DHt;
70
+ constexpr uint32_t k_chunk_tiles = Sk_chunk_t * DHt;
71
+ constexpr uint32_t v_chunk_tiles = Sk_chunk_t * vDHt;
72
+ constexpr uint32_t qk_chunk_tiles = Sq_chunk_t * Sk_chunk_t;
73
+ constexpr uint32_t out_chunk_tiles = Sq_chunk_t * vDHt;
74
+
75
+ constexpr uint32_t cb_arg_offset = 34;
76
+ constexpr uint32_t cb_q_in = get_compile_time_arg_val(cb_arg_offset + 0);
77
+ constexpr uint32_t cb_k_in = get_compile_time_arg_val(cb_arg_offset + 1);
78
+ constexpr uint32_t cb_v_in = get_compile_time_arg_val(cb_arg_offset + 2);
79
+ constexpr uint32_t cb_mask_in = get_compile_time_arg_val(cb_arg_offset + 3);
80
+ constexpr uint32_t cb_attention_sink = get_compile_time_arg_val(cb_arg_offset + 4);
81
+ constexpr uint32_t cb_identity_scale_in = get_compile_time_arg_val(cb_arg_offset + 5);
82
+ constexpr uint32_t cb_col_identity = get_compile_time_arg_val(cb_arg_offset + 6);
83
+ constexpr uint32_t cb_chunk_start_idx = get_compile_time_arg_val(cb_arg_offset + 7);
84
+ constexpr uint32_t cb_recip_scratch = get_compile_time_arg_val(cb_arg_offset + 8);
85
+ constexpr uint32_t cb_out = get_compile_time_arg_val(cb_arg_offset + 9);
86
+ constexpr uint32_t cb_qk_im = get_compile_time_arg_val(cb_arg_offset + 10);
87
+ constexpr uint32_t cb_out_im_A = get_compile_time_arg_val(cb_arg_offset + 11);
88
+ constexpr uint32_t cb_out_im_B = get_compile_time_arg_val(cb_arg_offset + 12);
89
+ constexpr uint32_t cb_max_A = get_compile_time_arg_val(cb_arg_offset + 13);
90
+ constexpr uint32_t cb_max_B = get_compile_time_arg_val(cb_arg_offset + 14);
91
+ constexpr uint32_t cb_sum_A = get_compile_time_arg_val(cb_arg_offset + 15);
92
+ constexpr uint32_t cb_sum_B = get_compile_time_arg_val(cb_arg_offset + 16);
93
+ constexpr uint32_t cb_exp_max_diff = get_compile_time_arg_val(cb_arg_offset + 17);
94
+ uint32_t chunked_q_chunk_offset = 0;
95
+ CircularBuffer cb_chunk_start_idx_obj(cb_chunk_start_idx);
96
+ CircularBuffer cb_identity_scale_in_obj(cb_identity_scale_in);
97
+ CircularBuffer cb_mask_in_obj(cb_mask_in);
98
+ compute_kernel_hw_startup<SrcOrder::Reverse>(cb_q_in, cb_k_in, cb_out);
99
+ matmul_init(cb_q_in, cb_k_in);
100
+
101
+ if constexpr (is_chunked) {
102
+ if (use_chunk_start_idx_tensor != 0) {
103
+ cb_chunk_start_idx_obj.wait_front(1);
104
+ uint32_t chunk_start_idx = ckernel::read_tile_value(cb_chunk_start_idx, 0, 0);
105
+ cb_chunk_start_idx_obj.pop_front(1);
106
+ const uint32_t q_chunk_size = Sq_chunk_t * TILE_HEIGHT;
107
+ chunked_q_chunk_offset_phase_1 = chunk_start_idx / q_chunk_size;
108
+ if (num_phases == 2) {
109
+ chunked_q_chunk_offset_phase_2 = chunked_q_chunk_offset_phase_1;
110
+ }
111
+ }
112
+ }
113
+
114
+ if constexpr (use_streaming_compute) {
115
+ // Streaming SDPA v2: direct cb_qkt_im writes via cb_push_back_hold_wr_ptr.
116
+ // No row buffers needed; a dedicated 1-tile CB is used as recip scratch.
117
+
118
+ // Wait once for identity scale; v2 removes per-call waits inside reduce_c_row_group
119
+ cb_identity_scale_in_obj.wait_front(1);
120
+
121
+ // Lightweight-mask context: writer pre-generates either [neginf, causal_diag, partial?]
122
+ // or, for sliding, [neginf, trailing_primary, leading_prev, leading_current, trailing_next, partial?].
123
+ // primary_diag_tile_idx is the per-layout tile used for the row-local diagonal stamp.
124
+ LightweightMaskContext lw_mask;
125
+ uint32_t lw_mask_tile_count = 1;
126
+ lw_mask.neginf_tile_idx = 0;
127
+ lw_mask.is_causal = is_causal;
128
+ if constexpr (sliding_window_size > 0) {
129
+ lw_mask.primary_diag_tile_idx = 1;
130
+ lw_mask.sliding_leading_prev_tile_idx = 2;
131
+ lw_mask.sliding_leading_tile_idx = 3;
132
+ lw_mask.sliding_trailing_next_tile_idx = 4;
133
+ lw_mask_tile_count = 5;
134
+ } else if constexpr (is_causal) {
135
+ lw_mask.causal_diag_tile_idx = lw_mask_tile_count++;
136
+ lw_mask.primary_diag_tile_idx = lw_mask.causal_diag_tile_idx;
137
+ }
138
+ if constexpr (k_partial_col > 0) {
139
+ lw_mask.global_n_partial_col = k_partial_col;
140
+ lw_mask.global_n_partial_tile_idx = lw_mask_tile_count++;
141
+ // global_n_padded_tiles = Sk_chunk_t - valid_tiles_in_last_chunk
142
+ constexpr uint32_t last_chunk_first_tile =
143
+ (valid_Skt > Sk_chunk_t) ? ((valid_Skt - 1) / Sk_chunk_t) * Sk_chunk_t : 0u;
144
+ constexpr uint32_t valid_tiles_in_last_chunk = valid_Skt - last_chunk_first_tile;
145
+ lw_mask.global_n_padded_tiles = Sk_chunk_t - valid_tiles_in_last_chunk;
146
+ }
147
+ // A user-provided dense mask is streamed per-chunk by the reader and consumed inside the
148
+ // inner loop — it does not use the writer-generated lightweight palette, so skip this wait.
149
+ if constexpr ((is_causal || sliding_window_size > 0 || k_partial_col > 0) && !use_provided_mask) {
150
+ cb_mask_in_obj.wait_front(lw_mask_tile_count);
151
+ }
152
+
153
+ // Global Q scheduling: sdpa_standard_v2 walks the per-core flat range over
154
+ // B*NQH*q_num_chunks chunks; the modulo inside its inner loop extracts the per-head q_chunk
155
+ // from each flat index. num_phases==1 is pinned for streaming, so the chunked offset comes
156
+ // from phase 1.
157
+ sdpa_standard_v2<
158
+ Sq_chunk_t,
159
+ Sk_chunk_t,
160
+ valid_Skt,
161
+ DHt,
162
+ vDHt,
163
+ scale_fp32,
164
+ qk_subblock_h,
165
+ qk_subblock_w,
166
+ out_subblock_h,
167
+ out_subblock_w,
168
+ use_padded_mask,
169
+ cb_q_in,
170
+ cb_k_in,
171
+ cb_v_in,
172
+ cb_qk_im,
173
+ cb_identity_scale_in,
174
+ cb_exp_max_diff,
175
+ cb_col_identity,
176
+ cb_recip_scratch,
177
+ cb_out, // normalized output goes directly to output CB
178
+ cb_mask_in,
179
+ sliding_window_size,
180
+ is_causal,
181
+ use_attention_sink,
182
+ cb_attention_sink,
183
+ use_provided_mask>(
184
+ global_q_count,
185
+ k_num_chunks,
186
+ cb_out_im_A,
187
+ cb_out_im_B,
188
+ cb_max_A,
189
+ cb_max_B,
190
+ cb_sum_A,
191
+ cb_sum_B,
192
+ global_q_start,
193
+ chunked_q_chunk_offset_phase_1,
194
+ lw_mask,
195
+ q_num_chunks,
196
+ use_zigzag_balancing);
197
+ } else {
198
+ // Standard SDPA path (causal, masked, chunked, etc.)
199
+ constexpr bool use_lightweight_causal_mask = is_causal && !use_provided_mask && (sliding_window_size == 0);
200
+
201
+ LightweightMaskContext lw_mask;
202
+ if constexpr (use_lightweight_causal_mask) {
203
+ lw_mask.is_causal = true;
204
+ lw_mask.neginf_tile_idx = 0;
205
+ lw_mask.causal_diag_tile_idx = 1;
206
+ lw_mask.primary_diag_tile_idx = lw_mask.causal_diag_tile_idx;
207
+ cb_mask_in_obj.wait_front(2);
208
+ }
209
+
210
+ for (uint32_t phase = 0; phase < num_phases; ++phase) {
211
+ if (phase == 0) {
212
+ chunked_q_chunk_offset = chunked_q_chunk_offset_phase_1;
213
+ } else {
214
+ chunked_q_chunk_offset = chunked_q_chunk_offset_phase_2;
215
+ }
216
+
217
+ // Global Q scheduling: sdpa_standard walks the per-core flat range over
218
+ // B*NQH*q_num_chunks chunks; the modulo inside its inner loop extracts the per-head
219
+ // q_chunk from each flat index.
220
+ sdpa_standard<
221
+ cb_qk_im,
222
+ cb_identity_scale_in,
223
+ cb_attention_sink,
224
+ Sq_chunk_t,
225
+ Sk_chunk_t,
226
+ DHt,
227
+ vDHt,
228
+ use_attention_sink,
229
+ is_causal,
230
+ use_provided_mask,
231
+ use_padded_mask,
232
+ is_chunked,
233
+ scale_fp32,
234
+ sliding_window_size,
235
+ use_lightweight_causal_mask>(
236
+ Skt,
237
+ qk_in0_block_w,
238
+ qk_subblock_w,
239
+ qk_subblock_h,
240
+ qk_in0_num_subblocks,
241
+ qk_in1_num_subblocks,
242
+ qk_num_blocks,
243
+ out_in0_block_w,
244
+ out_subblock_w,
245
+ out_subblock_h,
246
+ out_in0_num_subblocks,
247
+ out_in1_num_subblocks,
248
+ out_num_blocks,
249
+ /*iter_q_start=*/0,
250
+ /*iter_q_end=*/global_q_count,
251
+ q_num_chunks,
252
+ /*local_q_start=*/global_q_start,
253
+ chunked_q_chunk_offset,
254
+ k_num_chunks,
255
+ q_chunk_tiles,
256
+ k_chunk_tiles,
257
+ v_chunk_tiles,
258
+ qk_chunk_tiles,
259
+ out_chunk_tiles,
260
+ cb_q_in,
261
+ cb_k_in,
262
+ cb_v_in,
263
+ cb_mask_in,
264
+ cb_col_identity,
265
+ cb_out_im_A,
266
+ cb_out_im_B,
267
+ cb_max_A,
268
+ cb_max_B,
269
+ cb_sum_A,
270
+ cb_sum_B,
271
+ cb_exp_max_diff,
272
+ cb_out,
273
+ lw_mask,
274
+ use_zigzag_balancing);
275
+ }
276
+ }
277
+ }
code/models/demos/mast3r/tt/kernels/heads_concat_reader.cpp ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): concatenate heads, reader.
3
+ // Unit u = (b, s, h) over [B, St, H]: reads the Wt tiles of head h, seq-tile s, batch b of the [B, H, S, Wt*32]
4
+ // TILE input (tile ((b*H + h)*St + s)*Wt) into CB 0. Units are processed in blocks of UB to batch NoC barriers.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ void kernel_main() {
9
+ constexpr uint32_t H = get_compile_time_arg_val(0);
10
+ constexpr uint32_t St = get_compile_time_arg_val(1);
11
+ constexpr uint32_t Wt = get_compile_time_arg_val(2);
12
+ constexpr uint32_t UB = get_compile_time_arg_val(3);
13
+ constexpr uint32_t LOG2_TILE = 11;
14
+ constexpr uint32_t tile_bytes = 1u << LOG2_TILE;
15
+ const uint32_t in_addr = get_arg_val<uint32_t>(0);
16
+ const uint32_t u_start = get_arg_val<uint32_t>(1);
17
+ const uint32_t u_count = get_arg_val<uint32_t>(2);
18
+ const InterleavedPow2AddrGen<IN_DRAM> g = {.bank_base_address = in_addr, .log_base_2_of_page_size = LOG2_TILE};
19
+ const uint32_t u_end = u_start + u_count;
20
+ for (uint32_t u0 = u_start; u0 < u_end; u0 += UB) {
21
+ const uint32_t n = (u_end - u0) < UB ? (u_end - u0) : UB;
22
+ cb_reserve_back(0, n * Wt);
23
+ uint32_t w = get_write_ptr(0);
24
+ for (uint32_t u = u0; u < u0 + n; ++u) {
25
+ const uint32_t b = u / (St * H);
26
+ const uint32_t r = u - b * (St * H);
27
+ const uint32_t s = r / H;
28
+ const uint32_t h = r - s * H;
29
+ const uint32_t it = ((b * H + h) * St + s) * Wt;
30
+ for (uint32_t j = 0; j < Wt; ++j) {
31
+ noc_async_read(g.get_noc_addr(it + j), w, tile_bytes);
32
+ w += tile_bytes;
33
+ }
34
+ }
35
+ noc_async_read_barrier();
36
+ cb_push_back(0, n * Wt);
37
+ }
38
+ }
code/models/demos/mast3r/tt/kernels/heads_concat_writer.cpp ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): concatenate heads, writer.
3
+ // Writes unit (b, s, h) to output tensor of batch b at tile out_base[b] + s*out_w[b] + h*Wt.
4
+ #include <stdint.h>
5
+ #include "api/dataflow/dataflow_api.h"
6
+
7
+ void kernel_main() {
8
+ constexpr uint32_t H = get_compile_time_arg_val(0);
9
+ constexpr uint32_t St = get_compile_time_arg_val(1);
10
+ constexpr uint32_t Wt = get_compile_time_arg_val(2);
11
+ constexpr uint32_t UB = get_compile_time_arg_val(3);
12
+ constexpr uint32_t NB = get_compile_time_arg_val(4);
13
+ constexpr uint32_t LOG2_TILE = 11;
14
+ constexpr uint32_t tile_bytes = 1u << LOG2_TILE;
15
+ uint32_t a = 0;
16
+ const uint32_t u_start = get_arg_val<uint32_t>(a++);
17
+ const uint32_t u_count = get_arg_val<uint32_t>(a++);
18
+ uint32_t o_addr[NB], o_base[NB], o_w[NB];
19
+ for (uint32_t b = 0; b < NB; ++b) {
20
+ o_addr[b] = get_arg_val<uint32_t>(a++); o_base[b] = get_arg_val<uint32_t>(a++); o_w[b] = get_arg_val<uint32_t>(a++);
21
+ }
22
+ const uint32_t u_end = u_start + u_count;
23
+ for (uint32_t u0 = u_start; u0 < u_end; u0 += UB) {
24
+ const uint32_t n = (u_end - u0) < UB ? (u_end - u0) : UB;
25
+ cb_wait_front(0, n * Wt);
26
+ uint32_t rd = get_read_ptr(0);
27
+ for (uint32_t u = u0; u < u0 + n; ++u) {
28
+ const uint32_t b = u / (St * H);
29
+ const uint32_t r = u - b * (St * H);
30
+ const uint32_t s = r / H;
31
+ const uint32_t h = r - s * H;
32
+ const InterleavedPow2AddrGen<OUT_DRAM> g = {.bank_base_address = o_addr[b], .log_base_2_of_page_size = LOG2_TILE};
33
+ const uint32_t ot = o_base[b] + s * o_w[b] + h * Wt;
34
+ for (uint32_t j = 0; j < Wt; ++j) {
35
+ noc_async_write(rd, g.get_noc_addr(ot + j), tile_bytes);
36
+ rd += tile_bytes;
37
+ }
38
+ }
39
+ noc_async_writes_flushed();
40
+ cb_pop_front(0, n * Wt);
41
+ }
42
+ noc_async_write_barrier();
43
+ }
code/models/demos/mast3r/tt/kernels/heads_rope_compute.cpp ADDED
@@ -0,0 +1,180 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): fused split-heads + RoPE, compute.
3
+ // Per unit, for q then k (Wt tiles each): the same op sequence as ttnn's rotary_embedding_llama compute kernel
4
+ // (rotated = x @ trans_mat; sin_i = rotated * sin; cos_i = x * cos; out = cos_i + sin_i, each packed to bf16),
5
+ // so the result is bit-identical to split_heads + rotary_embedding_llama with the same compute config.
6
+ #include <cstdint>
7
+ #include "api/compute/common.h"
8
+ #include "api/compute/eltwise_binary.h"
9
+ #include "api/compute/bcast.h"
10
+ #include "api/compute/matmul.h"
11
+ #include "api/compute/compute_kernel_hw_startup.h"
12
+ #include "api/compute/reconfig_data_format.h"
13
+ #include "api/compute/pack.h"
14
+ #include "api/compute/cb_api.h"
15
+
16
+ constexpr uint32_t CB_IN = 0, CB_COS = 1, CB_SIN = 2, CB_TRANS = 3;
17
+ constexpr uint32_t CB_OUT = 16, CB_ROT = 24, CB_COSI = 25, CB_SINI = 26;
18
+
19
+ #ifdef D_NOCOMP
20
+ #define MATH_OP(x) ((void)0)
21
+ #else
22
+ #define MATH_OP(x) x
23
+ #endif
24
+ ALWI void ACQ() { tile_regs_acquire(); tile_regs_wait(); }
25
+ ALWI void REL() { tile_regs_commit(); tile_regs_release(); }
26
+
27
+ void kernel_main() {
28
+ constexpr uint32_t H = get_compile_time_arg_val(0);
29
+ constexpr uint32_t St = get_compile_time_arg_val(1);
30
+ constexpr uint32_t Wt = get_compile_time_arg_val(2);
31
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
32
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
33
+ if (u_count == 0) {
34
+ return;
35
+ }
36
+ compute_kernel_hw_startup<SrcOrder::Reverse>(CB_IN, CB_TRANS, CB_OUT);
37
+ compute_kernel_hw_startup(CB_COSI, CB_SINI, CB_OUT);
38
+ cb_wait_front(CB_TRANS, 1);
39
+ uint32_t last_s = 0xFFFFFFFF;
40
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
41
+ const uint32_t b = u / (St * H);
42
+ const uint32_t r = u - b * (St * H);
43
+ const uint32_t s = r / H;
44
+ if (s != last_s) {
45
+ if (last_s != 0xFFFFFFFF) {
46
+ cb_pop_front(CB_SIN, Wt);
47
+ cb_pop_front(CB_COS, Wt);
48
+ }
49
+ cb_wait_front(CB_SIN, Wt);
50
+ cb_wait_front(CB_COS, Wt);
51
+ last_s = s;
52
+ }
53
+ #ifdef HROPE_QK2
54
+ {
55
+ // MAST3R_OPT "hrqk": q and k of the unit in ONE pass per stage (2*Wt tiles per dest acquire; fp32 dest
56
+ // half = 4 tiles = 2*Wt for dh = 64). Same per-tile op sequence -> bit-identical.
57
+ constexpr uint32_t NT = 2 * Wt;
58
+ static_assert(NT <= 4, "q + k of one unit must fit one fp32 dest half (4 tiles)");
59
+ cb_wait_front(CB_IN, NT);
60
+ cb_reserve_back(CB_ROT, NT);
61
+ cb_reserve_back(CB_SINI, NT);
62
+ cb_reserve_back(CB_COSI, NT);
63
+ cb_reserve_back(CB_OUT, NT);
64
+
65
+ reconfig_data_format(CB_COSI, CB_TRANS, CB_SINI, CB_IN);
66
+ pack_reconfig_data_format(CB_OUT, CB_ROT);
67
+ matmul_init(CB_IN, CB_TRANS);
68
+ ACQ();
69
+ for (uint32_t j = 0; j < NT; ++j) {
70
+ MATH_OP(matmul_tiles(CB_IN, CB_TRANS, j, 0, j));
71
+ pack_tile(j, CB_ROT, j);
72
+ }
73
+ REL();
74
+ cb_push_back(CB_ROT, NT);
75
+ cb_wait_front(CB_ROT, NT);
76
+
77
+ reconfig_data_format(CB_TRANS, CB_ROT, CB_IN, CB_SIN);
78
+ pack_reconfig_data_format(CB_ROT, CB_SINI);
79
+ mul_init(CB_ROT, CB_SIN);
80
+ ACQ();
81
+ for (uint32_t j = 0; j < NT; ++j) {
82
+ MATH_OP(mul_tiles(CB_ROT, CB_SIN, j, j % Wt, j));
83
+ pack_tile(j, CB_SINI, j);
84
+ }
85
+ REL();
86
+ cb_push_back(CB_SINI, NT);
87
+ cb_pop_front(CB_ROT, NT);
88
+
89
+ reconfig_data_format(CB_ROT, CB_IN, CB_SIN, CB_COS);
90
+ pack_reconfig_data_format(CB_SINI, CB_COSI);
91
+ mul_init(CB_IN, CB_COS);
92
+ ACQ();
93
+ for (uint32_t j = 0; j < NT; ++j) {
94
+ MATH_OP(mul_tiles(CB_IN, CB_COS, j, j % Wt, j));
95
+ pack_tile(j, CB_COSI, j);
96
+ }
97
+ REL();
98
+ cb_push_back(CB_COSI, NT);
99
+ cb_pop_front(CB_IN, NT);
100
+
101
+ cb_wait_front(CB_SINI, NT);
102
+ cb_wait_front(CB_COSI, NT);
103
+ reconfig_data_format(CB_IN, CB_COSI, CB_COS, CB_SINI);
104
+ pack_reconfig_data_format(CB_COSI, CB_OUT);
105
+ add_init(CB_COSI, CB_SINI);
106
+ ACQ();
107
+ for (uint32_t j = 0; j < NT; ++j) {
108
+ MATH_OP(add_tiles(CB_COSI, CB_SINI, j, j, j));
109
+ pack_tile(j, CB_OUT, j);
110
+ }
111
+ REL();
112
+ cb_push_back(CB_OUT, NT);
113
+ cb_pop_front(CB_SINI, NT);
114
+ cb_pop_front(CB_COSI, NT);
115
+ }
116
+ #else
117
+ for (uint32_t t = 0; t < 2; ++t) { // q, k
118
+ cb_wait_front(CB_IN, Wt);
119
+ cb_reserve_back(CB_ROT, Wt);
120
+ cb_reserve_back(CB_SINI, Wt);
121
+ cb_reserve_back(CB_COSI, Wt);
122
+ cb_reserve_back(CB_OUT, Wt);
123
+
124
+ reconfig_data_format(CB_COSI, CB_TRANS, CB_SINI, CB_IN);
125
+ pack_reconfig_data_format(CB_OUT, CB_ROT);
126
+ matmul_init(CB_IN, CB_TRANS);
127
+ ACQ();
128
+ for (uint32_t j = 0; j < Wt; ++j) {
129
+ MATH_OP(matmul_tiles(CB_IN, CB_TRANS, j, 0, j));
130
+ pack_tile(j, CB_ROT, j);
131
+ }
132
+ REL();
133
+ cb_push_back(CB_ROT, Wt);
134
+ cb_wait_front(CB_ROT, Wt);
135
+
136
+ reconfig_data_format(CB_TRANS, CB_ROT, CB_IN, CB_SIN);
137
+ pack_reconfig_data_format(CB_ROT, CB_SINI);
138
+ mul_init(CB_ROT, CB_SIN);
139
+ ACQ();
140
+ for (uint32_t j = 0; j < Wt; ++j) {
141
+ MATH_OP(mul_tiles(CB_ROT, CB_SIN, j, j, j));
142
+ pack_tile(j, CB_SINI, j);
143
+ }
144
+ REL();
145
+ cb_push_back(CB_SINI, Wt);
146
+ cb_pop_front(CB_ROT, Wt);
147
+
148
+ reconfig_data_format(CB_ROT, CB_IN, CB_SIN, CB_COS);
149
+ pack_reconfig_data_format(CB_SINI, CB_COSI);
150
+ mul_init(CB_IN, CB_COS);
151
+ ACQ();
152
+ for (uint32_t j = 0; j < Wt; ++j) {
153
+ MATH_OP(mul_tiles(CB_IN, CB_COS, j, j, j));
154
+ pack_tile(j, CB_COSI, j);
155
+ }
156
+ REL();
157
+ cb_push_back(CB_COSI, Wt);
158
+ cb_pop_front(CB_IN, Wt);
159
+
160
+ cb_wait_front(CB_SINI, Wt);
161
+ cb_wait_front(CB_COSI, Wt);
162
+ reconfig_data_format(CB_IN, CB_COSI, CB_COS, CB_SINI);
163
+ pack_reconfig_data_format(CB_COSI, CB_OUT);
164
+ add_init(CB_COSI, CB_SINI);
165
+ ACQ();
166
+ for (uint32_t j = 0; j < Wt; ++j) {
167
+ MATH_OP(add_tiles(CB_COSI, CB_SINI, j, j, j));
168
+ pack_tile(j, CB_OUT, j);
169
+ }
170
+ REL();
171
+ cb_push_back(CB_OUT, Wt);
172
+ cb_pop_front(CB_SINI, Wt);
173
+ cb_pop_front(CB_COSI, Wt);
174
+ }
175
+ #endif
176
+ }
177
+ cb_pop_front(CB_SIN, Wt);
178
+ cb_pop_front(CB_COS, Wt);
179
+ cb_pop_front(CB_TRANS, 1);
180
+ }
code/models/demos/mast3r/tt/kernels/heads_rope_reader.cpp ADDED
@@ -0,0 +1,110 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): fused split-heads + RoPE, reader.
3
+ // Work unit u = (b, s, h) in row-major order over [B, St, H]. For each unit it reads the Wt q tiles and
4
+ // Wt k tiles of head h at seq-tile s of batch b (from per-batch source tensors with arbitrary row width and
5
+ // column offset) into CB_IN, the Wt v tiles into CB_V, and (when s changes) the cos/sin rows into CB_COS/CB_SIN.
6
+ #include <stdint.h>
7
+ #include "api/dataflow/dataflow_api.h"
8
+
9
+ constexpr uint32_t CB_IN = 0, CB_COS = 1, CB_SIN = 2, CB_TRANS = 3, CB_V = 4;
10
+
11
+ void kernel_main() {
12
+ constexpr uint32_t H = get_compile_time_arg_val(0);
13
+ constexpr uint32_t St = get_compile_time_arg_val(1);
14
+ constexpr uint32_t Wt = get_compile_time_arg_val(2);
15
+ constexpr uint32_t NB = get_compile_time_arg_val(3); // batch entries described in the rt args
16
+ constexpr uint32_t LOG2_TILE = 11; // bf16 32x32 tile = 2048 B
17
+
18
+ // per-core runtime args: unit range; common runtime args (same for every core): addresses / source layout
19
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
20
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
21
+ uint32_t a = 0;
22
+ const uint32_t cos_addr = get_common_arg_val<uint32_t>(a++);
23
+ const uint32_t sin_addr = get_common_arg_val<uint32_t>(a++);
24
+ const uint32_t trans_addr = get_common_arg_val<uint32_t>(a++);
25
+ // per batch entry and role (q, k, v): address, tile base, row width (tiles), source kind
26
+ // (kind 0 = interleaved, 1 / 2 = the 1st / 2nd TensorAccessor layout in the compile-time args: MAST3R_OPT "obs",
27
+ // block-sharded matmul outputs)
28
+ uint32_t q_addr[NB], q_base[NB], q_w[NB], k_addr[NB], k_base[NB], k_w[NB], v_addr[NB], v_base[NB], v_w[NB];
29
+ uint32_t q_kind[NB], k_kind[NB], v_kind[NB];
30
+ for (uint32_t b = 0; b < NB; ++b) {
31
+ q_addr[b] = get_common_arg_val<uint32_t>(a++); q_base[b] = get_common_arg_val<uint32_t>(a++); q_w[b] = get_common_arg_val<uint32_t>(a++);
32
+ q_kind[b] = get_common_arg_val<uint32_t>(a++);
33
+ k_addr[b] = get_common_arg_val<uint32_t>(a++); k_base[b] = get_common_arg_val<uint32_t>(a++); k_w[b] = get_common_arg_val<uint32_t>(a++);
34
+ k_kind[b] = get_common_arg_val<uint32_t>(a++);
35
+ v_addr[b] = get_common_arg_val<uint32_t>(a++); v_base[b] = get_common_arg_val<uint32_t>(a++); v_w[b] = get_common_arg_val<uint32_t>(a++);
36
+ v_kind[b] = get_common_arg_val<uint32_t>(a++);
37
+ }
38
+ if (u_count == 0) {
39
+ return;
40
+ }
41
+ #if NTA >= 1
42
+ constexpr auto ta1_args = TensorAccessorArgs<4>();
43
+ #endif
44
+ #if NTA >= 2
45
+ constexpr auto ta2_args = TensorAccessorArgs<ta1_args.next_compile_time_args_offset()>();
46
+ #endif
47
+ auto src_noc = [&](uint32_t kind, uint32_t addr, uint32_t tile) -> uint64_t {
48
+ #if NTA >= 1
49
+ if (kind == 1) {
50
+ const auto ta = TensorAccessor(ta1_args, addr, 1u << LOG2_TILE);
51
+ return ta.get_noc_addr(tile);
52
+ }
53
+ #endif
54
+ #if NTA >= 2
55
+ if (kind == 2) {
56
+ const auto ta = TensorAccessor(ta2_args, addr, 1u << LOG2_TILE);
57
+ return ta.get_noc_addr(tile);
58
+ }
59
+ #endif
60
+ const InterleavedPow2AddrGen<SRC_DRAM> g = {.bank_base_address = addr, .log_base_2_of_page_size = LOG2_TILE};
61
+ return g.get_noc_addr(tile);
62
+ };
63
+ const InterleavedPow2AddrGen<COS_DRAM> cos_g = {.bank_base_address = cos_addr, .log_base_2_of_page_size = LOG2_TILE};
64
+ const InterleavedPow2AddrGen<COS_DRAM> sin_g = {.bank_base_address = sin_addr, .log_base_2_of_page_size = LOG2_TILE};
65
+ const InterleavedPow2AddrGen<TRANS_DRAM> tr_g = {.bank_base_address = trans_addr, .log_base_2_of_page_size = LOG2_TILE};
66
+ constexpr uint32_t tile_bytes = 1u << LOG2_TILE;
67
+
68
+ cb_reserve_back(CB_TRANS, 1);
69
+ noc_async_read(tr_g.get_noc_addr(0), get_write_ptr(CB_TRANS), tile_bytes);
70
+ noc_async_read_barrier();
71
+ cb_push_back(CB_TRANS, 1);
72
+
73
+ uint32_t last_s = 0xFFFFFFFF;
74
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
75
+ const uint32_t b = u / (St * H);
76
+ const uint32_t r = u - b * (St * H);
77
+ const uint32_t s = r / H;
78
+ const uint32_t h = r - s * H;
79
+ if (s != last_s) {
80
+ cb_reserve_back(CB_COS, Wt);
81
+ cb_reserve_back(CB_SIN, Wt);
82
+ uint32_t cw = get_write_ptr(CB_COS), sw = get_write_ptr(CB_SIN);
83
+ for (uint32_t j = 0; j < Wt; ++j) {
84
+ noc_async_read(cos_g.get_noc_addr(s * Wt + j), cw + j * tile_bytes, tile_bytes);
85
+ noc_async_read(sin_g.get_noc_addr(s * Wt + j), sw + j * tile_bytes, tile_bytes);
86
+ }
87
+ }
88
+ cb_reserve_back(CB_IN, 2 * Wt);
89
+ cb_reserve_back(CB_V, Wt);
90
+ uint32_t iw = get_write_ptr(CB_IN), vw = get_write_ptr(CB_V);
91
+ const uint32_t qt = q_base[b] + s * q_w[b] + h * Wt;
92
+ const uint32_t kt = k_base[b] + s * k_w[b] + h * Wt;
93
+ const uint32_t vt = v_base[b] + s * v_w[b] + h * Wt;
94
+ #ifndef D_NOREAD
95
+ for (uint32_t j = 0; j < Wt; ++j) {
96
+ noc_async_read(src_noc(q_kind[b], q_addr[b], qt + j), iw + j * tile_bytes, tile_bytes);
97
+ noc_async_read(src_noc(k_kind[b], k_addr[b], kt + j), iw + (Wt + j) * tile_bytes, tile_bytes);
98
+ noc_async_read(src_noc(v_kind[b], v_addr[b], vt + j), vw + j * tile_bytes, tile_bytes);
99
+ }
100
+ #endif
101
+ noc_async_read_barrier();
102
+ if (s != last_s) {
103
+ cb_push_back(CB_COS, Wt);
104
+ cb_push_back(CB_SIN, Wt);
105
+ last_s = s;
106
+ }
107
+ cb_push_back(CB_IN, 2 * Wt);
108
+ cb_push_back(CB_V, Wt);
109
+ }
110
+ }
code/models/demos/mast3r/tt/kernels/heads_rope_writer.cpp ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): fused split-heads + RoPE, writer.
3
+ // Writes rotated q (Wt tiles), rotated k (Wt tiles) from CB_OUT and v (Wt tiles) from CB_V of unit (b, s, h)
4
+ // to the [B, H, St*32, Wt*32] TILE outputs at tile ((b*H + h)*St + s)*Wt.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ constexpr uint32_t CB_V = 4, CB_OUT = 16;
9
+ #ifdef D_NOWRITE
10
+ #define WR(a, b, c) ((void)(b))
11
+ #else
12
+ #define WR(a, b, c) noc_async_write(a, b, c)
13
+ #endif
14
+
15
+ void kernel_main() {
16
+ constexpr uint32_t H = get_compile_time_arg_val(0);
17
+ constexpr uint32_t St = get_compile_time_arg_val(1);
18
+ constexpr uint32_t Wt = get_compile_time_arg_val(2);
19
+ constexpr uint32_t LOG2_TILE = 11;
20
+ constexpr uint32_t tile_bytes = 1u << LOG2_TILE;
21
+ const uint32_t q_addr = get_common_arg_val<uint32_t>(0);
22
+ const uint32_t k_addr = get_common_arg_val<uint32_t>(1);
23
+ const uint32_t v_addr = get_common_arg_val<uint32_t>(2);
24
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
25
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
26
+ const InterleavedPow2AddrGen<OUT_DRAM> qg = {.bank_base_address = q_addr, .log_base_2_of_page_size = LOG2_TILE};
27
+ const InterleavedPow2AddrGen<OUT_DRAM> kg = {.bank_base_address = k_addr, .log_base_2_of_page_size = LOG2_TILE};
28
+ const InterleavedPow2AddrGen<OUT_DRAM> vg = {.bank_base_address = v_addr, .log_base_2_of_page_size = LOG2_TILE};
29
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
30
+ const uint32_t b = u / (St * H);
31
+ const uint32_t r = u - b * (St * H);
32
+ const uint32_t s = r / H;
33
+ const uint32_t h = r - s * H;
34
+ const uint32_t ot = ((b * H + h) * St + s) * Wt;
35
+ cb_wait_front(CB_V, Wt);
36
+ uint32_t rd = get_read_ptr(CB_V);
37
+ for (uint32_t j = 0; j < Wt; ++j) {
38
+ WR(rd + j * tile_bytes, vg.get_noc_addr(ot + j), tile_bytes);
39
+ }
40
+ cb_wait_front(CB_OUT, Wt);
41
+ rd = get_read_ptr(CB_OUT);
42
+ for (uint32_t j = 0; j < Wt; ++j) {
43
+ WR(rd + j * tile_bytes, qg.get_noc_addr(ot + j), tile_bytes);
44
+ }
45
+ noc_async_writes_flushed();
46
+ cb_pop_front(CB_OUT, Wt);
47
+ cb_wait_front(CB_OUT, Wt);
48
+ rd = get_read_ptr(CB_OUT);
49
+ for (uint32_t j = 0; j < Wt; ++j) {
50
+ WR(rd + j * tile_bytes, kg.get_noc_addr(ot + j), tile_bytes);
51
+ }
52
+ noc_async_writes_flushed();
53
+ cb_pop_front(CB_OUT, Wt);
54
+ cb_pop_front(CB_V, Wt);
55
+ }
56
+ noc_async_write_barrier();
57
+ }
code/models/demos/mast3r/tt/kernels/ln_add_compute.cpp ADDED
@@ -0,0 +1,436 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc.
2
+ //
3
+ // SPDX-License-Identifier: Apache-2.0
4
+ //
5
+ // mast3r-p150 model-local copy of ttnn's layernorm.cpp compute kernel (tt-metal 8b98410e730), used through
6
+ // ttnn.generic_op. Only change: with FUSE_PRE_ADD the pre-add sum x = a + b is ALSO packed (bf16) into CB
7
+ // "cb_res_out", which the model-local writer stores as the new residual stream. One program replaces
8
+ // ttnn.add(residual, y) + ttnn.layer_norm(sum).
9
+
10
+ #include <cstdint>
11
+
12
+ #define BCAST_LLKOP EltwiseBinaryType::ELWMUL
13
+ #define BCAST_DIM BroadcastType::COL
14
+
15
+ #include "api/compute/compute_kernel_api.h"
16
+ #include "api/compute/bcast.h"
17
+ #include "api/compute/eltwise_binary.h"
18
+ #include "api/compute/layernorm.h"
19
+ #ifdef TILIZE_IN
20
+ #include "api/compute/tilize.h"
21
+ #endif
22
+ #ifdef UNTILIZE_OUT
23
+ #include "api/compute/pack_untilize.h"
24
+ #endif
25
+ #include <tt-metalium/constants.hpp>
26
+ #include "ttnn/operations/normalization/kernel_util/compute/numeric.h"
27
+ #include "ttnn/operations/normalization/kernel_util/generic/blocked_range.h"
28
+ #include "ttnn/operations/normalization/kernel_util/generic/bit.h"
29
+ #include "ttnn/operations/normalization/layernorm/device/kernels/layernorm_scaler_tiles.h"
30
+ #include "api/compute/eltwise_unary/sfpu_split_includes.h"
31
+ #include "api/compute/tile_move_copy.h"
32
+ #include "api/compute/eltwise_unary/eltwise_unary.h"
33
+ #include "api/dataflow/dataflow_buffer.h"
34
+ #include "ttnn/operations/normalization/layernorm/device/kernels/compute/layernorm_compute_utils.h"
35
+
36
+ namespace generic = norm::kernel_util::generic;
37
+ namespace kutil = norm::kernel_util;
38
+ namespace numeric = kutil::compute::numeric;
39
+ namespace policies = kutil::compute::policies;
40
+
41
+ constexpr uint32_t CB_RES_OUT = 17;
42
+
43
+ void kernel_main() {
44
+ uint32_t NCHt = get_arg_val<uint32_t>(0);
45
+ constexpr uint32_t Wt = get_compile_time_arg_val(0);
46
+ constexpr uint32_t block_size = get_compile_time_arg_val(1);
47
+ constexpr uint32_t do_gamma = get_compile_time_arg_val(2);
48
+ constexpr uint32_t do_beta = get_compile_time_arg_val(3);
49
+ constexpr bool FLOAT32_DTYPE = get_compile_time_arg_val(4) == 1;
50
+ constexpr bool FLOAT32_REDUCTION = get_compile_time_arg_val(5) == 1;
51
+ constexpr bool LEGACY_RSQRT = get_compile_time_arg_val(6) == 1;
52
+ constexpr uint32_t W = get_compile_time_arg_val(7);
53
+ constexpr uint32_t tile_width = get_compile_time_arg_val(8);
54
+
55
+ // CB indices - configurable via named compile-time args for kernel chaining support
56
+ constexpr auto dfb_scaler_id = get_named_compile_time_arg_val("cb_scaler"); // single tile generated by the reader
57
+ constexpr auto dfb_eps_id = get_named_compile_time_arg_val("cb_eps"); // single tile generated by the reader
58
+ constexpr auto dfb_in_id = get_named_compile_time_arg_val("cb_in"); // input x or a for fused pre-add (x=a+b)
59
+ constexpr auto dfb_inb_id = get_named_compile_time_arg_val("cb_inb"); // input b for fused pre-add
60
+ constexpr auto dfb_out_id = get_named_compile_time_arg_val("cb_out"); // output
61
+ constexpr auto dfb_gamma_id = get_named_compile_time_arg_val("cb_gamma");
62
+ constexpr auto dfb_beta_id = get_named_compile_time_arg_val("cb_beta");
63
+ #if defined RMSNORM and not defined FUSE_PRE_ADD
64
+ constexpr uint32_t dfb_xmm_id = dfb_in_id; // x minus mean
65
+ #else
66
+ constexpr uint32_t dfb_xmm_id = get_named_compile_time_arg_val("cb_xmm"); // x minus mean
67
+ #endif
68
+ DataflowBuffer dfb_xmm(dfb_xmm_id);
69
+ constexpr auto dfb_ex_id = get_named_compile_time_arg_val("cb_ex"); // E[x]
70
+ constexpr auto dfb_ex2_id = get_named_compile_time_arg_val("cb_ex2"); // E[(x-E[x])^2]
71
+ constexpr auto dfb_xmm2_id = get_named_compile_time_arg_val("cb_xmm2"); // xmm^2
72
+ constexpr auto dfb_ex2pe_id = get_named_compile_time_arg_val("cb_ex2pe"); // E[(x-E[x])^2]+eps
73
+ constexpr auto dfb_fusion_id = get_named_compile_time_arg_val("cb_fusion"); // stream gamma/beta
74
+ DataflowBuffer dfb_eps(dfb_eps_id);
75
+ DataflowBuffer dfb_in(dfb_in_id);
76
+ DataflowBuffer dfb_inb(dfb_inb_id);
77
+ DataflowBuffer dfb_out(dfb_out_id);
78
+ DataflowBuffer dfb_gamma(dfb_gamma_id);
79
+ DataflowBuffer dfb_beta(dfb_beta_id);
80
+ DataflowBuffer dfb_ex(dfb_ex_id);
81
+ DataflowBuffer dfb_ex2(dfb_ex2_id);
82
+ DataflowBuffer dfb_xmm2(dfb_xmm2_id);
83
+ DataflowBuffer dfb_ex2pe(dfb_ex2pe_id);
84
+ DataflowBuffer dfb_fusion(dfb_fusion_id);
85
+ DataflowBuffer dfb_scaler(dfb_scaler_id);
86
+
87
+ constexpr auto dfb_in_rm_id =
88
+ get_named_compile_time_arg_val("cb_in_rm"); // input row-major (if row-major input, otherwise unused)
89
+ DataflowBuffer dfb_in_rm(dfb_in_rm_id);
90
+
91
+ constexpr int onetile = 1;
92
+ constexpr int dst0 = 0;
93
+ constexpr int dst1 = 1;
94
+ constexpr auto scaler0 = 0;
95
+
96
+ #ifdef FUSE_PRE_ADD
97
+ #ifdef RMSNORM
98
+ constexpr uint32_t dfb_x_id = dfb_xmm_id;
99
+ #else
100
+ constexpr uint32_t dfb_x_id = get_named_compile_time_arg_val("cb_x");
101
+ #endif
102
+ #else
103
+ constexpr uint32_t dfb_x_id = dfb_in_id;
104
+ #endif
105
+ DataflowBuffer dfb_x(dfb_x_id);
106
+
107
+ #ifdef TILIZE_IN
108
+ compute_kernel_hw_startup(dfb_in_rm_id, dfb_in_rm_id, dfb_in_id);
109
+ #elif defined(FUSE_PRE_ADD)
110
+ compute_kernel_hw_startup(dfb_in_id, dfb_inb_id, dfb_x_id);
111
+ #elif defined(RMSNORM)
112
+ compute_kernel_hw_startup(dfb_xmm_id, dfb_xmm_id, dfb_xmm2_id);
113
+ #else
114
+ compute_kernel_hw_startup(dfb_x_id, dfb_scaler_id, dfb_ex_id);
115
+ #endif
116
+
117
+ dfb_eps.wait_front(1); // comes from the reader
118
+
119
+ constexpr int dfb_im_or_out_id = (do_gamma | do_beta) ? dfb_fusion_id : dfb_out_id;
120
+ DataflowBuffer dfb_im_or_out(dfb_im_or_out_id);
121
+
122
+ // Intermediate buffers need to be reserved/pushed/popped
123
+ // in full blocks
124
+ const auto total_buffer_size = generic::blocks(Wt, block_size).total_with_remainder();
125
+
126
+ for (uint32_t ncht = 0; ncht < NCHt; ncht++) {
127
+ #ifdef TILIZE_IN
128
+ tilize_all_blocks_to_cb<block_size>(dfb_in_rm, dfb_in, Wt);
129
+ // Re-init binary ops after tilize hardware reconfiguration.
130
+ #ifdef FUSE_PRE_ADD
131
+ // TODO(#52395): compute_kernel_hw_startup is a call-once API; this mid-kernel re-init (preserving the
132
+ // pre-cleanup full-init behaviour) should become a targeted DST re-arm.
133
+ compute_kernel_hw_startup(dfb_in_id, dfb_inb_id, dfb_x_id);
134
+ #elif defined(RMSNORM)
135
+ // TODO(#52395): compute_kernel_hw_startup is a call-once API; this mid-kernel re-init (preserving the
136
+ // pre-cleanup full-init behaviour) should become a targeted DST re-arm.
137
+ compute_kernel_hw_startup(dfb_xmm_id, dfb_xmm_id, dfb_xmm2_id);
138
+ #else
139
+ // TODO(#52395): compute_kernel_hw_startup is a call-once API; this mid-kernel re-init (preserving the
140
+ // pre-cleanup full-init behaviour) should become a targeted DST re-arm.
141
+ compute_kernel_hw_startup(dfb_x_id, dfb_scaler_id, dfb_ex_id);
142
+ #endif
143
+ #endif
144
+ /*
145
+ * X + Y
146
+ */
147
+ #ifdef FUSE_PRE_ADD
148
+ reconfig_data_format(dfb_in_id, dfb_inb_id);
149
+ pack_reconfig_data_format(dfb_x_id);
150
+ add_init(dfb_in_id, dfb_inb_id);
151
+ for (auto block : generic::blocks(Wt, block_size)) {
152
+ // In/inb come from the reader and need to be
153
+ // synced on full block size. Keep cb_x_id aligned
154
+ // to full block size as well so pre-add/no-pre-add
155
+ // can be handled the same way.
156
+ dfb_in.wait_front(block.full_block_size());
157
+ dfb_inb.wait_front(block.full_block_size());
158
+
159
+ tile_regs_acquire();
160
+ for (auto i : block.local()) {
161
+ add_tiles(dfb_in_id, dfb_inb_id, i, i, i);
162
+ }
163
+ tile_regs_commit();
164
+
165
+ dfb_in.pop_front(block.full_block_size());
166
+ dfb_inb.pop_front(block.full_block_size());
167
+
168
+ dfb_x.reserve_back(block.full_block_size());
169
+ cb_reserve_back(CB_RES_OUT, block.full_block_size());
170
+
171
+ tile_regs_wait();
172
+ for (auto i : block.local()) {
173
+ pack_tile(i, dfb_x_id);
174
+ }
175
+ pack_reconfig_data_format(dfb_x_id, CB_RES_OUT);
176
+ for (auto i : block.local()) {
177
+ pack_tile(i, CB_RES_OUT);
178
+ }
179
+ pack_reconfig_data_format(CB_RES_OUT, dfb_x_id);
180
+ tile_regs_release();
181
+
182
+ dfb_x.push_back(block.full_block_size()); // push the sum into the same buffer
183
+ cb_push_back(CB_RES_OUT, block.full_block_size());
184
+ }
185
+ #ifndef RMSNORM
186
+ reconfig_data_format(dfb_in_id, dfb_x_id, dfb_inb_id, dfb_scaler_id);
187
+ #else
188
+ reconfig_data_format(dfb_in_id, dfb_x_id, dfb_inb_id, dfb_x_id);
189
+ #endif
190
+ // by the end of this loop we should end up with Wt tiles in cb_x_id
191
+ #else
192
+ #ifdef RMSNORM
193
+ reconfig_data_format(dfb_in_id, dfb_in_id);
194
+ pack_reconfig_data_format(dfb_xmm2_id);
195
+ #endif
196
+ #endif
197
+
198
+ #ifndef RMSNORM
199
+ // E[x]
200
+ numeric::
201
+ row_wise_mean<PoolType::SUM, ReduceDim::REDUCE_ROW, FLOAT32_REDUCTION, policies::FullBlockWithoutPopPolicy>(
202
+ dfb_x, dfb_scaler, dfb_ex, W, Wt, block_size, tile_width);
203
+
204
+ // x - E[x]
205
+ reconfig_data_format(dfb_x_id, dfb_ex_id);
206
+ dfb_xmm.reserve_back(total_buffer_size);
207
+ sub_bcast_cols_init(dfb_x_id, dfb_ex_id);
208
+ for (auto block : generic::blocks(Wt, block_size)) {
209
+ tile_regs_acquire();
210
+ for (auto i : block.local()) {
211
+ sub_tiles_bcast_cols(dfb_x_id, dfb_ex_id, i, 0, i);
212
+ }
213
+ tile_regs_commit();
214
+
215
+ dfb_x.pop_front(block.full_block_size());
216
+
217
+ tile_regs_wait();
218
+ for (auto i : block.local()) {
219
+ pack_tile(i, dfb_xmm_id);
220
+ }
221
+ tile_regs_release();
222
+
223
+ dfb_xmm.push_back(block.full_block_size());
224
+ }
225
+ dfb_ex.pop_front(1);
226
+
227
+ // model-local: cb_x is bf16 here (the model passes a bf16 cb_x so the LN sees exactly the bf16 sum),
228
+ // so srcA must be reconfigured to the fp32 xmm format as in the non-fused path
229
+ reconfig_data_format_srca(dfb_x_id, dfb_xmm_id);
230
+ #endif
231
+
232
+ /* (x - E[x])^2
233
+ * compute temp = xmm*xmm = (x-E[x])^2
234
+ */
235
+ mul_init(dfb_xmm_id, dfb_xmm_id);
236
+ for (auto block : generic::blocks(Wt, block_size)) {
237
+ #ifndef RMSNORM
238
+ dfb_xmm.wait_front(block.start() + block.size());
239
+ #else
240
+ dfb_xmm.wait_front(block.start() + block.full_block_size());
241
+ #endif
242
+ tile_regs_acquire();
243
+ for (auto i : block.local()) {
244
+ const auto global_i = block.to_global(i);
245
+ mul_tiles(dfb_xmm_id, dfb_xmm_id, global_i, global_i, i);
246
+ }
247
+ tile_regs_commit();
248
+
249
+ dfb_xmm2.reserve_back(block.full_block_size());
250
+
251
+ tile_regs_wait();
252
+ for (auto i : block.local()) {
253
+ pack_tile(i, dfb_xmm2_id);
254
+ }
255
+ tile_regs_release();
256
+
257
+ dfb_xmm2.push_back(block.full_block_size());
258
+ }
259
+ #if defined RMSNORM and not defined FUSED_PRE_ADD
260
+ reconfig_data_format(dfb_xmm_id, dfb_xmm2_id, dfb_xmm_id, dfb_scaler_id);
261
+ #endif
262
+
263
+ // Var[x]
264
+ numeric::
265
+ row_wise_mean<PoolType::SUM, ReduceDim::REDUCE_ROW, FLOAT32_REDUCTION, policies::FullBlockWithPopPolicy>(
266
+ dfb_xmm2, dfb_scaler, dfb_ex2, W, Wt, block_size, tile_width);
267
+
268
+ // Var[x] + eps
269
+ dfb_ex2.wait_front(1);
270
+ reconfig_data_format(dfb_ex2_id, dfb_eps_id);
271
+
272
+ tile_regs_acquire();
273
+ add_init(dfb_ex2_id, dfb_eps_id);
274
+ add_tiles(dfb_ex2_id, dfb_eps_id, 0, 0, dst0);
275
+ rsqrt_tile_init<LEGACY_RSQRT>();
276
+ rsqrt_tile<LEGACY_RSQRT>(dst0);
277
+ tile_regs_commit();
278
+
279
+ dfb_ex2.pop_front(1);
280
+
281
+ dfb_ex2pe.reserve_back(1);
282
+ pack_reconfig_data_format(dfb_ex2pe_id);
283
+
284
+ tile_regs_wait();
285
+ pack_tile(dst0, dfb_ex2pe_id);
286
+ tile_regs_release();
287
+
288
+ dfb_ex2pe.push_back(1);
289
+
290
+ // (x-E[x]) / sqrt(Var[x] + eps) * gamma + beta
291
+ dfb_ex2pe.wait_front(1);
292
+ for (auto block : generic::blocks(Wt, block_size)) {
293
+ reconfig_data_format(dfb_xmm_id, dfb_ex2pe_id);
294
+ if constexpr (do_gamma == 0 && do_beta == 0) {
295
+ pack_reconfig_data_format(dfb_out_id);
296
+ } else {
297
+ pack_reconfig_data_format(dfb_fusion_id);
298
+ }
299
+ dfb_im_or_out.reserve_back(block.full_block_size());
300
+ #if defined RMSNORM and not defined FUSE_PRE_ADD
301
+ reconfig_data_format_srca(dfb_fusion_id, dfb_xmm_id);
302
+ #endif
303
+ tile_regs_acquire();
304
+ mul_bcast_cols_init(dfb_xmm_id, dfb_ex2pe_id);
305
+ for (auto i : block.local()) {
306
+ mul_tiles_bcast_cols(dfb_xmm_id, dfb_ex2pe_id, block.to_global(i), 0, i); // tile *= 1/(sum(exp(x)))
307
+ #ifdef SFPU_OP_INIT_ACTIVATION
308
+ // Activation must be applied last. If do_gamma != 0 or do_beta != 0 then
309
+ // activation will be applied after the gamma/beta multiplication/addition.
310
+ // Otherwise, we can apply the activation here.
311
+ if constexpr (!(do_gamma == 1 || do_beta == 1)) {
312
+ SFPU_OP_INIT_ACTIVATION
313
+ SFPU_OP_FUNC_ACTIVATION
314
+ }
315
+ #endif
316
+ }
317
+ tile_regs_commit();
318
+
319
+ tile_regs_wait();
320
+ for (auto i : block.local()) {
321
+ pack_tile(i, dfb_im_or_out_id); // pack either to intermediate (dfb_fusion or out0)
322
+ }
323
+ tile_regs_release();
324
+
325
+ dfb_im_or_out.push_back(
326
+ block.full_block_size()); // if no gamma/beta are provided, this will be passed on to the writer
327
+
328
+ if constexpr (!(do_gamma == 0 && do_beta == 0)) {
329
+ #if defined RMSNORM and not defined FUSE_PRE_ADD
330
+ reconfig_data_format_srca(dfb_xmm_id, dfb_fusion_id);
331
+ #endif
332
+ }
333
+
334
+ if constexpr (do_gamma) {
335
+ if constexpr (do_beta == 0) {
336
+ pack_reconfig_data_format(dfb_out_id);
337
+ }
338
+ reconfig_data_format_srcb(dfb_ex2pe_id, dfb_gamma_id);
339
+ uint32_t dfb_outg_id = do_beta ? dfb_fusion_id : dfb_out_id;
340
+ DataflowBuffer dfb_outg(dfb_outg_id);
341
+ dfb_gamma.wait_front(
342
+ block.start() + block.full_block_size()); // we don't pop, TODO: only wait on first ht
343
+ dfb_fusion.wait_front(block.full_block_size());
344
+
345
+ tile_regs_acquire();
346
+ mul_bcast_rows_init(dfb_fusion_id, dfb_gamma_id);
347
+ for (auto i : block.local()) {
348
+ mul_tiles_bcast_rows(
349
+ dfb_fusion_id, dfb_gamma_id, i, block.to_global(i), i); // tile *= 1/(sum(exp(x)))
350
+ #ifdef SFPU_OP_INIT_ACTIVATION
351
+ // Activation must be applied last. If do_beta != 0 then
352
+ // activation will be applied after the beta addition.
353
+ // Otherwise, we can apply the activation here.
354
+ if constexpr (!(do_beta == 1)) {
355
+ SFPU_OP_INIT_ACTIVATION
356
+ SFPU_OP_FUNC_ACTIVATION
357
+ }
358
+ #endif
359
+ }
360
+ tile_regs_commit();
361
+
362
+ dfb_fusion.pop_front(block.full_block_size());
363
+ // We don't pop gamma since it's 1,1,1,Wt and we reuse it for all NCHt
364
+
365
+ dfb_outg.reserve_back(block.full_block_size());
366
+
367
+ tile_regs_wait();
368
+ for (auto i : block.local()) {
369
+ pack_tile(i, dfb_outg_id); // pack either to intermediate (dfb_fusion or out0)
370
+ }
371
+ tile_regs_release();
372
+
373
+ dfb_outg.push_back(block.full_block_size());
374
+ }
375
+ if constexpr (do_beta) {
376
+ pack_reconfig_data_format(dfb_out_id);
377
+ if constexpr (do_gamma) {
378
+ reconfig_data_format_srcb(dfb_gamma_id, dfb_beta_id);
379
+ } else {
380
+ reconfig_data_format_srcb(dfb_ex2pe_id, dfb_beta_id);
381
+ }
382
+ dfb_beta.wait_front(
383
+ block.start() + block.full_block_size()); // TODO: optimization - only wait on first ht
384
+ dfb_fusion.wait_front(block.full_block_size());
385
+
386
+ tile_regs_acquire();
387
+ add_bcast_rows_init(dfb_fusion_id, dfb_beta_id);
388
+ for (auto i : block.local()) {
389
+ add_tiles_bcast_rows(
390
+ dfb_fusion_id, dfb_beta_id, i, block.to_global(i), i); // tile *= 1/(sum(exp(x)))
391
+ #ifdef SFPU_OP_INIT_ACTIVATION
392
+ SFPU_OP_INIT_ACTIVATION
393
+ SFPU_OP_FUNC_ACTIVATION
394
+ #endif
395
+ }
396
+ tile_regs_commit();
397
+
398
+ dfb_fusion.pop_front(block.full_block_size());
399
+ // We don't pop beta since it's 1,1,1,Wt and we reuse it for all NCHt
400
+
401
+ dfb_out.reserve_back(block.full_block_size());
402
+
403
+ tile_regs_wait();
404
+ for (auto i : block.local()) {
405
+ pack_tile(i, dfb_out_id);
406
+ }
407
+ tile_regs_release();
408
+
409
+ dfb_out.push_back(block.full_block_size());
410
+ }
411
+ }
412
+ dfb_ex2pe.pop_front(1);
413
+ dfb_xmm.pop_front(total_buffer_size);
414
+
415
+ #ifdef UNTILIZE_OUT
416
+ constexpr auto dfb_out_rm_id = get_named_compile_time_arg_val("cb_out_rm");
417
+ DataflowBuffer dfb_out_rm(dfb_out_rm_id);
418
+ untilize_all_blocks_from_cb<block_size>(dfb_out, dfb_out_rm, Wt);
419
+ #endif
420
+ } // NCHt loop
421
+ // The reduce scaler is generated once by the reader and reused (waited inside row_wise_mean)
422
+ // across every NCHt iteration but never popped. Pop the producer's tile count once here to
423
+ // balance the CB. The reader pushes a second scaler tile only when the last column tile is
424
+ // partial (W not a multiple of tile_width), matching row_wise_mean's wait count.
425
+ //
426
+ // The reader generates the scalers using tt::constants::TILE_WIDTH; this kernel must use the
427
+ // same width for the count to match, so derive both from the shared helper. (tile_width is the
428
+ // tensor's tile width, which equals TILE_WIDTH for every supported layernorm config — see the
429
+ // partial-column handling in row_wise_mean above.)
430
+ static_assert(
431
+ tile_width == tt::constants::TILE_WIDTH,
432
+ "layernorm reader generates reduce scalers using TILE_WIDTH; compute must use the same tile "
433
+ "width or cb_scaler push/pop counts diverge (issue #48487)");
434
+ constexpr uint32_t num_scaler_tiles = norm::layernorm::reduce_scaler_tile_count(W, tile_width);
435
+ DataflowBuffer(dfb_scaler).pop_front(num_scaler_tiles);
436
+ }
code/models/demos/mast3r/tt/kernels/ln_add_writer.cpp ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local writer for ln_add_compute.cpp (generic_op): per tile row, first the residual sum
3
+ // (CB 17, Wt tiles in blocks) to res_out, then the normalized row (CB cb_out) to out. Same blocking as ttnn's
4
+ // writer_unary_interleaved_start_id_blocked.cpp.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ void kernel_main() {
9
+ const uint32_t dst_addr = get_arg_val<uint32_t>(0);
10
+ const uint32_t Wt = get_arg_val<uint32_t>(1);
11
+ const uint32_t num_tile_rows = get_arg_val<uint32_t>(2);
12
+ const uint32_t tile_offset = get_arg_val<uint32_t>(3);
13
+ const uint32_t res_addr = get_arg_val<uint32_t>(4);
14
+ constexpr uint32_t blk = get_compile_time_arg_val(0);
15
+ constexpr uint32_t CB_OUT = 16, CB_RES = 17;
16
+ constexpr uint32_t LOG2_TILE = 11;
17
+ constexpr uint32_t tile_bytes = 1u << LOG2_TILE;
18
+ const InterleavedPow2AddrGen<OUT_DRAM> og = {.bank_base_address = dst_addr, .log_base_2_of_page_size = LOG2_TILE};
19
+ const InterleavedPow2AddrGen<RES_DRAM> rg = {.bank_base_address = res_addr, .log_base_2_of_page_size = LOG2_TILE};
20
+ for (uint32_t h = 0; h < num_tile_rows; h++) {
21
+ const uint32_t row0 = tile_offset + h * Wt;
22
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
23
+ const uint32_t n = (Wt - w0) < blk ? (Wt - w0) : blk;
24
+ cb_wait_front(CB_RES, blk);
25
+ uint32_t rd = get_read_ptr(CB_RES);
26
+ for (uint32_t i = 0; i < n; ++i) {
27
+ noc_async_write(rd + i * tile_bytes, rg.get_noc_addr(row0 + w0 + i), tile_bytes);
28
+ }
29
+ noc_async_write_barrier();
30
+ cb_pop_front(CB_RES, blk);
31
+ }
32
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
33
+ const uint32_t n = (Wt - w0) < blk ? (Wt - w0) : blk;
34
+ cb_wait_front(CB_OUT, blk);
35
+ uint32_t rd = get_read_ptr(CB_OUT);
36
+ for (uint32_t i = 0; i < n; ++i) {
37
+ noc_async_write(rd + i * tile_bytes, og.get_noc_addr(row0 + w0 + i), tile_bytes);
38
+ }
39
+ noc_async_write_barrier();
40
+ cb_pop_front(CB_OUT, blk);
41
+ }
42
+ }
43
+ }
code/models/demos/mast3r/tt/kernels/ln_add_writer_rb.cpp ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local writer for the fused add + LN (MAST3R_OPT "lnsplit"): ln_add_writer.cpp plus the read of the
3
+ // pre-add operand b (CB 1, in blocks, through a TensorAccessor: b may be block-sharded), so a (reader, one NoC) and b
4
+ // (this RISC, the other NoC) stream in parallel. Per tile row: read b block i, then write residual-sum block i-1
5
+ // (CB 17) as soon as the compute has produced it; then the normalized row (CB 16).
6
+ #include <stdint.h>
7
+ #include "api/dataflow/dataflow_api.h"
8
+
9
+ void kernel_main() {
10
+ const uint32_t dst_addr = get_arg_val<uint32_t>(0);
11
+ const uint32_t Wt = get_arg_val<uint32_t>(1);
12
+ const uint32_t num_tile_rows = get_arg_val<uint32_t>(2);
13
+ const uint32_t tile_offset = get_arg_val<uint32_t>(3);
14
+ const uint32_t res_addr = get_arg_val<uint32_t>(4);
15
+ const uint32_t b_addr = get_arg_val<uint32_t>(5);
16
+ constexpr uint32_t blk = get_compile_time_arg_val(0);
17
+ constexpr auto b_args = TensorAccessorArgs<1>();
18
+ constexpr uint32_t CB_INB = 1, CB_OUT = 16, CB_RES = 17;
19
+ constexpr uint32_t LOG2_TILE = 11;
20
+ constexpr uint32_t tile_bytes = 1u << LOG2_TILE;
21
+ const InterleavedPow2AddrGen<OUT_DRAM> og = {.bank_base_address = dst_addr, .log_base_2_of_page_size = LOG2_TILE};
22
+ const InterleavedPow2AddrGen<RES_DRAM> rg = {.bank_base_address = res_addr, .log_base_2_of_page_size = LOG2_TILE};
23
+ const auto bg = TensorAccessor(b_args, b_addr, tile_bytes);
24
+ for (uint32_t h = 0; h < num_tile_rows; h++) {
25
+ const uint32_t row0 = tile_offset + h * Wt;
26
+ auto write_res = [&](uint32_t w0) {
27
+ const uint32_t n = (Wt - w0) < blk ? (Wt - w0) : blk;
28
+ cb_wait_front(CB_RES, blk);
29
+ uint32_t rd = get_read_ptr(CB_RES);
30
+ for (uint32_t i = 0; i < n; ++i) {
31
+ noc_async_write(rd + i * tile_bytes, rg.get_noc_addr(row0 + w0 + i), tile_bytes);
32
+ }
33
+ noc_async_write_barrier();
34
+ cb_pop_front(CB_RES, blk);
35
+ };
36
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
37
+ const uint32_t n = (Wt - w0) < blk ? (Wt - w0) : blk;
38
+ cb_reserve_back(CB_INB, blk);
39
+ uint32_t wp = get_write_ptr(CB_INB);
40
+ for (uint32_t i = 0; i < n; ++i) {
41
+ noc_async_read(bg.get_noc_addr(row0 + w0 + i), wp + i * tile_bytes, tile_bytes);
42
+ }
43
+ noc_async_read_barrier();
44
+ cb_push_back(CB_INB, blk);
45
+ // Interleave: the residual-sum block of the previous step must be drained here. Reading ALL of b before
46
+ // writing any residual block DEADLOCKS (compute blocks on CB 17 while this kernel blocks on cb_inb):
47
+ // that variant hung chip 9 in round 7.
48
+ if (w0 > 0) {
49
+ write_res(w0 - blk);
50
+ }
51
+ }
52
+ write_res(((Wt - 1) / blk) * blk);
53
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
54
+ const uint32_t n = (Wt - w0) < blk ? (Wt - w0) : blk;
55
+ cb_wait_front(CB_OUT, blk);
56
+ uint32_t rd = get_read_ptr(CB_OUT);
57
+ for (uint32_t i = 0; i < n; ++i) {
58
+ noc_async_write(rd + i * tile_bytes, og.get_noc_addr(row0 + w0 + i), tile_bytes);
59
+ }
60
+ noc_async_write_barrier();
61
+ cb_pop_front(CB_OUT, blk);
62
+ }
63
+ }
64
+ }
code/models/demos/mast3r/tt/kernels/ln_fast_compute.cpp ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // mast3r-p150 model-local fused residual-add + affine-free LayerNorm compute kernel (MAST3R_OPT "lnf"),
4
+ // a drop-in for ln_add_compute.cpp with the same reader (ttnn's reader_unary_interleaved_ln, FUSE_PRE_ADD) and
5
+ // writer (ln_add_writer.cpp) and the same per-row CB protocol (cb_in / cb_inb in blocks, cb_out + CB_RES_OUT in
6
+ // blocks). Fewer passes than ttnn's layernorm.cpp (6 passes, 3 of them over fp32 tiles at HiFi4):
7
+ // A: s = a + b (fp32 dest) -> cb_x (bf16, the value the LN normalises, = ttnn.add's bf16 sum) + CB_RES_OUT
8
+ // M: mean_full = (sum_t s_t @ ONES) / W (matmul with an all-ones tile: the row sum lands in all 32 columns,
9
+ // exact bf16 inputs, fp32 dest accumulation over the Wt tiles)
10
+ // B: d = s - mean_full; d^2 on the SFPU (fp32); the d^2 tiles are summed elementwise by the PACKER (L1
11
+ // accumulation into one fp32 tile) -> no xmm / xmm^2 round trips through L1
12
+ // V: rstd_full = rsqrt((acc @ ONES) / W + eps) (one matmul, SFPU scale / add / rsqrt)
13
+ // C: y = (s - mean_full) * rstd_full (FPU sub, then FPU mul with the DST operand reused as SrcA) -> cb_out
14
+ // Precision: same operand precisions as the stock kernel for the subtraction and the final multiply (fp32 CB
15
+ // operands enter SrcA/SrcB as tf32), the variance uses an exact fp32 SFPU square instead of a tf32 FPU square.
16
+
17
+ #include <cstdint>
18
+
19
+ #include "api/compute/compute_kernel_api.h"
20
+ #include "api/compute/compute_kernel_hw_startup.h"
21
+ #include "api/compute/eltwise_binary.h"
22
+ #include "api/compute/matmul.h"
23
+ #include "api/compute/pack.h"
24
+ #include "api/compute/reconfig_data_format.h"
25
+ #include "api/compute/eltwise_unary/eltwise_unary.h"
26
+ #include "api/compute/eltwise_unary/fill.h"
27
+ #include "api/compute/eltwise_unary/binop_with_scalar.h"
28
+ #include "api/compute/eltwise_unary/rsqrt.h"
29
+ #ifdef LNF_PROF
30
+ #include "tools/profiler/kernel_profiler.hpp"
31
+ #define LZ(n) DeviceZoneScopedN(n)
32
+ #else
33
+ #define LZ(n)
34
+ #endif
35
+
36
+ void kernel_main() {
37
+ const uint32_t NCHt = get_arg_val<uint32_t>(0);
38
+ constexpr uint32_t Wt = get_compile_time_arg_val(0);
39
+ constexpr uint32_t blk = get_compile_time_arg_val(1);
40
+ constexpr uint32_t inv_w_bits = get_compile_time_arg_val(2);
41
+ constexpr uint32_t eps_bits = get_compile_time_arg_val(3);
42
+
43
+ constexpr uint32_t CB_IN = 0, CB_INB = 1, CB_OUT = 16, CB_RES = 17, CB_MEAN = 18, CB_ACC2 = 20, CB_RSTD = 21,
44
+ CB_X = 23, CB_ONES = 24;
45
+
46
+ compute_kernel_hw_startup(CB_IN, CB_INB, CB_X);
47
+
48
+ // all-ones bf16 tile (row-sum broadcast through a matmul)
49
+ cb_reserve_back(CB_ONES, 1);
50
+ tile_regs_acquire();
51
+ fill_tile_init();
52
+ fill_tile(0, 1.0f);
53
+ tile_regs_commit();
54
+ tile_regs_wait();
55
+ pack_reconfig_data_format(CB_ONES);
56
+ pack_tile(0, CB_ONES);
57
+ tile_regs_release();
58
+ cb_push_back(CB_ONES, 1);
59
+ cb_wait_front(CB_ONES, 1);
60
+
61
+ for (uint32_t row = 0; row < NCHt; ++row) {
62
+ {
63
+ LZ("LNF_A");
64
+ // ---- A: s = a + b ----
65
+ reconfig_data_format(CB_IN, CB_INB);
66
+ pack_reconfig_data_format(CB_X);
67
+ add_init(CB_IN, CB_INB);
68
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
69
+ cb_wait_front(CB_IN, blk);
70
+ cb_wait_front(CB_INB, blk);
71
+ tile_regs_acquire();
72
+ for (uint32_t i = 0; i < blk; ++i) {
73
+ add_tiles(CB_IN, CB_INB, i, i, i);
74
+ }
75
+ tile_regs_commit();
76
+ cb_pop_front(CB_IN, blk);
77
+ cb_pop_front(CB_INB, blk);
78
+ cb_reserve_back(CB_X, blk);
79
+ cb_reserve_back(CB_RES, blk);
80
+ tile_regs_wait();
81
+ for (uint32_t i = 0; i < blk; ++i) {
82
+ pack_tile(i, CB_X);
83
+ }
84
+ for (uint32_t i = 0; i < blk; ++i) {
85
+ pack_tile(i, CB_RES);
86
+ }
87
+ tile_regs_release();
88
+ cb_push_back(CB_X, blk);
89
+ cb_push_back(CB_RES, blk);
90
+ }
91
+
92
+ }
93
+ {
94
+ LZ("LNF_M");
95
+ // ---- M: mean_full ----
96
+ cb_wait_front(CB_X, Wt);
97
+ reconfig_data_format(CB_ONES, CB_X); // matmul: in1 (ONES) -> SrcA, in0 (X) -> SrcB
98
+ matmul_init(CB_X, CB_ONES);
99
+ tile_regs_acquire();
100
+ for (uint32_t t = 0; t < Wt; ++t) {
101
+ matmul_tiles(CB_X, CB_ONES, t, 0, 0);
102
+ }
103
+ binop_with_scalar_tile_init();
104
+ mul_unary_tile(0, inv_w_bits);
105
+ tile_regs_commit();
106
+ cb_reserve_back(CB_MEAN, 1);
107
+ pack_reconfig_data_format(CB_MEAN);
108
+ tile_regs_wait();
109
+ pack_tile(0, CB_MEAN);
110
+ tile_regs_release();
111
+ cb_push_back(CB_MEAN, 1);
112
+ cb_wait_front(CB_MEAN, 1);
113
+
114
+ }
115
+ {
116
+ LZ("LNF_B");
117
+ // ---- B: sum of (s - mean)^2, packer L1 accumulation into one fp32 tile ----
118
+ reconfig_data_format(CB_X, CB_MEAN);
119
+ sub_init(CB_X, CB_MEAN);
120
+ square_tile_init();
121
+ pack_reconfig_data_format(CB_ACC2);
122
+ cb_reserve_back(CB_ACC2, 1);
123
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
124
+ tile_regs_acquire();
125
+ for (uint32_t i = 0; i < blk; ++i) {
126
+ sub_tiles(CB_X, CB_MEAN, w0 + i, 0, i);
127
+ square_tile(i);
128
+ }
129
+ tile_regs_commit();
130
+ tile_regs_wait();
131
+ for (uint32_t i = 0; i < blk; ++i) {
132
+ if (w0 + i == 0) {
133
+ pack_reconfig_l1_acc(0);
134
+ }
135
+ pack_tile<true>(i, CB_ACC2, 0);
136
+ if (w0 + i == 0) {
137
+ pack_reconfig_l1_acc(1);
138
+ }
139
+ }
140
+ tile_regs_release();
141
+ }
142
+ pack_reconfig_l1_acc(0);
143
+ cb_push_back(CB_ACC2, 1);
144
+
145
+ }
146
+ {
147
+ LZ("LNF_V");
148
+ // ---- V: rstd_full ----
149
+ cb_wait_front(CB_ACC2, 1);
150
+ reconfig_data_format(CB_ONES, CB_ACC2);
151
+ matmul_init(CB_ACC2, CB_ONES);
152
+ tile_regs_acquire();
153
+ matmul_tiles(CB_ACC2, CB_ONES, 0, 0, 0);
154
+ binop_with_scalar_tile_init();
155
+ mul_unary_tile(0, inv_w_bits);
156
+ add_unary_tile(0, eps_bits);
157
+ rsqrt_tile_init<false>();
158
+ rsqrt_tile<false>(0);
159
+ tile_regs_commit();
160
+ cb_pop_front(CB_ACC2, 1);
161
+ cb_reserve_back(CB_RSTD, 1);
162
+ pack_reconfig_data_format(CB_RSTD);
163
+ tile_regs_wait();
164
+ pack_tile(0, CB_RSTD);
165
+ tile_regs_release();
166
+ cb_push_back(CB_RSTD, 1);
167
+ cb_wait_front(CB_RSTD, 1);
168
+
169
+ }
170
+ {
171
+ LZ("LNF_C");
172
+ // ---- C: y = (s - mean) * rstd ----
173
+ pack_reconfig_data_format(CB_OUT);
174
+ for (uint32_t w0 = 0; w0 < Wt; w0 += blk) {
175
+ tile_regs_acquire();
176
+ reconfig_data_format(CB_X, CB_MEAN);
177
+ sub_init(CB_X, CB_MEAN);
178
+ for (uint32_t i = 0; i < blk; ++i) {
179
+ sub_tiles(CB_X, CB_MEAN, w0 + i, 0, i);
180
+ }
181
+ // SrcA takes the fp32 DST values: configure it for an fp32 operand (tf32 in SrcA, as the stock kernel's
182
+ // fp32 xmm CB) instead of the bf16 cb_x format
183
+ reconfig_data_format(CB_RSTD, CB_RSTD);
184
+ mul_reuse_dest_init<EltwiseBinaryReuseDestType::DEST_TO_SRCA>(CB_RSTD);
185
+ for (uint32_t i = 0; i < blk; ++i) {
186
+ mul_reuse_dest_tiles<EltwiseBinaryReuseDestType::DEST_TO_SRCA>(CB_RSTD, 0, i);
187
+ }
188
+ tile_regs_commit();
189
+ cb_reserve_back(CB_OUT, blk);
190
+ tile_regs_wait();
191
+ for (uint32_t i = 0; i < blk; ++i) {
192
+ pack_tile(i, CB_OUT);
193
+ }
194
+ tile_regs_release();
195
+ cb_push_back(CB_OUT, blk);
196
+ }
197
+ }
198
+ cb_pop_front(CB_X, Wt);
199
+ cb_pop_front(CB_MEAN, 1);
200
+ cb_pop_front(CB_RSTD, 1);
201
+ }
202
+ }
code/models/demos/mast3r/tt/kernels/ln_reader_split.cpp ADDED
@@ -0,0 +1,209 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // mast3r-p150 model-local copy of ttnn's reader_unary_interleaved_ln.cpp (tt-metal 8b98410e730). Only change: with
2
+ // LN_SPLIT_B the second (pre-add) input b is NOT read here; the model-local writer ln_add_writer_rb.cpp reads it on
3
+ // the other RISC / NoC, so the two operand streams of the fused add + LN load in parallel (MAST3R_OPT "lnsplit").
4
+ // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
5
+ //
6
+ // SPDX-License-Identifier: Apache-2.0
7
+
8
+ // Merged reader kernel for layernorm / dit_rms_norm_unary_fused (standard / non-large-tensor path).
9
+ //
10
+ // Handles both TILE-layout input (default) and ROW_MAJOR input (#ifdef TILIZE_IN).
11
+ // The loop structure, scaler/eps generation, gamma/beta reads, and FUSE_PRE_ADD reads
12
+ // are shared between both paths. Only the input accessor setup and the per-ncht input
13
+ // read call branch on TILIZE_IN.
14
+ //
15
+ // Replaces the two separate files:
16
+ // - reader_unary_interleaved_ln.cpp (TILE path, 105 lines)
17
+ // - reader_unary_interleaved_ln_rm_input.cpp (ROW_MAJOR path, 143 lines)
18
+ //
19
+ // Compile-time args:
20
+ // CTA[0] = block_size
21
+ // CTA[1] = use_welford (0 for TILIZE_IN / RMSNORM)
22
+ // CTA[2] = W (logical width in elements)
23
+ // CTA[3..] = TensorAccessorArgs for input a
24
+ // ... = TensorAccessorArgs for b / residual (may be null)
25
+ // ... = TensorAccessorArgs for gamma (may be null)
26
+ // ... = TensorAccessorArgs for beta (may be null)
27
+ // CTA[N] = elem_size_bytes (TILIZE_IN only; unused for TILE path)
28
+ //
29
+ // Runtime args:
30
+ // arg[0] = src_addr
31
+ // arg[1] = NCHt (number of tile-rows assigned to this core)
32
+ // arg[2] = Wt (width in tiles)
33
+ // arg[3] = start_tile_row (tile-row index of first row for this core)
34
+ // TILE: previously passed as tile_offset = start_tile_row * Wt
35
+ // RM: previously passed as start_row; start_tile_row = start_row / TILE_H
36
+ // arg[4] = packed_one_value (legacy; unused, scaler is generated in-kernel)
37
+ // arg[5] = eps (epsilon as bit-cast uint32)
38
+ // arg[6] = gamma_dram_addr
39
+ // arg[7] = beta_dram_addr
40
+ // arg[8] = b_dram_addr (residual, unused if no FUSE_PRE_ADD)
41
+ // arg[9] = H_logical (TILIZE_IN only: total valid rows; unused for TILE path)
42
+
43
+ #include <stdint.h>
44
+ #include "api/dataflow/dataflow_api.h"
45
+ #include "ttnn/cpp/ttnn/kernel_lib/reduce_helpers_dataflow.hpp"
46
+ #include "ttnn/kernel/dataflow/generate_bcast_scalar.hpp"
47
+ #include "ttnn/operations/normalization/kernel_util/generic/blocked_range.h"
48
+ #include "ttnn/operations/normalization/layernorm/device/kernels/layernorm_scaler_tiles.h"
49
+ #include "ttnn/operations/normalization/layernorm/device/kernels/dataflow/layernorm_dataflow_utils.h"
50
+
51
+ namespace generic = norm::kernel_util::generic;
52
+ namespace layernorm_dataflow_utils = norm::layernorm::device::kernels::dataflow;
53
+
54
+ void kernel_main() {
55
+ const uint32_t src_addr = get_arg_val<uint32_t>(0);
56
+ const uint32_t NCHt = get_arg_val<uint32_t>(1);
57
+ const uint32_t Wt = get_arg_val<uint32_t>(2);
58
+ const uint32_t start_tile_row = get_arg_val<uint32_t>(3);
59
+ const uint32_t gamma_addr = get_arg_val<uint32_t>(6);
60
+ const uint32_t beta_addr = get_arg_val<uint32_t>(7);
61
+ const uint32_t b_addr = get_arg_val<uint32_t>(8);
62
+ #ifdef TILIZE_IN
63
+ const uint32_t H_logical = get_arg_val<uint32_t>(9);
64
+ #endif
65
+
66
+ constexpr uint32_t dfb_id_in0 = get_named_compile_time_arg_val("cb_in");
67
+ // Welford-fp32 alias of cb_in (non-fused) or cb_x (fused). Shares SRAM with the
68
+ // primary CB but has its own read/write pointers, so we must push_back on it whenever we
69
+ // push to the primary CB. When welford_fp32_alias is 0, dfb_x_welford == cb_in.
70
+ constexpr uint32_t dfb_id_x_welford = get_named_compile_time_arg_val("cb_x_welford");
71
+ constexpr bool welford_fp32_alias = get_named_compile_time_arg_val("welford_fp32_alias") != 0;
72
+ constexpr uint32_t dfb_id_in1 = get_named_compile_time_arg_val("cb_inb");
73
+ constexpr uint32_t dfb_id_gamma = get_named_compile_time_arg_val("cb_gamma");
74
+ constexpr uint32_t dfb_id_beta = get_named_compile_time_arg_val("cb_beta");
75
+
76
+ Noc noc;
77
+ DataflowBuffer dfb_in0(dfb_id_in0);
78
+ DataflowBuffer dfb_x_welford(dfb_id_x_welford);
79
+ #ifdef FUSE_PRE_ADD
80
+ DataflowBuffer dfb_in1(dfb_id_in1);
81
+ #endif
82
+ #ifdef FUSE_GAMMA
83
+ DataflowBuffer dfb_gamma(dfb_id_gamma);
84
+ #endif
85
+ #ifdef FUSE_BETA
86
+ DataflowBuffer dfb_beta(dfb_id_beta);
87
+ #endif
88
+
89
+ constexpr uint32_t block_size = get_compile_time_arg_val(0);
90
+ constexpr bool use_welford = get_compile_time_arg_val(1) == 1;
91
+ constexpr uint32_t W = get_compile_time_arg_val(2);
92
+ constexpr auto src0_args = TensorAccessorArgs<3>();
93
+ [[maybe_unused]] constexpr auto src1_args = TensorAccessorArgs<src0_args.next_compile_time_args_offset()>();
94
+ [[maybe_unused]] constexpr auto gamma_args = TensorAccessorArgs<src1_args.next_compile_time_args_offset()>();
95
+ [[maybe_unused]] constexpr auto beta_args = TensorAccessorArgs<gamma_args.next_compile_time_args_offset()>();
96
+
97
+ constexpr uint32_t TILE_H = tt::constants::TILE_HEIGHT;
98
+ constexpr uint32_t TILE_W = tt::constants::TILE_WIDTH;
99
+
100
+ #ifdef TILIZE_IN
101
+ // ROW_MAJOR path: input a is a row-major tensor.
102
+ // The compute kernel tilizes dfb_in_rm (c_27) → cb_in (c_0) before processing.
103
+ constexpr uint32_t elem_size_bytes = get_compile_time_arg_val(beta_args.next_compile_time_args_offset());
104
+
105
+ constexpr uint32_t rm_row_stride_bytes = block_size * TILE_W * elem_size_bytes;
106
+ constexpr uint32_t dfb_id_in_rm = get_named_compile_time_arg_val("cb_in_rm");
107
+ DataflowBuffer dfb_in_rm(dfb_id_in_rm);
108
+
109
+ const uint32_t src0_page_bytes = W * elem_size_bytes;
110
+ #else
111
+ // TILE path: input a is already in tile layout.
112
+ const uint32_t src0_page_bytes = dfb_in0.get_tile_size();
113
+ #endif
114
+
115
+ const auto src_a = TensorAccessor(src0_args, src_addr);
116
+
117
+ #ifdef FUSE_GAMMA
118
+ const uint32_t gamma_tile_bytes = dfb_gamma.get_tile_size();
119
+ const auto addrg = TensorAccessor(gamma_args, gamma_addr);
120
+ #endif
121
+ #ifdef FUSE_BETA
122
+ const uint32_t beta_tile_bytes = dfb_beta.get_tile_size();
123
+ const auto addrb = TensorAccessor(beta_args, beta_addr);
124
+ #endif
125
+ #ifdef FUSE_PRE_ADD
126
+ const uint32_t src1_tile_bytes = dfb_in1.get_tile_size();
127
+ const auto src_b = TensorAccessor(src1_args, b_addr);
128
+ #endif
129
+
130
+ // Generate constant tiles for layernorm compute
131
+ constexpr uint32_t dfb_scaler = get_named_compile_time_arg_val("cb_scaler");
132
+ constexpr uint32_t dfb_eps = get_named_compile_time_arg_val("cb_eps");
133
+
134
+ if constexpr (!use_welford) {
135
+ constexpr uint32_t partial_last_tile_cols = W % tt::constants::TILE_WIDTH;
136
+ // Push count shared with the compute kernel's cb_scaler pop count (issue #48487).
137
+ constexpr uint32_t num_scaler_tiles = norm::layernorm::reduce_scaler_tile_count(W, tt::constants::TILE_WIDTH);
138
+
139
+ dataflow_kernel_lib::calculate_and_prepare_reduce_scaler<
140
+ dfb_scaler,
141
+ ckernel::PoolType::SUM,
142
+ ckernel::ReduceDim::REDUCE_ROW,
143
+ dataflow_kernel_lib::SUM_AND_MAX_REDUCE_FACTOR>();
144
+
145
+ if constexpr (num_scaler_tiles == 2) {
146
+ dataflow_kernel_lib::calculate_and_prepare_reduce_scaler<
147
+ dfb_scaler,
148
+ ckernel::PoolType::SUM,
149
+ ckernel::ReduceDim::REDUCE_ROW,
150
+ dataflow_kernel_lib::SUM_AND_MAX_REDUCE_FACTOR>(partial_last_tile_cols);
151
+ }
152
+ }
153
+
154
+ const uint32_t eps = get_arg_val<uint32_t>(5);
155
+ generate_bcast_col_scalar(CircularBuffer(dfb_eps), eps);
156
+
157
+ for (uint32_t ncht = 0; ncht < NCHt; ncht++) {
158
+ const uint32_t curr_tile_row = start_tile_row + ncht;
159
+
160
+ // --- Input read: branches on layout ---
161
+ #ifdef TILIZE_IN
162
+ // ROW_MAJOR: push one tile-row of row-major data into dfb_in_rm (block-by-block).
163
+ // The compute kernel's TILIZE_IN block converts dfb_in_rm → cb_in before processing.
164
+ layernorm_dataflow_utils::push_row_major_blocks_to_cb<decltype(src_a), TILE_W, TILE_H>(
165
+ noc, dfb_in_rm, src_a, Wt, block_size, curr_tile_row, elem_size_bytes, rm_row_stride_bytes, H_logical);
166
+
167
+ #ifdef FUSE_PRE_ADD
168
+ for (auto block : generic::blocks(Wt, block_size)) {
169
+ layernorm_dataflow_utils::read_block_to_cb(
170
+ noc, dfb_in1, src_b, src1_tile_bytes, curr_tile_row * Wt + block.start(), block);
171
+ }
172
+ #endif
173
+ #else
174
+ // TILE: read input a and b (if present) interleaved per block.
175
+ for (auto block : generic::blocks(Wt, block_size)) {
176
+ const uint32_t flat_offset = curr_tile_row * Wt + block.start();
177
+ layernorm_dataflow_utils::read_block_to_cb(noc, dfb_in0, src_a, src0_page_bytes, flat_offset, block);
178
+ #if defined(FUSE_PRE_ADD) && !defined(LN_SPLIT_B)
179
+ layernorm_dataflow_utils::read_block_to_cb(noc, dfb_in1, src_b, src1_tile_bytes, flat_offset, block);
180
+ #elif !defined(FUSE_PRE_ADD)
181
+ // Non-fused welford-fp32 alias: dfb_x_welford shares dfb_in0's memory but has its own
182
+ // read/write pointers. After the data lands in dfb_in0, push
183
+ // dfb_x_welford by the same amount so compute can wait_front on the alias separately
184
+ // for welford reads. Skipped when no alias is active (dfb_x_welford == dfb_in0; the
185
+ // duplicate push would double-count dfb_in0's semaphore).
186
+ if constexpr (welford_fp32_alias) {
187
+ dfb_x_welford.reserve_back(block.full_block_size());
188
+ dfb_x_welford.push_back(block.full_block_size());
189
+ }
190
+ #endif
191
+ }
192
+ #endif
193
+
194
+ // --- Gamma / beta (shared): read once at ncht == 0 ---
195
+ #if defined FUSE_GAMMA || defined FUSE_BETA
196
+ if (ncht == 0) {
197
+ for (auto block : generic::blocks(Wt, block_size)) {
198
+ #ifdef FUSE_GAMMA
199
+ layernorm_dataflow_utils::read_block_to_cb(
200
+ noc, dfb_gamma, addrg, gamma_tile_bytes, block.start(), block);
201
+ #endif
202
+ #ifdef FUSE_BETA
203
+ layernorm_dataflow_utils::read_block_to_cb(noc, dfb_beta, addrb, beta_tile_bytes, block.start(), block);
204
+ #endif
205
+ } // wt loop
206
+ }
207
+ #endif
208
+ } // ncht loop
209
+ }
code/models/demos/mast3r/tt/kernels/mast3r_gelu_poly.h ADDED
@@ -0,0 +1,116 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local SFPU GELU for the fc1 matmul epilogue (MAST3R_OPT "gpoly"), used in the GELU_TANH slot of
3
+ // the model-local mm_gelu_activation.hpp.
4
+ //
5
+ // The reference model (DUSt3R / CroCo Mlp) uses the exact erf GELU, gelu(x) = x Phi(x). Here
6
+ // gelu(x) = 0.5 x + |x| s q(s^2), s = min(|x|, X), X = 4.25,
7
+ // where s q(s^2) is a degree-17 odd minimax polynomial (9 coefficients in s^2, weighted LP fit, tools_prof/gelu_fit.py)
8
+ // of Phi(s) - 0.5 on [0, X] with s q(s^2) = 0.5 exactly at s = X (so |x| > X gives relu(x) up to |x| Phi(-X) <= 1e-5 |x|).
9
+ // Max |error| vs the erf GELU on [-8, 8]: 6e-5 before the bf16 rounding (the tanh form, gelut, is 4.7e-4 away from
10
+ // the erf GELU); fp32 SFPU arithmetic, one bf16 round-to-nearest at the end (as ttnn's gelu_tanh does without fp32
11
+ // dest). ~30 SFPU instructions per row vs ~40 for ttnn's accurate gelu_tanh.
12
+ #pragma once
13
+ #include "api/compute/common_globals.h"
14
+ #if defined(TRISC_MATH) || defined(TRISC_PACK)
15
+ #include "llk_math_eltwise_unary_sfpu_macros.h"
16
+ #include "sfpi.h"
17
+
18
+ namespace ckernel::sfpu {
19
+
20
+ constexpr float MG_X = 4.25f;
21
+ constexpr int MG_N = 9;
22
+ constexpr float MG_C[MG_N] = {3.989576101e-01f, -6.634455919e-02f, 9.786477312e-03f, -1.095801592e-03f, 9.108911036e-05f,
23
+ -5.402808256e-06f, 2.137231121e-07f, -5.002082304e-09f, 5.199383632e-11f};
24
+
25
+ #ifndef MAST3R_GPOLY_X1
26
+ // Two dst rows per step: every coefficient is loaded once for both Horner chains (halves the SFPLOADI pairs) and the
27
+ // two independent SFPMAD chains hide each other's latency. Same arithmetic per element as the one-row form.
28
+ template <bool is_fp32_dest_acc_en, int ITERATIONS = 8>
29
+ inline void calculate_mast3r_gelu_poly() {
30
+ // vConstFloatPrgm0 = X, vConstFloatPrgm1 = 0.5, vConstFloatPrgm2 = MG_C[MG_N - 1] (set in init)
31
+ static_assert(ITERATIONS % 2 == 0, "two rows per step");
32
+ #pragma GCC unroll 4
33
+ for (int d = 0; d < ITERATIONS / 2; d++) {
34
+ sfpi::vFloat x0 = sfpi::dst_reg[0];
35
+ sfpi::vFloat x1 = sfpi::dst_reg[1];
36
+ sfpi::vFloat a0 = sfpi::abs(x0);
37
+ sfpi::vFloat a1 = sfpi::abs(x1);
38
+ sfpi::vFloat xb = sfpi::vConstFloatPrgm0;
39
+ sfpi::vFloat s0 = sfpi::min(a0, xb);
40
+ sfpi::vFloat s1 = sfpi::min(a1, xb);
41
+ sfpi::vFloat u0 = s0 * s0;
42
+ sfpi::vFloat u1 = s1 * s1;
43
+ sfpi::vFloat as0 = a0 * s0;
44
+ sfpi::vFloat as1 = a1 * s1;
45
+ sfpi::vFloat c = MG_C[MG_N - 2];
46
+ sfpi::vFloat p0 = sfpi::vConstFloatPrgm2 * u0 + c;
47
+ sfpi::vFloat p1 = sfpi::vConstFloatPrgm2 * u1 + c;
48
+ #pragma GCC unroll 8
49
+ for (int k = MG_N - 3; k >= 0; --k) {
50
+ c = MG_C[k];
51
+ p0 = p0 * u0 + c;
52
+ p1 = p1 * u1 + c;
53
+ }
54
+ p0 = p0 * as0;
55
+ p1 = p1 * as1;
56
+ x0 = sfpi::dst_reg[0];
57
+ x1 = sfpi::dst_reg[1];
58
+ sfpi::vFloat y0 = x0 * sfpi::vConstFloatPrgm1 + p0;
59
+ sfpi::vFloat y1 = x1 * sfpi::vConstFloatPrgm1 + p1;
60
+ if constexpr (!is_fp32_dest_acc_en) {
61
+ y0 = sfpi::convert<sfpi::vFloat16b>(y0, sfpi::RoundMode::Nearest);
62
+ y1 = sfpi::convert<sfpi::vFloat16b>(y1, sfpi::RoundMode::Nearest);
63
+ }
64
+ sfpi::dst_reg[0] = y0;
65
+ sfpi::dst_reg[1] = y1;
66
+ sfpi::dst_reg += 2;
67
+ }
68
+ }
69
+ #else
70
+ template <bool is_fp32_dest_acc_en, int ITERATIONS = 8>
71
+ inline void calculate_mast3r_gelu_poly() {
72
+ // vConstFloatPrgm0 = X, vConstFloatPrgm1 = 0.5, vConstFloatPrgm2 = MG_C[MG_N - 1] (set in init)
73
+ sfpi::vFloat h0 = MG_C[0], h1 = MG_C[1];
74
+ #pragma GCC unroll 8
75
+ for (int d = 0; d < ITERATIONS; d++) {
76
+ sfpi::vFloat x = sfpi::dst_reg[0];
77
+ sfpi::vFloat a = sfpi::abs(x);
78
+ sfpi::vFloat xb = sfpi::vConstFloatPrgm0;
79
+ sfpi::vFloat s = sfpi::min(a, xb);
80
+ sfpi::vFloat u = s * s;
81
+ sfpi::vFloat p = sfpi::vConstFloatPrgm2 * u + MG_C[MG_N - 2];
82
+ #pragma GCC unroll 8
83
+ for (int k = MG_N - 3; k >= 2; --k) {
84
+ p = p * u + MG_C[k];
85
+ }
86
+ p = p * u + h1;
87
+ p = p * u + h0;
88
+ sfpi::vFloat as = a * s;
89
+ p = p * as;
90
+ sfpi::vFloat y = x * sfpi::vConstFloatPrgm1 + p;
91
+ if constexpr (!is_fp32_dest_acc_en) {
92
+ y = sfpi::convert<sfpi::vFloat16b>(y, sfpi::RoundMode::Nearest);
93
+ }
94
+ sfpi::dst_reg[0] = y;
95
+ sfpi::dst_reg++;
96
+ }
97
+ }
98
+ #endif
99
+
100
+ template <bool is_fp32_dest_acc_en>
101
+ inline void mast3r_gelu_poly_init() {
102
+ math::reset_counters(p_setrwc::SET_ABD_F);
103
+ sfpi::vConstFloatPrgm0 = MG_X;
104
+ sfpi::vConstFloatPrgm1 = 0.5f;
105
+ sfpi::vConstFloatPrgm2 = MG_C[MG_N - 1];
106
+ }
107
+
108
+ } // namespace ckernel::sfpu
109
+ #endif
110
+
111
+ ALWI void mast3r_gelu_tanh_tile_init_pack() {
112
+ PACK(SFPU_UNARY_INIT_FN(gelu_tanh, sfpu::mast3r_gelu_poly_init, (DST_ACCUM_MODE)));
113
+ }
114
+ ALWI void mast3r_gelu_tanh_tile_pack(uint32_t idst) {
115
+ PACK(SFPU_UNARY_CALL(DST_SYNC_MODE, DST_ACCUM_MODE, calculate_mast3r_gelu_poly, (DST_ACCUM_MODE), idst, VectorMode::RC));
116
+ }
code/models/demos/mast3r/tt/kernels/mm_gelu_activation.hpp ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
2
+ //
3
+ // SPDX-License-Identifier: Apache-2.0
4
+ #pragma once
5
+ #include "ttnn/operations/matmul/shared_with_host/activation_type.hpp"
6
+ #include "api/compute/compute_kernel_api.h"
7
+ #include "api/compute/eltwise_unary/gelu.h"
8
+ #include "api/compute/eltwise_unary/relu.h"
9
+ #include "api/compute/eltwise_unary/activations.h"
10
+ #include "api/compute/eltwise_unary/hardtanh.h"
11
+ #include "api/compute/eltwise_unary/selu.h"
12
+ #include "api/compute/eltwise_unary/softplus.h"
13
+ #include "internal/risc_attribs.h"
14
+ #include "mast3r_gelu_poly.h" // mast3r-p150 model-local: GELU_TANH slot -> erf-GELU minimax polynomial (MAST3R_OPT "gpoly")
15
+ #include <cstring> // for memcpy
16
+
17
+ using ttnn::operations::matmul::KernelActivation;
18
+
19
+ // Helper templates to select activation variants based on parameters
20
+ template <KernelActivation ACT, uint32_t PARAM0 = 0, uint32_t PARAM1 = 0>
21
+ struct ActivationInitHelper {
22
+ // Compile-time validation
23
+ static_assert(
24
+ ACT == KernelActivation::NONE || ACT == KernelActivation::SILU || ACT == KernelActivation::TANH ||
25
+ ACT == KernelActivation::GELU || ACT == KernelActivation::GELU_TANH || ACT == KernelActivation::RELU6 ||
26
+ ACT == KernelActivation::SIGMOID || ACT == KernelActivation::HARDSIGMOID ||
27
+ ACT == KernelActivation::HARDTANH || ACT == KernelActivation::SELU || ACT == KernelActivation::SOFTPLUS,
28
+ "Unsupported KernelActivation type for fused activation init");
29
+
30
+ FORCE_INLINE static void init() {
31
+ if constexpr (ACT == KernelActivation::SILU) {
32
+ silu_tile_init_pack();
33
+ } else if constexpr (ACT == KernelActivation::TANH) {
34
+ // PARAM0: 0 = accurate, non-zero = fast
35
+ tanh_tile_init_pack<PARAM0 != 0>();
36
+ } else if constexpr (ACT == KernelActivation::GELU) {
37
+ // PARAM0: 0 = accurate, non-zero = fast
38
+ gelu_tile_init_pack<PARAM0 != 0>();
39
+ } else if constexpr (ACT == KernelActivation::GELU_TANH) {
40
+ mast3r_gelu_tanh_tile_init_pack();
41
+ } else if constexpr (ACT == KernelActivation::RELU6) {
42
+ relu_max_tile_init_pack();
43
+ } else if constexpr (ACT == KernelActivation::SIGMOID) {
44
+ // Enhanced: PARAM1 is fast_approximate flag
45
+ sigmoid_tile_init_pack<PARAM1 != 0>();
46
+ } else if constexpr (ACT == KernelActivation::HARDSIGMOID) {
47
+ hardsigmoid_tile_init_pack();
48
+ } else if constexpr (ACT == KernelActivation::HARDTANH) {
49
+ hardtanh_tile_init_pack();
50
+ } else if constexpr (ACT == KernelActivation::SELU) {
51
+ selu_tile_init_pack();
52
+ } else if constexpr (ACT == KernelActivation::SOFTPLUS) {
53
+ softplus_tile_init_pack();
54
+ }
55
+ }
56
+ };
57
+
58
+ template <KernelActivation ACT, uint32_t PARAM0 = 0, uint32_t PARAM1 = 0, uint32_t PARAM2 = 0>
59
+ struct ActivationApplyHelper {
60
+ // Compile-time validation
61
+ static_assert(
62
+ ACT == KernelActivation::NONE || ACT == KernelActivation::SILU || ACT == KernelActivation::TANH ||
63
+ ACT == KernelActivation::GELU || ACT == KernelActivation::GELU_TANH || ACT == KernelActivation::RELU6 ||
64
+ ACT == KernelActivation::SIGMOID || ACT == KernelActivation::HARDSIGMOID ||
65
+ ACT == KernelActivation::HARDTANH || ACT == KernelActivation::SELU || ACT == KernelActivation::SOFTPLUS,
66
+ "Unsupported KernelActivation type for fused activation apply");
67
+
68
+ // Parameter-specific validation
69
+ static_assert(
70
+ ACT != KernelActivation::SOFTPLUS || PARAM0 != 0,
71
+ "SOFTPLUS PARAM0 (beta) must be non-zero to avoid division by zero");
72
+
73
+ FORCE_INLINE static void apply(uint32_t tile_index) {
74
+ if constexpr (ACT == KernelActivation::SILU) {
75
+ silu_tile_pack(tile_index);
76
+ } else if constexpr (ACT == KernelActivation::TANH) {
77
+ // PARAM0: 0 = accurate, non-zero = fast
78
+ tanh_tile_pack<PARAM0 != 0>(tile_index);
79
+ } else if constexpr (ACT == KernelActivation::GELU) {
80
+ // PARAM0: 0 = accurate, non-zero = fast
81
+ gelu_tile_pack<PARAM0 != 0>(tile_index);
82
+ } else if constexpr (ACT == KernelActivation::GELU_TANH) {
83
+ mast3r_gelu_tanh_tile_pack(tile_index);
84
+ } else if constexpr (ACT == KernelActivation::RELU6) {
85
+ // PARAM0 is the max value (as uint32_t bit pattern)
86
+ // Default to 6.0 if PARAM0 is 0
87
+ constexpr uint32_t max = (PARAM0 != 0) ? PARAM0 : 0x40c00000u;
88
+ relu_max_tile_pack(tile_index, max);
89
+ } else if constexpr (ACT == KernelActivation::SIGMOID) {
90
+ // Enhanced: PARAM0 is vector mode, PARAM1 is fast_approximate
91
+ constexpr VectorMode vec_mode = (PARAM0 == 1) ? VectorMode::R
92
+ : (PARAM0 == 2) ? VectorMode::C
93
+ : VectorMode::RC;
94
+ sigmoid_tile_pack<vec_mode, PARAM1 != 0>(tile_index);
95
+ } else if constexpr (ACT == KernelActivation::HARDSIGMOID) {
96
+ hardsigmoid_tile_pack(tile_index);
97
+ } else if constexpr (ACT == KernelActivation::HARDTANH) {
98
+ hardtanh_tile_pack(tile_index, PARAM0, PARAM1);
99
+ } else if constexpr (ACT == KernelActivation::SELU) {
100
+ // PARAM0 is alpha, PARAM1 is lambda
101
+ selu_tile_pack(tile_index, PARAM0, PARAM1);
102
+ } else if constexpr (ACT == KernelActivation::SOFTPLUS) {
103
+ // PARAM0 is beta, PARAM2 beta reciprocal, PARAM1 is threshold
104
+ softplus_tile_pack(tile_index, PARAM0, PARAM2, PARAM1);
105
+ }
106
+ }
107
+ };
108
+
109
+ template <KernelActivation ACT, uint32_t PARAM0 = 0, uint32_t PARAM1 = 0, uint32_t PARAM2 = 0>
110
+ FORCE_INLINE void apply_activation_from_pack(uint32_t out_subblock_num_tiles) {
111
+ PACK(TTI_SEMWAIT(
112
+ p_stall::STALL_TDMA | p_stall::STALL_CFG, semaphore::t6_sem(semaphore::MATH_PACK), p_stall::STALL_ON_ZERO));
113
+
114
+ // Flip destination register offset for PACKER access
115
+ PACK(TT_SETC16(DEST_TARGET_REG_CFG_MATH_Offset_ADDR32, ckernel::packer::get_packer_dest_offset()));
116
+
117
+ for (uint32_t i = 0; i < out_subblock_num_tiles; i++) {
118
+ ActivationApplyHelper<ACT, PARAM0, PARAM1, PARAM2>::apply(i);
119
+ }
120
+
121
+ // Wait for SFPU completion before packing
122
+ PACK(TTI_STALLWAIT(p_stall::STALL_PACK, p_stall::WAIT_SFPU));
123
+ }
code/models/demos/mast3r/tt/kernels/mm_gelu_compute.cpp ADDED
@@ -0,0 +1,657 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc.
2
+ //
3
+ // SPDX-License-Identifier: Apache-2.0
4
+ // mast3r-p150 model-local copy of ttnn bmm_large_block_zm_fused_bias_activation.cpp (tt-metal 8b98410e730), unchanged
5
+ // except that its "mm_gelu_activation.hpp" resolves to the model-local copy (GELU_TANH -> mast3r_gelu_poly.h, MAST3R_OPT "gpoly")
6
+ // and the optional fused residual add (MAST3R_RESID_CB, MAST3R_OPT "resmm").
7
+
8
+ #include <cstdint>
9
+
10
+ #include "api/compute/matmul.h"
11
+ #include "api/compute/compute_kernel_hw_startup.h"
12
+ #include "api/compute/pack_untilize.h"
13
+ #include "api/compute/tile_move_copy.h"
14
+ #include "api/compute/transpose.h"
15
+ #include "api/dataflow/dataflow_buffer.h"
16
+ #include "internal/mod_div_lib.h"
17
+
18
+ #ifdef FUSE_BIAS
19
+ #include "api/compute/bcast.h"
20
+ #endif
21
+
22
+ #include "api/compute/eltwise_binary.h"
23
+ #ifdef SFPU_ACTIVATION
24
+ #include "mm_gelu_activation.hpp"
25
+ #endif
26
+
27
+ // Please update
28
+ // tests/tt_metal/tt_metal/perf_microbenchmark/1_compute_mm/kernels/bmm_large_block_zm_fused_bias_activation_copy.cpp
29
+ // when making any changes to this file.
30
+ // Have to keep a copy because cannot import ttnn into tests/tt_metal.
31
+ // With FUSE_BIAS: row_broadcast_bias (row-broadcast vs elementwise add_tiles) is compile-time arg 18 here;
32
+ // the perf copy uses index 14 (different compile-time arg layout).
33
+
34
+ /**
35
+ * @brief Transposes a block of tiles from one circular buffer to another.
36
+ *
37
+ * This function reads a block of tiles from the input circular buffer (cb), performs a width-height
38
+ * (WH) transpose on each tile, and writes the transposed tiles to the output circular buffer.
39
+ * The operation is performed in blocks of `block_size` tiles for efficiency, with a separate loop
40
+ * at the end to handle any leftover tiles when the total tile count is not divisible by
41
+ * `block_size`. The default block size is 4, since there are guaranteed to be 4 tiles in the dst
42
+ * regs irrespective of dst sync mode or data format.
43
+ *
44
+ * @tparam in0_block_num_tiles The number of tiles in the block to be transposed.
45
+ * @tparam block_size The number of tiles in each block to be transposed.
46
+ * @param in0_transpose_dfb_id Circular buffer ID to read the original tiles from.
47
+ * @param in0_dfb_id Circular buffer ID to which the transposed tiles are written.
48
+ */
49
+ template <uint32_t in0_block_num_tiles, uint32_t block_size = 4>
50
+ FORCE_INLINE void transpose_tile_block(uint32_t in0_transpose_dfb_id, uint32_t in0_dfb_id) {
51
+ DataflowBuffer in0_transpose_dfb(in0_transpose_dfb_id);
52
+ DataflowBuffer in0_dfb(in0_dfb_id);
53
+ constexpr uint32_t num_blocks = in0_block_num_tiles / block_size;
54
+ constexpr uint32_t last_block_size = in0_block_num_tiles % block_size;
55
+ // Lets do 2 passes: One loop until last and one last for the left overs
56
+ for (uint32_t block_idx = 0; block_idx < num_blocks; ++block_idx) {
57
+ in0_transpose_dfb.wait_front(block_size);
58
+ tile_regs_acquire();
59
+ for (uint32_t tile_idx = 0; tile_idx < block_size; tile_idx++) {
60
+ transpose_tile(in0_transpose_dfb_id, tile_idx, tile_idx);
61
+ }
62
+ tile_regs_commit();
63
+ in0_transpose_dfb.pop_front(block_size);
64
+
65
+ in0_dfb.reserve_back(block_size);
66
+ tile_regs_wait();
67
+ for (uint32_t tile_idx = 0; tile_idx < block_size; tile_idx++) {
68
+ pack_tile(tile_idx, in0_dfb_id);
69
+ }
70
+ tile_regs_release();
71
+ in0_dfb.push_back(block_size);
72
+ }
73
+
74
+ if constexpr (last_block_size > 0) {
75
+ in0_transpose_dfb.wait_front(last_block_size);
76
+ tile_regs_acquire();
77
+ for (uint32_t tile_idx = 0; tile_idx < last_block_size; tile_idx++) {
78
+ transpose_tile(in0_transpose_dfb_id, tile_idx, tile_idx);
79
+ }
80
+ tile_regs_commit();
81
+ in0_transpose_dfb.pop_front(last_block_size);
82
+
83
+ in0_dfb.reserve_back(last_block_size);
84
+ tile_regs_wait();
85
+ for (uint32_t tile_idx = 0; tile_idx < last_block_size; tile_idx++) {
86
+ pack_tile(tile_idx, in0_dfb_id);
87
+ }
88
+ tile_regs_release();
89
+ in0_dfb.push_back(last_block_size);
90
+ }
91
+ }
92
+
93
+ FORCE_INLINE void reload_from_cb_to_dst(
94
+ uint32_t in0_dfb_id,
95
+ uint32_t in1_dfb_id,
96
+ uint32_t mm_partials_dfb_id,
97
+ uint32_t mm_partials_reload_dfb_id,
98
+ bool in1_transpose_tile,
99
+ uint32_t out_subblock_num_tiles,
100
+ uint32_t out_subblock_w,
101
+ uint32_t out_subblock_h,
102
+ uint32_t in0_block_w) {
103
+ DataflowBuffer mm_partials_dfb(mm_partials_dfb_id);
104
+ // mm_partials_reload_dfb_id is the CB view the reload copies through. It equals mm_partials_dfb_id
105
+ // unless the partials CB is also read as an FPU operand elsewhere (the fused bias add reads it via
106
+ // SrcA), in which case UnpackToDestFp32 cannot be set on it directly; instead a second buffer index
107
+ // aliases the same SRAM with UnpackToDestFp32 set, and the reload copies through that alias while the
108
+ // FPU consumer keeps the original view. The alias has its own read pointer, so align it with the
109
+ // partials CB's current read position before copying.
110
+ // Reconfigure input
111
+ copy_tile_to_dst_init_short_with_dt(in1_dfb_id, mm_partials_reload_dfb_id);
112
+ mm_partials_dfb.wait_front(out_subblock_num_tiles);
113
+
114
+ if (mm_partials_reload_dfb_id != mm_partials_dfb_id) {
115
+ // Only the unpacker owns cb_interface / the read pointer; keep this off the MATH/PACK threads.
116
+ UNPACK(
117
+ (get_local_cb_interface(mm_partials_reload_dfb_id).fifo_rd_ptr =
118
+ get_local_cb_interface(mm_partials_dfb_id).fifo_rd_ptr));
119
+ }
120
+
121
+ uint32_t start_dst_index = 0;
122
+ uint32_t start_tile_index = 0;
123
+ copy_block(mm_partials_reload_dfb_id, start_tile_index, start_dst_index, out_subblock_num_tiles);
124
+
125
+ mm_partials_dfb.pop_front(out_subblock_num_tiles);
126
+ // Reconfigure srcA back
127
+ reconfig_data_format_srca(mm_partials_reload_dfb_id, in1_dfb_id);
128
+ matmul_block_init(in0_dfb_id, in1_dfb_id, in1_transpose_tile, out_subblock_w, out_subblock_h, in0_block_w);
129
+ }
130
+
131
+ template <uint32_t out_subblock_w, uint32_t out_block_w>
132
+ inline void reblock_and_untilize(
133
+ uint32_t num_out_subblocks_in_col,
134
+ uint32_t out_subblock_num_tiles,
135
+ uint32_t out_subblock_h,
136
+ uint32_t interm_dfb_id,
137
+ uint32_t out_dfb_id) {
138
+ DataflowBuffer interm_dfb(interm_dfb_id);
139
+ DataflowBuffer out_dfb(out_dfb_id);
140
+ uint32_t num_tiles_in_row_of_subblocks = mulsi3(out_subblock_num_tiles, num_out_subblocks_in_col);
141
+ interm_dfb.wait_front(num_tiles_in_row_of_subblocks);
142
+
143
+ uint32_t within_block_index = 0;
144
+ for (uint32_t h = 0; h < out_subblock_h; h++) {
145
+ uint32_t block_offset = 0;
146
+
147
+ out_dfb.reserve_back(out_block_w);
148
+ for (uint32_t n = 0; n < num_out_subblocks_in_col; n++) {
149
+ tile_regs_acquire();
150
+ for (uint32_t w = 0; w < out_subblock_w; w++) {
151
+ uint32_t tile_index = block_offset + within_block_index + w;
152
+ copy_tile(interm_dfb_id, tile_index, w);
153
+ }
154
+ tile_regs_commit();
155
+ tile_regs_wait();
156
+ pack_untilize_dest<out_subblock_w, out_block_w>(out_dfb_id, 1, n);
157
+ tile_regs_release();
158
+ block_offset += out_subblock_num_tiles;
159
+ }
160
+ out_dfb.push_back(out_block_w);
161
+
162
+ within_block_index += out_subblock_w;
163
+ }
164
+ interm_dfb.pop_front(num_tiles_in_row_of_subblocks);
165
+ }
166
+
167
+ void kernel_main() {
168
+ // RUNTIME ARGS
169
+ #ifdef MATMUL_DRAM_SHARDED
170
+ const bool is_worker_core = get_arg_val<uint32_t>(0) == 1;
171
+ // if not worker core, skip
172
+ if (not is_worker_core) {
173
+ return;
174
+ }
175
+ #endif
176
+
177
+ constexpr uint32_t in0_block_w = get_compile_time_arg_val(0); // inner block size in tiles
178
+ constexpr uint32_t in0_num_subblocks = get_compile_time_arg_val(1); // outer row block size (in inner row blocks)
179
+ constexpr uint32_t in0_block_num_tiles =
180
+ get_compile_time_arg_val(2); // out_subblock_h*in0_block_w*in0_num_subblocks;
181
+ constexpr uint32_t in0_subblock_num_tiles = get_compile_time_arg_val(3); // out_subblock_h*in0_block_w
182
+ constexpr uint32_t in1_num_subblocks =
183
+ get_compile_time_arg_val(4); // outer column block size (in inner column blocks)
184
+ constexpr uint32_t in1_block_num_tiles =
185
+ get_compile_time_arg_val(5); // out_subblock_w*in0_block_w* in1_num_subblocks;
186
+ constexpr uint32_t in1_block_w = get_compile_time_arg_val(6); // out_subblock_w*in1_num_subblocks
187
+ constexpr uint32_t num_blocks_inner_dim = get_compile_time_arg_val(7); // outer inner dim (in inner dim blocks)
188
+ constexpr uint32_t num_blocks_w_dim = get_compile_time_arg_val(8); // outer inner dim (in inner dim blocks)
189
+ constexpr uint32_t num_blocks_h_dim = get_compile_time_arg_val(9); // outer inner dim (in inner dim blocks)
190
+ constexpr uint32_t out_subblock_h = get_compile_time_arg_val(10); // inner row block size in tiles
191
+ constexpr uint32_t out_subblock_w = get_compile_time_arg_val(11); // inner column block size in tiles
192
+ constexpr uint32_t out_subblock_num_tiles = get_compile_time_arg_val(12); // out_subblock_h * out_subblock_w;
193
+ constexpr uint32_t batch = get_compile_time_arg_val(13); // batch dim
194
+ constexpr uint32_t out_block_num_tiles = get_compile_time_arg_val(14); // number of tiles in out_block
195
+ constexpr bool untilize_out = get_compile_time_arg_val(15); // untilize output
196
+ // This boolean is set when the number of batches is only known at runtime, typically based on a sparsity tensor.
197
+ constexpr bool get_batch_from_reader = (bool)get_compile_time_arg_val(16);
198
+ constexpr bool in0_transpose_tile = (bool)get_compile_time_arg_val(17);
199
+
200
+ constexpr uint32_t out_block_w = out_subblock_w * in1_num_subblocks;
201
+
202
+ constexpr uint32_t in0_dfb_id = in0_transpose_tile ? get_named_compile_time_arg_val("cb_in0_transposed")
203
+ : get_named_compile_time_arg_val("cb_in0");
204
+ constexpr uint32_t in1_dfb_id = get_named_compile_time_arg_val("cb_in1");
205
+ constexpr uint32_t out_dfb_id = get_named_compile_time_arg_val("cb_out");
206
+ constexpr uint32_t mm_partials_dfb_id = get_named_compile_time_arg_val("cb_intermed0");
207
+ // CB view the cross-block reload copies through: the UnpackToDestFp32-marked alias of the partials
208
+ // CB when it is also read as an FPU operand (fused bias), otherwise the partials CB itself.
209
+ #ifdef MM_PARTIALS_RELOAD_ALIAS_CB
210
+ // The partials CB is also read as an FPU operand (fused bias) and so cannot carry UnpackToDestFp32;
211
+ // the reload instead copies through this alias view of the same SRAM, which does carry the flag.
212
+ constexpr uint32_t mm_partials_reload_dfb_id = MM_PARTIALS_RELOAD_ALIAS_CB;
213
+ #else
214
+ constexpr uint32_t mm_partials_reload_dfb_id = mm_partials_dfb_id;
215
+ #endif
216
+ constexpr uint32_t untilize_mode_out_dfb_id = untilize_out ? mm_partials_dfb_id : out_dfb_id;
217
+ // When in0 needs to be transposed, the original data is read from cb_in0 (in0_transpose_dfb_id),
218
+ // transposed, and the result is written to cb_in0_transposed (in0_dfb_id), which is then used
219
+ // as input for the matmul call.
220
+ constexpr uint32_t in0_transpose_dfb_id = get_named_compile_time_arg_val("cb_in0");
221
+
222
+ DataflowBuffer in0_dfb(in0_dfb_id);
223
+ DataflowBuffer in1_dfb(in1_dfb_id);
224
+ DataflowBuffer mm_partials_dfb(mm_partials_dfb_id);
225
+ DataflowBuffer untilize_mode_out_dfb(untilize_mode_out_dfb_id);
226
+
227
+ #ifdef FUSE_BIAS
228
+ constexpr uint32_t bias_dfb_id = get_named_compile_time_arg_val("cb_bias");
229
+ constexpr uint32_t bias_ntiles = get_named_compile_time_arg_val("bias_ntiles");
230
+ constexpr uint32_t mm_out_dfb_id = mm_partials_dfb_id;
231
+ // true: row-0 broadcast ([N] / [...,1,N]); false: elementwise add_tiles (bias has multiple M rows).
232
+ constexpr bool row_broadcast_bias = (bool)get_compile_time_arg_val(18);
233
+ DataflowBuffer bias_dfb(bias_dfb_id);
234
+ #ifdef MAST3R_RESID_CB
235
+ // MAST3R_OPT "resmm": out = in0 @ in1 + bias + resid. The residual CB is globally allocated on a BLOCK_SHARDED
236
+ // tensor whose shard == this core's output block (row-major per_core_M x per_core_N tiles, same spec as the
237
+ // output). Only valid for batch == 1, one output block per core (asserted on the host).
238
+ constexpr uint32_t resid_dfb_id = MAST3R_RESID_CB;
239
+ DataflowBuffer resid_dfb(resid_dfb_id);
240
+ #endif
241
+ #else
242
+ constexpr uint32_t mm_out_dfb_id = untilize_mode_out_dfb_id;
243
+ #endif
244
+ DataflowBuffer mm_out_dfb(mm_out_dfb_id);
245
+
246
+ // Number of valid in1 columns in the last in1 subblock. For the DRAM-sharded variant the
247
+ // planner may pad per_core_N_compute beyond per_core_N_in1_sender so that out_subblock_w can be
248
+ // larger; the reader only pushes per_core_N_in1_sender tiles per block into cb_in1. To avoid
249
+ // reading those non-existent (padded) cb_in1 tiles, the compute kernel narrows the matmul_block
250
+ // call on the last in1 subblock to last_subblock_w_valid lanes. When no padding occurs this
251
+ // equals out_subblock_w and the original full-width path is preserved.
252
+ #ifdef MATMUL_DRAM_SHARDED
253
+ constexpr uint32_t last_subblock_w_valid = get_named_compile_time_arg_val("last_subblock_w_valid");
254
+ #else
255
+ constexpr uint32_t last_subblock_w_valid = out_subblock_w;
256
+ #endif
257
+ constexpr bool last_subblock_padded = last_subblock_w_valid < out_subblock_w;
258
+
259
+ #ifdef SFPU_ACTIVATION
260
+ constexpr KernelActivation activation_type =
261
+ static_cast<KernelActivation>(get_named_compile_time_arg_val("activation_type"));
262
+ constexpr uint32_t activation_param0 = get_named_compile_time_arg_val("activation_param0");
263
+ constexpr uint32_t activation_param1 = get_named_compile_time_arg_val("activation_param1");
264
+ constexpr uint32_t activation_param2 = get_named_compile_time_arg_val("activation_param2");
265
+
266
+ ActivationInitHelper<activation_type, activation_param0, activation_param1>::init();
267
+ #endif
268
+
269
+ #ifdef IN1_TRANSPOSE_TILE
270
+ constexpr uint32_t in1_transpose_tile = true;
271
+ #else
272
+ constexpr uint32_t in1_transpose_tile = false;
273
+ #endif
274
+
275
+ constexpr bool spill = num_blocks_inner_dim > 1;
276
+
277
+ compute_kernel_hw_startup<SrcOrder::Reverse>(in0_dfb_id, in1_dfb_id, mm_partials_dfb_id);
278
+ matmul_block_init(in0_dfb_id, in1_dfb_id, in1_transpose_tile, out_subblock_w, out_subblock_h, in0_block_w);
279
+ for (uint32_t b = 0; b < batch; b++) {
280
+ if constexpr (get_batch_from_reader) {
281
+ // Check whether this batch is valid
282
+ bool is_batch_valid = false;
283
+ UNPACK(is_batch_valid = (bool)mailbox_read(ckernel::ThreadId::BriscThreadId);)
284
+ MATH(is_batch_valid = (bool)mailbox_read(ckernel::ThreadId::BriscThreadId);)
285
+ PACK(is_batch_valid = (bool)mailbox_read(ckernel::ThreadId::BriscThreadId);)
286
+ if (!is_batch_valid) {
287
+ continue;
288
+ }
289
+ }
290
+
291
+ for (uint32_t bh = 0; bh < num_blocks_h_dim; ++bh) {
292
+ for (uint32_t bw = 0; bw < num_blocks_w_dim; ++bw) {
293
+ bool enable_reload = false;
294
+
295
+ #ifdef PACK_RELU
296
+ // for each batch we start with relu disabled so that intermediate results are not relu'd
297
+ if constexpr (batch > 1 || num_blocks_h_dim > 1 || num_blocks_w_dim > 1) {
298
+ PACK((llk_pack_relu_config(ReluConfig::none())));
299
+ }
300
+ #endif
301
+
302
+ if constexpr (batch > 1 || num_blocks_h_dim > 1 || num_blocks_w_dim > 1) {
303
+ PACK((pack_reconfig_data_format(mm_partials_dfb_id)));
304
+ }
305
+
306
+ for (uint32_t block = 0; block < num_blocks_inner_dim; block++) {
307
+ bool last_out = block == (num_blocks_inner_dim - 1);
308
+ // Configure packer once for pack out without Bias
309
+ #if not defined FUSE_BIAS and defined PACK_RELU
310
+ if (last_out) {
311
+ // if last block we pack the final result with relu enabled
312
+ PACK((llk_pack_relu_config(ReluConfig::zero())));
313
+ }
314
+ #endif
315
+
316
+ if constexpr (in0_transpose_tile) {
317
+ reconfig_data_format_srca(in1_dfb_id, in0_transpose_dfb_id);
318
+ transpose_init(in0_transpose_dfb_id);
319
+ PACK((pack_reconfig_data_format(in0_dfb_id)));
320
+ #ifdef PACKER_L1_ACC
321
+ PACK((llk_pack_reconfig_l1_acc(0)));
322
+ #endif
323
+ transpose_tile_block<in0_block_num_tiles>(in0_transpose_dfb_id, in0_dfb_id);
324
+ reconfig_data_format_srca(in0_transpose_dfb_id, in1_dfb_id);
325
+ matmul_block_init(
326
+ in0_dfb_id, in1_dfb_id, in1_transpose_tile, out_subblock_w, out_subblock_h, in0_block_w);
327
+ PACK((pack_reconfig_data_format(mm_partials_dfb_id)));
328
+ }
329
+
330
+ in0_dfb.wait_front(in0_block_num_tiles);
331
+ in1_dfb.wait_front(in1_block_num_tiles);
332
+
333
+ int in0_index_subblock_offset = 0;
334
+ for (uint32_t in0_subblock = 0; in0_subblock < in0_num_subblocks; in0_subblock++) {
335
+ int in1_index_subblock_offset = 0;
336
+ for (uint32_t in1_subblock = 0; in1_subblock < in1_num_subblocks; in1_subblock++) {
337
+ // When last_subblock_padded is true the last in1 subblock has
338
+ // (out_subblock_w - last_subblock_w_valid) padded lanes whose cb_in1 tiles were
339
+ // never pushed by the reader. Narrow matmul_block so the unpacker only touches
340
+ // tiles that exist; the padded dst lanes are left at whatever the previous
341
+ // operation wrote there and the output writer (BRISC) drops those columns.
342
+ const bool is_last_in1_subblock_padded =
343
+ last_subblock_padded && (in1_subblock == in1_num_subblocks - 1);
344
+ const uint32_t effective_subblock_w =
345
+ is_last_in1_subblock_padded ? last_subblock_w_valid : out_subblock_w;
346
+
347
+ tile_regs_acquire();
348
+ if (enable_reload) {
349
+ reload_from_cb_to_dst(
350
+ in0_dfb_id,
351
+ in1_dfb_id,
352
+ mm_partials_dfb_id,
353
+ mm_partials_reload_dfb_id,
354
+ in1_transpose_tile,
355
+ out_subblock_num_tiles,
356
+ out_subblock_w,
357
+ out_subblock_h,
358
+ in0_block_w);
359
+ }
360
+
361
+ #ifndef SKIP_COMPUTE
362
+ // Compute output sub-block
363
+ uint32_t dst_index =
364
+ 0; // start at 0, each call to matmul_block internally increments dst_index
365
+ uint32_t in0_index = in0_index_subblock_offset; // offset into in0 block
366
+ uint32_t in1_index = in1_index_subblock_offset; // offset into in1 block
367
+ // inner dim that we accumulate is the inner dim of in0/in1, which is in0_block_w
368
+ for (uint32_t inner_dim_idx = 0; inner_dim_idx < in0_block_w; ++inner_dim_idx) {
369
+ // matmul outer product of (out_subblock_h x out_subblock_w) tiles that fill dst
370
+ // accumulation is done by iterating matmul_block across inner dim
371
+ // in0_block_w is passed as innder dim (kt) to matmul_block, internally used to stride
372
+ // in0
373
+ matmul_block(
374
+ in0_dfb_id,
375
+ in1_dfb_id,
376
+ in0_index,
377
+ in1_index,
378
+ dst_index,
379
+ in1_transpose_tile,
380
+ effective_subblock_w,
381
+ out_subblock_h,
382
+ in0_block_w);
383
+ in0_index++; // stride right by 1
384
+ in1_index += in1_block_w; // to stride down by 1 need to stride by in_per_core_w
385
+ // (should be called in1_block_w)
386
+ }
387
+
388
+ #endif // SKIP_COMPUTE
389
+
390
+ if (last_out) {
391
+ tile_regs_commit();
392
+ mm_out_dfb.reserve_back(out_subblock_num_tiles);
393
+
394
+ #if defined SFPU_ACTIVATION and not defined FUSE_BIAS
395
+ apply_activation_from_pack<
396
+ activation_type,
397
+ activation_param0,
398
+ activation_param1,
399
+ activation_param2>(out_subblock_num_tiles);
400
+ #else
401
+ tile_regs_wait();
402
+ #endif
403
+
404
+ #if defined FP32_DEST_ACC_EN or defined PACKER_L1_ACC
405
+ PACK((pack_reconfig_data_format(mm_out_dfb_id)));
406
+ #endif
407
+
408
+ #ifdef PACKER_L1_ACC
409
+ #ifdef FUSE_BIAS
410
+ if (block == 0) { // no accumulation for first iteration
411
+ PACK((llk_pack_reconfig_l1_acc(0)));
412
+ } else {
413
+ PACK((llk_pack_reconfig_l1_acc(1)));
414
+ }
415
+ #else
416
+ PACK((llk_pack_reconfig_l1_acc(0)));
417
+ #endif
418
+ #endif
419
+ uint32_t start_dst_index = 0;
420
+ pack_block(start_dst_index, mm_out_dfb_id, out_subblock_num_tiles);
421
+
422
+ tile_regs_release();
423
+ mm_out_dfb.push_back(out_subblock_num_tiles);
424
+
425
+ } else {
426
+ tile_regs_commit();
427
+ mm_partials_dfb.reserve_back(out_subblock_num_tiles);
428
+ tile_regs_wait();
429
+
430
+ #ifdef PACKER_L1_ACC
431
+ if (block == 0) { // no accumulation for first iteration
432
+ PACK((llk_pack_reconfig_l1_acc(0)));
433
+ } else if (block == 1) {
434
+ PACK((llk_pack_reconfig_l1_acc(1)));
435
+ } else if (in0_transpose_tile) {
436
+ // For each block, l1_acc would have been enabled during the
437
+ // transpose stage. So let us put it back here.
438
+ PACK((llk_pack_reconfig_l1_acc(1)));
439
+ }
440
+ #endif
441
+
442
+ uint32_t start_dst_index = 0;
443
+ pack_block(start_dst_index, mm_partials_dfb_id, out_subblock_num_tiles);
444
+
445
+ tile_regs_release();
446
+ mm_partials_dfb.push_back(out_subblock_num_tiles);
447
+ }
448
+
449
+ in1_index_subblock_offset += out_subblock_w;
450
+ }
451
+ in0_index_subblock_offset += in0_subblock_num_tiles;
452
+ }
453
+
454
+ #ifdef PACKER_L1_ACC
455
+ #ifdef FUSE_BIAS
456
+ if (block < num_blocks_inner_dim - 1) {
457
+ // Wait/pop in subblock-sized steps so the step size
458
+ // matches the bias section's wait_front(out_subblock_num_tiles),
459
+ // satisfying the CB API requirement that all wait_front
460
+ // increments on a given CB are identical.
461
+ for (uint32_t s = 0; s < out_block_num_tiles; s += out_subblock_num_tiles) {
462
+ mm_partials_dfb.wait_front(out_subblock_num_tiles);
463
+ mm_partials_dfb.pop_front(out_subblock_num_tiles);
464
+ }
465
+ }
466
+ // never reload when with bias, bias uses intermediate buffer
467
+ enable_reload = false;
468
+ #else
469
+ // Last iteration does spill and reload to output buffer
470
+ if (block < num_blocks_inner_dim - 2) {
471
+ for (uint32_t s = 0; s < out_block_num_tiles; s += out_subblock_num_tiles) {
472
+ mm_partials_dfb.wait_front(out_subblock_num_tiles);
473
+ mm_partials_dfb.pop_front(out_subblock_num_tiles);
474
+ }
475
+ }
476
+ if (block == num_blocks_inner_dim - 2) {
477
+ enable_reload = true;
478
+ } // reload when last iteration
479
+ #endif
480
+ #else
481
+ if constexpr (spill) {
482
+ enable_reload = true;
483
+ }
484
+ #endif
485
+
486
+ in0_dfb.pop_front(in0_block_num_tiles);
487
+ in1_dfb.pop_front(in1_block_num_tiles);
488
+ }
489
+
490
+ #ifdef FUSE_BIAS
491
+ #ifdef PACK_RELU
492
+ // if last block we pack the final result with relu enabled
493
+ PACK((llk_pack_relu_config(ReluConfig::zero())));
494
+ #endif
495
+ #if defined FP32_DEST_ACC_EN or defined PACKER_L1_ACC
496
+ PACK((pack_reconfig_data_format(out_dfb_id)));
497
+ #endif
498
+ #ifdef PACKER_L1_ACC
499
+ PACK((llk_pack_reconfig_l1_acc(0)));
500
+ #endif
501
+ reconfig_data_format(in1_dfb_id, mm_partials_dfb_id, in0_dfb_id, bias_dfb_id);
502
+ if constexpr (row_broadcast_bias) {
503
+ add_bcast_rows_init(mm_partials_dfb_id, bias_dfb_id);
504
+ } else {
505
+ add_init(mm_partials_dfb_id, bias_dfb_id);
506
+ }
507
+ // Reader only pushes bias once when num_blocks_w_dim == 1;
508
+ // the tiles stay in the CB for reuse across bh/batch iterations.
509
+ if ((b == 0 && bh == 0) || num_blocks_w_dim > 1) {
510
+ bias_dfb.wait_front(bias_ntiles);
511
+ }
512
+ #ifdef MAST3R_RESID_CB
513
+ // nothing produces the sharded residual CB: publish its (already resident) tiles once
514
+ resid_dfb.reserve_back(out_block_num_tiles);
515
+ resid_dfb.push_back(out_block_num_tiles);
516
+ resid_dfb.wait_front(out_block_num_tiles);
517
+ #endif
518
+ for (uint32_t in0_subblock = 0; in0_subblock < in0_num_subblocks; in0_subblock++) {
519
+ int in1_index_subblock_offset = 0;
520
+ for (uint32_t in1_subblock = 0; in1_subblock < in1_num_subblocks; in1_subblock++) {
521
+ // See matmul stage: the last in1 subblock has padded lanes whose bias tile was
522
+ // never pushed by the reader. Redirect those out-of-range bias_tile_idx reads to
523
+ // tile 0 of cb_bias to keep them in-bounds; the resulting padded output columns
524
+ // are dropped by the writer.
525
+ const bool is_last_in1_subblock_padded =
526
+ last_subblock_padded && (in1_subblock == in1_num_subblocks - 1);
527
+ // Redundant wait since we know data was just pushed
528
+ mm_partials_dfb.wait_front(out_subblock_num_tiles);
529
+ tile_regs_acquire();
530
+ for (uint32_t i = 0, j = 0; j < out_subblock_h; j++) {
531
+ #ifdef BIAS_FULL_BLOCK
532
+ // The bias CB holds a full [M, N] tile block. m_tile is this output tile's
533
+ // row within that block; bias_tile_idx is the position of the matching bias
534
+ // tile in the CB (row m_tile, column in1_index_subblock_offset). Only
535
+ // matmul_multicore_reuse_optimized loads the full block; other callers load a
536
+ // single bias row and use the N-only index below.
537
+ const uint32_t m_tile = in0_subblock * out_subblock_h + j;
538
+ uint32_t bias_tile_idx = m_tile * in1_block_w + in1_index_subblock_offset;
539
+ #else
540
+ uint32_t bias_tile_idx = in1_index_subblock_offset;
541
+ #endif
542
+ for (uint32_t k = 0; k < out_subblock_w; k++, i++) {
543
+ const uint32_t safe_bias_tile_idx =
544
+ (is_last_in1_subblock_padded && k >= last_subblock_w_valid)
545
+ ? 0u // Padded output columns with tile 0 of cb_bias added are
546
+ : bias_tile_idx; // dropped by the writer.
547
+
548
+ if constexpr (row_broadcast_bias) {
549
+ add_tiles_bcast_rows(mm_partials_dfb_id, bias_dfb_id, i, safe_bias_tile_idx, i);
550
+ } else {
551
+ add_tiles(mm_partials_dfb_id, bias_dfb_id, i, safe_bias_tile_idx, i);
552
+ }
553
+ bias_tile_idx++;
554
+ }
555
+ }
556
+ #ifdef MAST3R_RESID_CB
557
+ // dst (fp32: partials + bias) += resid tile; DST -> SrcA, resid -> SrcB (bf16, same as bias)
558
+ add_reuse_dest_init<EltwiseBinaryReuseDestType::DEST_TO_SRCA>(resid_dfb_id);
559
+ for (uint32_t i = 0, j = 0; j < out_subblock_h; j++) {
560
+ const uint32_t r0 = (in0_subblock * out_subblock_h + j) * in1_block_w + in1_index_subblock_offset;
561
+ for (uint32_t k = 0; k < out_subblock_w; k++, i++) {
562
+ add_reuse_dest_tiles<EltwiseBinaryReuseDestType::DEST_TO_SRCA>(resid_dfb_id, r0 + k, i);
563
+ }
564
+ }
565
+ if constexpr (row_broadcast_bias) {
566
+ add_bcast_rows_init(mm_partials_dfb_id, bias_dfb_id);
567
+ } else {
568
+ add_init(mm_partials_dfb_id, bias_dfb_id);
569
+ }
570
+ #endif
571
+ tile_regs_commit();
572
+
573
+ mm_partials_dfb.pop_front(out_subblock_num_tiles);
574
+
575
+ // Pack out to output buffer
576
+ untilize_mode_out_dfb.reserve_back(out_subblock_num_tiles);
577
+
578
+ #ifdef SFPU_ACTIVATION
579
+ PACK(TTI_SEMWAIT(
580
+ p_stall::STALL_TDMA | p_stall::STALL_CFG,
581
+ semaphore::t6_sem(semaphore::MATH_PACK),
582
+ p_stall::STALL_ON_ZERO));
583
+
584
+ // Flip destination register offset for PACKER access
585
+ PACK(TT_SETC16(
586
+ DEST_TARGET_REG_CFG_MATH_Offset_ADDR32, ckernel::packer::get_packer_dest_offset()));
587
+
588
+ for (uint32_t i = 0; i < out_subblock_num_tiles; i++) {
589
+ ActivationApplyHelper<activation_type, activation_param0, activation_param1>::apply(i);
590
+ }
591
+
592
+ PACK(TTI_STALLWAIT(p_stall::STALL_PACK, p_stall::WAIT_SFPU));
593
+ #else
594
+ tile_regs_wait();
595
+ #endif
596
+ for (uint32_t i = 0; i < out_subblock_num_tiles; i++) {
597
+ pack_tile(i, untilize_mode_out_dfb_id);
598
+ }
599
+ tile_regs_release();
600
+ untilize_mode_out_dfb.push_back(out_subblock_num_tiles);
601
+
602
+ in1_index_subblock_offset += out_subblock_w;
603
+ }
604
+ }
605
+ if constexpr (num_blocks_w_dim > 1) {
606
+ bias_dfb.pop_front(bias_ntiles);
607
+ }
608
+ #endif // FUSE_BIAS
609
+ if constexpr (untilize_out) {
610
+ #ifdef PACK_RELU
611
+ PACK((llk_pack_relu_config(ReluConfig::none())));
612
+ #endif // PACK_RELU
613
+ #ifndef FUSE_BIAS
614
+ reconfig_data_format_srca(in1_dfb_id, mm_partials_dfb_id);
615
+ #if defined FP32_DEST_ACC_EN or defined PACKER_L1_ACC
616
+ PACK((pack_reconfig_data_format(out_dfb_id)));
617
+ #endif
618
+ #ifdef PACKER_L1_ACC
619
+ PACK((llk_pack_reconfig_l1_acc(0)));
620
+ #endif
621
+ #endif // FUSE_BIAS
622
+ pack_untilize_dest_init<out_subblock_w, out_block_w>(out_dfb_id);
623
+ copy_tile_to_dst_init_short(mm_partials_dfb_id);
624
+ for (uint32_t in0_subblock_i = 0; in0_subblock_i < in0_num_subblocks; ++in0_subblock_i) {
625
+ reblock_and_untilize<out_subblock_w, out_block_w>(
626
+ in1_num_subblocks, out_subblock_num_tiles, out_subblock_h, mm_partials_dfb_id, out_dfb_id);
627
+ }
628
+ pack_untilize_uninit(mm_partials_dfb_id);
629
+ }
630
+ if constexpr (batch > 1 || num_blocks_w_dim > 1 || num_blocks_h_dim > 1) {
631
+ #ifdef FUSE_BIAS
632
+ // reconfigure unpacker df for src A and src B
633
+ reconfig_data_format(mm_partials_dfb_id, in1_dfb_id, bias_dfb_id, in0_dfb_id);
634
+ #else
635
+ // reconfigure unpacker df for src A
636
+ reconfig_data_format_srca(mm_partials_dfb_id, in1_dfb_id);
637
+ #endif
638
+ // reconfigure init for matmul
639
+ matmul_block_init(
640
+ in0_dfb_id, in1_dfb_id, in1_transpose_tile, out_subblock_w, out_subblock_h, in0_block_w);
641
+ }
642
+ }
643
+ }
644
+ }
645
+ #ifdef FUSE_BIAS
646
+ // For num_blocks_w_dim == 1 the reader pushes bias once and the kernel holds it resident,
647
+ // reusing it across all batch/bh/block iterations without popping. Pop it once here, after the
648
+ // last use, so the CB is balanced. (For num_blocks_w_dim > 1 the per-block pop above already
649
+ // balances each re-pushed bias block.)
650
+ if constexpr (num_blocks_w_dim == 1) {
651
+ bias_dfb.pop_front(bias_ntiles);
652
+ }
653
+ #ifdef MAST3R_RESID_CB
654
+ resid_dfb.pop_front(out_block_num_tiles);
655
+ #endif
656
+ #endif
657
+ }
code/models/demos/mast3r/tt/kernels/mm_in0_heads_reader.cpp ADDED
@@ -0,0 +1,472 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
2
+ //
3
+ // SPDX-License-Identifier: Apache-2.0
4
+ // mast3r-p150 model-local copy of ttnn's reader_bmm_tile_layout_in0_sender_padding.cpp (tt-metal 8b98410e730) for
5
+ // MAST3R_OPT "pcat": the only change is the in0 page id remap below, so the projection matmul reads the SDPA output
6
+ // [B, H, N, dh] (TILE, interleaved) directly as the concatenated-heads matrix [B*N, H*dh] (no concat-heads program).
7
+ // in0 tile (m, k) of the [M = NB*N, K = H*dh] view -> tile ((b*H + h)*NT + n)*DT + d of the head-major tensor, with
8
+ // b = m / NT + B0, n = m % NT, h = k / DT, d = k % DT.
9
+
10
+ #include <stdint.h>
11
+
12
+ #include "api/dataflow/dataflow_api.h"
13
+ #include "api/debug/assert.h"
14
+ #include "hostdevcommon/common_values.hpp"
15
+ #include "ttnn/operations/ccl/kernel_common/worker_sync_utils.hpp"
16
+ #include "ttnn/operations/kernel_helper_functions/pad_tile.hpp"
17
+ #include "ckernel.h"
18
+ #include "ckernel_defs.h"
19
+ #include "api/dataflow/noc.h"
20
+ #include "api/dataflow/dataflow_buffer.h"
21
+ #include "api/dataflow/noc_semaphore.h"
22
+ #include "api/tensor/noc_traits.h"
23
+ #include "api/dataflow/endpoints.h"
24
+ #include "api/core_local_mem.h"
25
+
26
+ #ifndef MAST3R_HEADS
27
+ #error "MAST3R_HEADS / MAST3R_NT / MAST3R_DT / MAST3R_B0 defines required"
28
+ #endif
29
+ static inline __attribute__((always_inline)) uint32_t mast3r_in0_remap(uint32_t t) {
30
+ constexpr uint32_t KT = MAST3R_HEADS * MAST3R_DT;
31
+ const uint32_t m = t / KT;
32
+ const uint32_t k = t - m * KT;
33
+ const uint32_t mb = m / MAST3R_NT;
34
+ const uint32_t n = m - mb * MAST3R_NT;
35
+ const uint32_t h = k / MAST3R_DT;
36
+ const uint32_t d = k - h * MAST3R_DT;
37
+ return (((mb + MAST3R_B0) * MAST3R_HEADS + h) * MAST3R_NT + n) * MAST3R_DT + d;
38
+ }
39
+
40
+ void kernel_main() {
41
+ uint32_t rt_args_idx = 0;
42
+ // in0 tensor args
43
+ const uint32_t in0_tensor_addr = get_arg_val<uint32_t>(rt_args_idx++);
44
+ uint32_t in0_tensor_start_tile_id = get_arg_val<uint32_t>(rt_args_idx++);
45
+ // in0 mcast args
46
+ const uint32_t in0_mcast_dest_noc_start_x = get_arg_val<uint32_t>(rt_args_idx++);
47
+ const uint32_t in0_mcast_dest_noc_start_y = get_arg_val<uint32_t>(rt_args_idx++);
48
+ const uint32_t in0_mcast_dest_noc_end_x = get_arg_val<uint32_t>(rt_args_idx++);
49
+ const uint32_t in0_mcast_dest_noc_end_y = get_arg_val<uint32_t>(rt_args_idx++);
50
+
51
+ // padding args
52
+ const uint32_t last_block_h = get_arg_val<uint32_t>(rt_args_idx++);
53
+ // sparsity args
54
+ const uint32_t sparsity_addr = get_arg_val<uint32_t>(rt_args_idx++);
55
+
56
+ // COMPILE TIME ARGS
57
+ // in0 tensor args
58
+ constexpr uint32_t in0_tensor_stride_w = get_compile_time_arg_val(0);
59
+ constexpr uint32_t in0_tensor_stride_h = get_compile_time_arg_val(1);
60
+ constexpr uint32_t in0_tensor_next_inner_dim_block_stride = get_compile_time_arg_val(2);
61
+ constexpr uint32_t in0_tensor_next_h_dim_block_stride = get_compile_time_arg_val(3);
62
+ // in0 block args
63
+ constexpr uint32_t in0_block_w = get_compile_time_arg_val(4);
64
+ constexpr uint32_t in0_block_h = get_compile_time_arg_val(5);
65
+ constexpr uint32_t in0_block_num_tiles = get_compile_time_arg_val(6);
66
+ constexpr uint32_t in0_last_ktile_w = get_compile_time_arg_val(7);
67
+ constexpr uint32_t in0_last_ktile_h = get_compile_time_arg_val(8);
68
+
69
+ constexpr bool extract_shard_sub_blocks = (bool)get_compile_time_arg_val(9);
70
+ constexpr uint32_t shard_width_in_tiles = get_compile_time_arg_val(10);
71
+ constexpr uint32_t shard_height_in_tiles = get_compile_time_arg_val(11);
72
+ // in0/in1 common args
73
+ constexpr uint32_t num_blocks_inner_dim = get_compile_time_arg_val(12);
74
+ constexpr uint32_t num_blocks_w_dim = get_compile_time_arg_val(13);
75
+ constexpr uint32_t num_blocks_h_dim = get_compile_time_arg_val(14);
76
+ // in0 mcast args
77
+ constexpr uint32_t in0_mcast_num_dests = get_compile_time_arg_val(17);
78
+ constexpr uint32_t in0_mcast_num_cores = get_compile_time_arg_val(18);
79
+ // batch args
80
+ constexpr uint32_t MtKt = get_compile_time_arg_val(19); // if 0
81
+ constexpr uint32_t in0_B = get_compile_time_arg_val(20);
82
+ constexpr uint32_t in1_B = get_compile_time_arg_val(21);
83
+ constexpr uint32_t in0_reuse_in_CB = get_compile_time_arg_val(22);
84
+
85
+ // sparsity args
86
+
87
+ constexpr uint32_t batchB = get_compile_time_arg_val(23);
88
+ constexpr uint32_t sparsity_pagesize = get_compile_time_arg_val(24);
89
+ // Boolean that is set when input A is sparse. If set, both input A and B are assumed to be sparse.
90
+ // Based on the sparsity tensor, the corresponding batch in input A and B are skipped.
91
+ constexpr bool bcast_A = (bool)get_compile_time_arg_val(25);
92
+ // This boolean is set when the number of batches is only known at runtime, typically based on a sparsity tensor.
93
+ constexpr bool get_batch_from_reader = (bool)get_compile_time_arg_val(26);
94
+
95
+ constexpr bool fuse_op = (bool)get_compile_time_arg_val(27);
96
+
97
+ constexpr auto in0_args = TensorAccessorArgs<28>();
98
+
99
+ constexpr auto sparsity_args = TensorAccessorArgs<in0_args.next_compile_time_args_offset()>();
100
+
101
+ // Number of valid (non-zero sparsity) batches the receiver and compute kernels are configured to
102
+ // process. When nnz is supplied (get_batch_from_reader == false), those kernels loop exactly
103
+ // num_batch_compute times, while this sender only multicasts once per non-zero sparsity entry, i.e.
104
+ // count_nonzero(sparsity) times. The op silently requires count_nonzero(sparsity) == num_batch_compute;
105
+ // if they disagree, the receivers wait on multicasts that never come (or the sender waits on receivers
106
+ // that already finished) and the device deadlocks. count_nonzero(sparsity) is data-dependent and only
107
+ // known here at runtime, so we validate it on-device by counting the multicasts we actually issue and
108
+ // asserting the contract holds -- surfacing a loud assert (under watcher) instead of a silent hang.
109
+ // See https://github.com/tenstorrent/tt-metal/issues/45943.
110
+ [[maybe_unused]] constexpr uint32_t num_batch_compute =
111
+ get_compile_time_arg_val(sparsity_args.next_compile_time_args_offset());
112
+
113
+ // 0 is used to specify "INVALID" state, i.e. when the multicasted data has not been received by the receiver.
114
+ // 0x1 is used to specify "VALID" state, i.e. when the batch is valid.
115
+ // 0x2 is used to specify "IGNORE_BATCH" state, i.e. when the batch is not valid.
116
+ constexpr uint32_t IGNORE_BATCH = 0x2;
117
+
118
+ // When sparsity is disabled, we just loop once
119
+ constexpr uint32_t batchB_lim = batchB == 0 ? 1u : batchB;
120
+
121
+ // Indexed/gather mode: iterate only the num_active selected sparse groups (the ids the caller
122
+ // passed in the `indices` operand). Every iterated batch is valid -- no sparsity scan, no skip,
123
+ // no validity multicast -- so this sender never has to look at the id list itself: A is either
124
+ // broadcast (bcast_A) or already compact and advanced sequentially per iteration (!bcast_A).
125
+ // Every factory that builds this kernel passes "num_active"; only the sparse matmul factory ever
126
+ // sets it non-zero. 0 means not indexed, i.e. the unchanged dense sparsity-scan path.
127
+ constexpr uint32_t num_active = get_named_compile_time_arg_val("num_active");
128
+ constexpr bool use_indices = num_active > 0;
129
+ constexpr uint32_t batch_loop_lim = use_indices ? num_active : batchB_lim;
130
+
131
+ MatmulOpReceiver fused_op_receiver;
132
+ if constexpr (fuse_op) {
133
+ fused_op_receiver = MatmulOpReceiver(
134
+ true, /* wait_for_op_signal */
135
+ rt_args_idx,
136
+ num_blocks_inner_dim,
137
+ in0_block_w /* tiles_per_block (in the same dimension as tensor slice) */
138
+ );
139
+ }
140
+
141
+ constexpr uint32_t dfb_id_in0 = get_named_compile_time_arg_val("cb_in0");
142
+ constexpr uint32_t in0_single_tile_size_bytes = get_tile_size(dfb_id_in0);
143
+ // Tiles whose size is not a multiple of the DRAM alignment are padded to it in DRAM, and the
144
+ // interleaved in0 CB pages are sized to match (see the program factory). The NOC reads the
145
+ // unpadded tile of data into each padded slot, and tiles are laid out / multicast at the padded
146
+ // stride. No-op when already aligned. The sharded path keeps the natural (unpadded) stride.
147
+ constexpr uint32_t in0_aligned_tile_size_bytes =
148
+ (in0_single_tile_size_bytes + (DRAM_ALIGNMENT - 1)) & ~(DRAM_ALIGNMENT - 1);
149
+ #ifdef IN0_SHARDED
150
+ constexpr uint32_t in0_block_size_bytes = in0_block_num_tiles * in0_single_tile_size_bytes;
151
+ #else
152
+ constexpr uint32_t in0_block_size_bytes = in0_block_num_tiles * in0_aligned_tile_size_bytes;
153
+ #endif
154
+
155
+ Noc noc;
156
+ DataflowBuffer dfb_in0(dfb_id_in0);
157
+ Semaphore<> sender_sem(get_compile_time_arg_val(15));
158
+ Semaphore<> receiver_sem(get_compile_time_arg_val(16));
159
+
160
+ #ifdef IN0_SHARDED
161
+ // In case we need to send multiple blocks per shard, in0 sharded cb is cb2 and we extract the sub-blocks to cb0
162
+ constexpr uint32_t shard_read_stride = shard_width_in_tiles * in0_single_tile_size_bytes;
163
+ constexpr uint32_t shard_read_width = in0_single_tile_size_bytes * in0_block_w;
164
+ constexpr uint32_t shard_num_tiles = shard_width_in_tiles * shard_height_in_tiles;
165
+ constexpr uint32_t in0_tensor_next_h_dim_block_stride_bytes =
166
+ in0_tensor_next_h_dim_block_stride * in0_single_tile_size_bytes;
167
+
168
+ uint32_t noc_shard_read_start_addr = 0;
169
+ if constexpr (extract_shard_sub_blocks) {
170
+ constexpr uint32_t dfb_id_in2 =
171
+ get_named_compile_time_arg_val("cb_in0_sharded"); // in0 sharded cb if extract_shard_sub_blocks
172
+ DataflowBuffer dfb_in2(dfb_id_in2);
173
+ noc_shard_read_start_addr = dfb_in2.get_read_ptr();
174
+ }
175
+
176
+ #else
177
+ const auto s0 = TensorAccessor(in0_args, in0_tensor_addr);
178
+ #endif // IN0_SHARDED
179
+
180
+ // sparsity accessor
181
+ constexpr uint32_t dfb_id_sparsity = get_named_compile_time_arg_val("cb_sparsity");
182
+ DataflowBuffer dfb_sparsity(dfb_id_sparsity);
183
+ const auto s_sparsity = TensorAccessor(sparsity_args, sparsity_addr);
184
+
185
+ #ifndef SKIP_MCAST
186
+ // Set ur local VALID value, to be mcasted to destinations flag address after the data has been mcasted
187
+ receiver_sem.set(VALID);
188
+ // local address that will be atomically incremented by mcast receivers, to know when all receivers are ready
189
+ // to receive the mcast
190
+
191
+ #ifdef IN0_SHARDED
192
+ uint32_t in0_start_address = dfb_in0.get_write_ptr();
193
+ #endif // IN0_SHARDED
194
+ #endif // SKIP_MCAST
195
+
196
+ uint32_t l1_write_addr_sparsity = 0;
197
+ if constexpr (batchB > 0 && !use_indices) {
198
+ dfb_sparsity.reserve_back(1);
199
+ l1_write_addr_sparsity = dfb_sparsity.get_write_ptr();
200
+ }
201
+
202
+ // Counts the in0 multicasts actually issued (one per non-zero sparsity entry). Used to validate
203
+ // count_nonzero(sparsity) == num_batch_compute when nnz is supplied (see num_batch_compute above).
204
+ [[maybe_unused]] uint32_t num_valid_batches = 0;
205
+
206
+ for (uint32_t b = 0; b < in0_B; ++b) {
207
+ if constexpr (batchB > 0 && !use_indices) {
208
+ noc.async_read(s_sparsity, dfb_sparsity, sparsity_pagesize, {.page_id = b}, {.offset_bytes = 0});
209
+ noc.async_read_barrier();
210
+ }
211
+
212
+ for (uint32_t bB = 0; bB < batch_loop_lim; ++bB) {
213
+ if constexpr (batchB > 0 && !use_indices) {
214
+ volatile auto is_batch_valid =
215
+ ((reinterpret_cast<volatile tt_l1_ptr uint16_t*>(l1_write_addr_sparsity))[bB]) != 0;
216
+
217
+ if constexpr (get_batch_from_reader) {
218
+ #ifndef SKIP_MCAST
219
+ // First broadcast this to other cores
220
+ sender_sem.wait(in0_mcast_num_dests);
221
+ sender_sem.set(0);
222
+ receiver_sem.set(is_batch_valid ? VALID : IGNORE_BATCH);
223
+ receiver_sem.set_multicast(
224
+ noc,
225
+ in0_mcast_dest_noc_start_x,
226
+ in0_mcast_dest_noc_start_y,
227
+ in0_mcast_dest_noc_end_x,
228
+ in0_mcast_dest_noc_end_y,
229
+ in0_mcast_num_cores);
230
+ noc.async_writes_flushed();
231
+ // Reset the semaphore value to VALID
232
+ receiver_sem.set(VALID);
233
+ #endif // SKIP_MCAST
234
+
235
+ // We need to pass the value to compute cores regardless of the value of is_batch_valid
236
+ ckernel::mailbox_write(ckernel::ThreadId::UnpackThreadId, is_batch_valid);
237
+ ckernel::mailbox_write(ckernel::ThreadId::MathThreadId, is_batch_valid);
238
+ ckernel::mailbox_write(ckernel::ThreadId::PackThreadId, is_batch_valid);
239
+ }
240
+
241
+ if (!is_batch_valid) {
242
+ if constexpr (!bcast_A) {
243
+ in0_tensor_start_tile_id += MtKt;
244
+ }
245
+ continue;
246
+ }
247
+
248
+ // This is a valid (non-zero) batch that we are about to multicast. When nnz was supplied,
249
+ // catch count_nonzero(sparsity) > num_batch_compute here, before the sender blocks below
250
+ // waiting on receivers that have already finished their num_batch_compute iterations.
251
+ if constexpr (!get_batch_from_reader) {
252
+ ++num_valid_batches;
253
+ ASSERT(num_valid_batches <= num_batch_compute);
254
+ }
255
+ }
256
+
257
+ #ifdef IN0_SHARDED
258
+ uint32_t in0_tensor_current_h_dim_block_start_addr = noc_shard_read_start_addr;
259
+ #endif // IN0_SHARDED
260
+ uint32_t in0_tensor_current_h_dim_block_tile_id = in0_tensor_start_tile_id;
261
+ for (uint32_t bh = 0; bh < num_blocks_h_dim; ++bh) {
262
+ for (uint32_t bw = 0; bw < num_blocks_w_dim; ++bw) {
263
+ #ifdef IN0_SHARDED
264
+ uint32_t in0_tensor_current_inner_dim_block_start_addr = in0_tensor_current_h_dim_block_start_addr;
265
+ #endif // IN0_SHARDED
266
+ uint32_t in0_tensor_current_inner_dim_block_start_tile_id = in0_tensor_current_h_dim_block_tile_id;
267
+ for (uint32_t block = 0; block < num_blocks_inner_dim; ++block) {
268
+ if constexpr (fuse_op) {
269
+ fused_op_receiver.update_current_block_start_tile_id(
270
+ block, in0_tensor_current_inner_dim_block_start_tile_id, in0_tensor_start_tile_id);
271
+ }
272
+
273
+ // Operand 0
274
+ // Common for sharded and interleaved paths
275
+ dfb_in0.reserve_back(in0_block_num_tiles);
276
+ #ifndef IN0_SHARDED
277
+
278
+ uint32_t in0_write_offset = 0;
279
+
280
+ #ifndef SKIP_MCAST
281
+ uint32_t in0_start_address =
282
+ dfb_in0.get_write_ptr(); // copy start address of block, to be used for mcasting
283
+ #endif // SKIP_MCAST
284
+
285
+ // Copy in0 block into CB, as the default kernel
286
+ uint32_t in0_tensor_row_start_tile_id = in0_tensor_current_inner_dim_block_start_tile_id;
287
+ for (uint32_t h = 0; h < in0_block_h; ++h) {
288
+ uint32_t in0_tensor_tile_id = in0_tensor_row_start_tile_id;
289
+ for (uint32_t w = 0; w < in0_block_w; ++w) {
290
+ if (bh < num_blocks_h_dim - 1 || h < last_block_h) {
291
+ noc.async_read(
292
+ s0,
293
+ dfb_in0,
294
+ in0_single_tile_size_bytes,
295
+ {.page_id = mast3r_in0_remap(in0_tensor_tile_id)},
296
+ {.offset_bytes = in0_write_offset});
297
+ }
298
+
299
+ // Zero out padded regions for the very last tile
300
+ if constexpr (in0_last_ktile_w > 0) {
301
+ if ((block == num_blocks_inner_dim - 1) && (w == in0_block_w - 1)) {
302
+ noc.async_read_barrier();
303
+ constexpr DataFormat in0_data_format = get_dataformat(dfb_id_in0);
304
+ pad_last_ktile<in0_data_format, in0_last_ktile_w>(
305
+ dfb_in0.get_write_ptr() + in0_write_offset);
306
+ }
307
+ }
308
+ if constexpr (in0_last_ktile_h > 0) {
309
+ if ((block == num_blocks_inner_dim - 1) && (w == in0_block_w - 1)) {
310
+ noc.async_read_barrier();
311
+ constexpr DataFormat in0_data_format = get_dataformat(dfb_id_in0);
312
+ pad_last_transposed_ktile<in0_data_format, in0_last_ktile_h>(
313
+ dfb_in0.get_write_ptr() + in0_write_offset);
314
+ }
315
+ }
316
+
317
+ in0_write_offset += in0_aligned_tile_size_bytes;
318
+ in0_tensor_tile_id += in0_tensor_stride_w;
319
+ }
320
+ in0_tensor_row_start_tile_id += in0_tensor_stride_h;
321
+ }
322
+ in0_tensor_current_inner_dim_block_start_tile_id += in0_tensor_next_inner_dim_block_stride;
323
+
324
+ // Barrier! make sure the reads are done
325
+ noc.async_read_barrier();
326
+ #else
327
+ if constexpr (extract_shard_sub_blocks) {
328
+ uint32_t l1_write_addr_in0 = dfb_in0.get_write_ptr();
329
+
330
+ #ifndef SKIP_MCAST
331
+ in0_start_address =
332
+ l1_write_addr_in0; // copy start address of block, to be used for mcasting
333
+ #endif // SKIP_MCAST
334
+
335
+ UnicastEndpoint self_ep;
336
+ uint32_t noc_shard_read_l1_addr = in0_tensor_current_inner_dim_block_start_addr;
337
+
338
+ for (uint32_t i = 0; i < in0_block_h; i++) {
339
+ noc.async_read(
340
+ self_ep,
341
+ CoreLocalMem<uint32_t>(l1_write_addr_in0),
342
+ shard_read_width,
343
+ {.noc_x = my_x[0], .noc_y = my_y[0], .addr = noc_shard_read_l1_addr},
344
+ {});
345
+
346
+ l1_write_addr_in0 += shard_read_width;
347
+ noc_shard_read_l1_addr += shard_read_stride;
348
+ }
349
+
350
+ in0_tensor_current_inner_dim_block_start_addr += shard_read_width;
351
+ noc.async_read_barrier();
352
+ }
353
+
354
+ {
355
+ constexpr DataFormat in0_data_format = get_dataformat(dfb_id_in0);
356
+ uint32_t in0_pad_base_addr = dfb_in0.get_write_ptr();
357
+ if constexpr (in0_last_ktile_w > 0) {
358
+ if ((block == num_blocks_inner_dim - 1)) {
359
+ for (uint32_t h = 0; h < in0_block_h; ++h) {
360
+ auto ptr = in0_pad_base_addr +
361
+ (h * in0_block_w + in0_block_w - 1) * in0_single_tile_size_bytes;
362
+ pad_last_ktile<in0_data_format, in0_last_ktile_w>(ptr);
363
+ }
364
+ }
365
+ }
366
+ if constexpr (in0_last_ktile_h > 0) {
367
+ if ((block == num_blocks_inner_dim - 1)) {
368
+ for (uint32_t w = 0; w < in0_block_w; ++w) {
369
+ auto ptr = in0_pad_base_addr +
370
+ ((in0_block_h - 1) * in0_block_w + w) * in0_single_tile_size_bytes;
371
+ pad_last_transposed_ktile<in0_data_format, in0_last_ktile_h>(ptr);
372
+ }
373
+ }
374
+ }
375
+ }
376
+ #endif // IN0_SHARDED
377
+
378
+ #ifndef SKIP_MCAST
379
+ // wait until all in0 mcast destinations have atomically incremented the in0 semaphore_addr
380
+ // (i.e. its value should be in0_mcast_num_dests), then reset the semaphore_addr value back to
381
+ // zero for the next block
382
+ sender_sem.wait(in0_mcast_num_dests);
383
+ sender_sem.set(0);
384
+
385
+ // Now we have the block in the CB address, we can mcast to dests!
386
+ MulticastEndpoint mcast_dst;
387
+ // num_dests must not include source, since we are NOT really doing a local copy!
388
+ noc.async_write_multicast(
389
+ CoreLocalMem<uint32_t>(in0_start_address),
390
+ mcast_dst,
391
+ in0_block_size_bytes,
392
+ in0_mcast_num_cores,
393
+ {},
394
+ {.noc_x_start = in0_mcast_dest_noc_start_x,
395
+ .noc_y_start = in0_mcast_dest_noc_start_y,
396
+ .noc_x_end = in0_mcast_dest_noc_end_x,
397
+ .noc_y_end = in0_mcast_dest_noc_end_y,
398
+ .addr = in0_start_address},
399
+ true);
400
+
401
+ // Note: no need for write barrier, since these two multicasts are done on the same noc id, same
402
+ // vc, same cmd_buf Also, this only works because we are setting VCs statically (using
403
+ // NOC_CMD_STATIC_VC).
404
+ #ifdef ARCH_BLACKHOLE
405
+ // On Blackhole the flush is needed because NoC latency is higher than L1 <-> RISCV
406
+ // latency which means data could be changed before write is issued.
407
+ noc.async_writes_flushed();
408
+ #endif // ARCH_BLACKHOLE
409
+
410
+ // We should also multicast the flag to destinations
411
+ // num_dests must not include source, since we are NOT really doing a local copy!
412
+ receiver_sem.set_multicast(
413
+ noc,
414
+ in0_mcast_dest_noc_start_x,
415
+ in0_mcast_dest_noc_start_y,
416
+ in0_mcast_dest_noc_end_x,
417
+ in0_mcast_dest_noc_end_y,
418
+ in0_mcast_num_cores);
419
+ #endif // SKIP_MCAST
420
+
421
+ // Common for sharded and interleaved paths
422
+ dfb_in0.push_back(in0_block_num_tiles);
423
+ }
424
+ }
425
+ #ifdef IN0_SHARDED
426
+ in0_tensor_current_h_dim_block_start_addr += in0_tensor_next_h_dim_block_stride_bytes;
427
+ #endif // IN0_SHARDED
428
+ in0_tensor_current_h_dim_block_tile_id += in0_tensor_next_h_dim_block_stride;
429
+ }
430
+
431
+ if constexpr (!bcast_A) {
432
+ in0_tensor_start_tile_id += MtKt;
433
+ }
434
+ }
435
+
436
+ if constexpr (bcast_A) {
437
+ in0_tensor_start_tile_id += MtKt;
438
+ }
439
+
440
+ // this is an optimization for the case when in0 is [1, 1, M, K] and in1 is [1, H, K, N], i.e. when in0_B ==
441
+ // 1 and in1_B > 1 in this case we originally had to replicate the in0 block for each batch, but with this
442
+ // optimization we just read the tensor slice once into each core's L1 and keep it there for all weight
443
+ // batches since the needed in0 data is already in L1 after batch 0, we can just move read pointer for this
444
+ // CB so compute kernel thinks it has new data
445
+ if (in0_reuse_in_CB) {
446
+ for (uint32_t fake_batch = 0; fake_batch < in1_B - in0_B; ++fake_batch) {
447
+ for (uint32_t blk = 0; blk < num_blocks_inner_dim; ++blk) {
448
+ dfb_in0.reserve_back(in0_block_num_tiles);
449
+ dfb_in0.push_back(in0_block_num_tiles);
450
+ }
451
+ }
452
+ }
453
+ }
454
+ noc.async_write_barrier();
455
+
456
+ // When nnz was supplied, the receiver and compute kernels loop exactly num_batch_compute times.
457
+ // If we issued fewer multicasts than that (count_nonzero(sparsity) < num_batch_compute), those
458
+ // kernels are now waiting on multicasts that will never arrive and the device would deadlock.
459
+ // Fail loudly instead. See https://github.com/tenstorrent/tt-metal/issues/45943.
460
+ // (Indexed/gather mode never counts: it multicasts exactly num_batch_compute == num_active times
461
+ // by construction, since every iterated group is active.)
462
+ if constexpr (!get_batch_from_reader && batchB > 0 && !use_indices) {
463
+ ASSERT(num_valid_batches == num_batch_compute);
464
+ }
465
+
466
+ // For completeness, we empty the sparsity CB if it was reserved earlier
467
+ if constexpr (batchB > 0 && !use_indices) {
468
+ dfb_sparsity.push_back(1);
469
+ dfb_sparsity.wait_front(1);
470
+ dfb_sparsity.pop_front(1);
471
+ }
472
+ }
code/models/demos/mast3r/tt/kernels/phase_il_rm_reader.cpp ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): polyphase interleave on ROW_MAJOR pixel rows, reader.
3
+ // X[(2i+a)*wl + 2j+b, :] = Y_ab[i*w0 + j, :]; one page = one pixel row (C bf16 = ROW_BYTES). Unit = a block of
4
+ // PB consecutive output pixels of one output row. Pure data movement: exact.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ void kernel_main() {
9
+ constexpr uint32_t W0 = get_compile_time_arg_val(0); // source row width (pixels)
10
+ constexpr uint32_t PB = get_compile_time_arg_val(1); // pixels per unit (divides 2*W0)
11
+ constexpr uint32_t LOG2_ROW = get_compile_time_arg_val(2); // log2(row bytes)
12
+ constexpr uint32_t ROW = 1u << LOG2_ROW;
13
+ constexpr uint32_t BPR = (2 * W0) / PB; // units per output pixel row
14
+ uint32_t y_addr[4];
15
+ for (uint32_t p = 0; p < 4; ++p) {
16
+ y_addr[p] = get_arg_val<uint32_t>(p);
17
+ }
18
+ const uint32_t u_start = get_arg_val<uint32_t>(4);
19
+ const uint32_t u_count = get_arg_val<uint32_t>(5);
20
+ const InterleavedPow2AddrGen<SRC_DRAM> g0 = {.bank_base_address = y_addr[0], .log_base_2_of_page_size = LOG2_ROW};
21
+ const InterleavedPow2AddrGen<SRC_DRAM> g1 = {.bank_base_address = y_addr[1], .log_base_2_of_page_size = LOG2_ROW};
22
+ const InterleavedPow2AddrGen<SRC_DRAM> g2 = {.bank_base_address = y_addr[2], .log_base_2_of_page_size = LOG2_ROW};
23
+ const InterleavedPow2AddrGen<SRC_DRAM> g3 = {.bank_base_address = y_addr[3], .log_base_2_of_page_size = LOG2_ROW};
24
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
25
+ const uint32_t y = u / BPR;
26
+ const uint32_t x0 = (u - y * BPR) * PB;
27
+ const uint32_t i = y >> 1;
28
+ const bool a = (y & 1) != 0;
29
+ cb_reserve_back(0, 1);
30
+ uint32_t dst = get_write_ptr(0);
31
+ for (uint32_t x = x0; x < x0 + PB; ++x) {
32
+ const uint32_t page = i * W0 + (x >> 1);
33
+ uint64_t src;
34
+ if (!a) {
35
+ src = (x & 1) ? g1.get_noc_addr(page) : g0.get_noc_addr(page);
36
+ } else {
37
+ src = (x & 1) ? g3.get_noc_addr(page) : g2.get_noc_addr(page);
38
+ }
39
+ noc_async_read(src, dst, ROW);
40
+ dst += ROW;
41
+ }
42
+ noc_async_read_barrier();
43
+ cb_push_back(0, 1);
44
+ }
45
+ }
code/models/demos/mast3r/tt/kernels/phase_il_rm_writer.cpp ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): polyphase interleave on ROW_MAJOR pixel rows, writer.
3
+ #include <stdint.h>
4
+ #include "api/dataflow/dataflow_api.h"
5
+
6
+ void kernel_main() {
7
+ constexpr uint32_t PB = get_compile_time_arg_val(0);
8
+ constexpr uint32_t LOG2_ROW = get_compile_time_arg_val(1);
9
+ constexpr uint32_t ROW = 1u << LOG2_ROW;
10
+ const uint32_t out_addr = get_arg_val<uint32_t>(0);
11
+ const uint32_t u_start = get_arg_val<uint32_t>(1);
12
+ const uint32_t u_count = get_arg_val<uint32_t>(2);
13
+ const InterleavedPow2AddrGen<OUT_DRAM> g = {.bank_base_address = out_addr, .log_base_2_of_page_size = LOG2_ROW};
14
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
15
+ cb_wait_front(0, 1);
16
+ uint32_t src = get_read_ptr(0);
17
+ for (uint32_t p = u * PB; p < (u + 1) * PB; ++p) {
18
+ noc_async_write(src, g.get_noc_addr(p), ROW);
19
+ src += ROW;
20
+ }
21
+ noc_async_writes_flushed();
22
+ cb_pop_front(0, 1);
23
+ }
24
+ noc_async_write_barrier();
25
+ }
code/models/demos/mast3r/tt/kernels/phase_il_rm_writer_hs.cpp ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): polyphase interleave on ROW_MAJOR pixel rows, writer straight into a
3
+ // HEIGHT_SHARDED L1 output (MAST3R_OPT "pilhs": the head.2 phase convs' own input shard spec, no interleaved->sharded
4
+ // copy). A unit = PB consecutive pixel rows (one CB page); shards hold UPS units, so a unit never straddles two shards and
5
+ // is written with one contiguous NoC write. Common runtime args: [out_addr, noc_x0, noc_y0, noc_x1, noc_y1, ...] per shard.
6
+ #include <stdint.h>
7
+ #include "api/dataflow/dataflow_api.h"
8
+
9
+ void kernel_main() {
10
+ constexpr uint32_t PB = get_compile_time_arg_val(0);
11
+ constexpr uint32_t LOG2_ROW = get_compile_time_arg_val(1);
12
+ constexpr uint32_t UPS = get_compile_time_arg_val(2); // units per shard
13
+ constexpr uint32_t ROW = 1u << LOG2_ROW;
14
+ constexpr uint32_t UNIT_BYTES = PB * ROW;
15
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
16
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
17
+ const uint32_t out_addr = get_common_arg_val<uint32_t>(0);
18
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
19
+ const uint32_t shard = u / UPS;
20
+ const uint32_t off = (u - shard * UPS) * UNIT_BYTES;
21
+ const uint32_t nx = get_common_arg_val<uint32_t>(1 + 2 * shard);
22
+ const uint32_t ny = get_common_arg_val<uint32_t>(2 + 2 * shard);
23
+ cb_wait_front(0, 1);
24
+ noc_async_write(get_read_ptr(0), get_noc_addr(nx, ny, out_addr + off), UNIT_BYTES);
25
+ noc_async_writes_flushed();
26
+ cb_pop_front(0, 1);
27
+ }
28
+ noc_async_write_barrier();
29
+ }
code/models/demos/mast3r/tt/kernels/strip_gather.cpp ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "sgat": gathers the DPT ring-strip inputs straight from
3
+ // refinenet1's HEIGHT_SHARDED TILE output xs [H0*W0, 32*CT] (pixel rows): tb [2, 3, W0, C] (image rows 0..2 and H0-3..H0-1)
4
+ // and lr [2, 3, H0, C] (columns 0..2 and W0-3..W0-1, already in the transposed strip order), ROW_MAJOR pixel pages in
5
+ // DRAM. Per output pixel: 2 face-row reads (16 channels = 32 B each) per channel tile from the owning shard. Pure data
6
+ // movement (exact). Replaces S2I to DRAM + untilize + 4 slices + 2 concats + permute.
7
+ #include <stdint.h>
8
+ #include "api/dataflow/dataflow_api.h"
9
+
10
+ void kernel_main() {
11
+ constexpr uint32_t CT = get_compile_time_arg_val(0); // channel tiles
12
+ constexpr uint32_t PER = get_compile_time_arg_val(1); // tile rows per shard
13
+ constexpr uint32_t W0 = get_compile_time_arg_val(2);
14
+ constexpr uint32_t H0 = get_compile_time_arg_val(3);
15
+ constexpr uint32_t PAGE = CT * 64; // one pixel, all channels (bf16)
16
+ const uint32_t p_start = get_arg_val<uint32_t>(0);
17
+ const uint32_t p_count = get_arg_val<uint32_t>(1);
18
+ if (p_count == 0) {
19
+ return;
20
+ }
21
+ const uint32_t xs_addr = get_common_arg_val<uint32_t>(0);
22
+ const InterleavedAddrGen<true> gtb = {.bank_base_address = get_common_arg_val<uint32_t>(1), .page_size = PAGE};
23
+ const InterleavedAddrGen<true> glr = {.bank_base_address = get_common_arg_val<uint32_t>(2), .page_size = PAGE};
24
+ const uint32_t buf = get_write_ptr(0);
25
+ uint32_t l1 = buf;
26
+ for (uint32_t q = p_start; q < p_start + p_count; ++q) {
27
+ uint32_t pix;
28
+ if (q < 6 * W0) {
29
+ const uint32_t sj = q / W0, x = q - sj * W0, s = sj / 3, j = sj - s * 3;
30
+ pix = (s == 0 ? j : H0 - 3 + j) * W0 + x;
31
+ } else {
32
+ const uint32_t q2 = q - 6 * W0;
33
+ const uint32_t sj = q2 / H0, i = q2 - sj * H0, s = sj / 3, j = sj - s * 3;
34
+ pix = i * W0 + (s == 0 ? j : W0 - 3 + j);
35
+ }
36
+ const uint32_t R = pix >> 5, r = pix & 31;
37
+ const uint32_t shard = R / PER, lt = R - shard * PER;
38
+ const uint32_t nx = get_common_arg_val<uint32_t>(3 + 2 * shard);
39
+ const uint32_t ny = get_common_arg_val<uint32_t>(4 + 2 * shard);
40
+ const uint32_t foff = (r >= 16 ? 1024u : 0u) + (r & 15) * 32;
41
+ for (uint32_t c = 0; c < CT; ++c) {
42
+ const uint64_t src = get_noc_addr(nx, ny, xs_addr + (lt * CT + c) * 2048 + foff);
43
+ noc_async_read(src, l1 + c * 64, 32);
44
+ noc_async_read(src + 512, l1 + c * 64 + 32, 32);
45
+ }
46
+ l1 += PAGE;
47
+ }
48
+ noc_async_read_barrier();
49
+ l1 = buf;
50
+ for (uint32_t q = p_start; q < p_start + p_count; ++q) {
51
+ const uint64_t dst = q < 6 * W0 ? gtb.get_noc_addr(q) : glr.get_noc_addr(q - 6 * W0);
52
+ noc_async_write(l1, dst, PAGE);
53
+ l1 += PAGE;
54
+ }
55
+ noc_async_write_barrier();
56
+ }
code/models/demos/mast3r/tt/kernels/tail_il_reader.cpp ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): head-tail polyphase interleave, reader.
3
+ // T [16, hl*wl] ROW_MAJOR (row (2a+b)*4 + c = phase (a,b) channel c, column i*wl + j) -> page (c, i) of the output
4
+ // [4*hl, 2*W] ROW_MAJOR = hi-res rows 2i and 2i+1 of channel c. For page (c, i) this reads the 4 source segments
5
+ // T[(2a+b)*4 + c][i*wl : (i+1)*wl] (a, b in {0,1}; wl*2 bytes each) into one CB slot, order (a, b).
6
+ #include <stdint.h>
7
+ #include "api/dataflow/dataflow_api.h"
8
+
9
+ void kernel_main() {
10
+ constexpr uint32_t HL = get_compile_time_arg_val(0);
11
+ constexpr uint32_t WL = get_compile_time_arg_val(1);
12
+ constexpr uint32_t LOG2_TROW = get_compile_time_arg_val(2); // log2(T row bytes) = log2(hl*wl*2)
13
+ constexpr uint32_t SEG = WL * 2;
14
+ const uint32_t t_addr = get_arg_val<uint32_t>(0);
15
+ const uint32_t u_start = get_arg_val<uint32_t>(1);
16
+ const uint32_t u_count = get_arg_val<uint32_t>(2);
17
+ const InterleavedPow2AddrGen<SRC_DRAM> g = {.bank_base_address = t_addr, .log_base_2_of_page_size = LOG2_TROW};
18
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
19
+ const uint32_t c = u / HL;
20
+ const uint32_t i = u - c * HL;
21
+ cb_reserve_back(0, 1);
22
+ uint32_t dst = get_write_ptr(0);
23
+ for (uint32_t p = 0; p < 4; ++p) {
24
+ noc_async_read(g.get_noc_addr(p * 4 + c) + i * SEG, dst, SEG);
25
+ dst += SEG;
26
+ }
27
+ noc_async_read_barrier();
28
+ cb_push_back(0, 1);
29
+ }
30
+ }
code/models/demos/mast3r/tt/kernels/tail_il_writer.cpp ADDED
@@ -0,0 +1,42 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): head-tail polyphase interleave, writer.
3
+ // out[c][2i+a][2j+b] = T[(2a+b)*4 + c][i*wl + j]: 16-bit element interleave of the (a,0) / (a,1) segments, then one
4
+ // page write (2 hi-res rows). Pure data movement: exact.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ void kernel_main() {
9
+ constexpr uint32_t WL = get_compile_time_arg_val(0);
10
+ constexpr uint32_t LOG2_PAGE = get_compile_time_arg_val(1); // log2(out page bytes) = log2(4*wl*2)
11
+ constexpr uint32_t SEG = WL * 2;
12
+ const uint32_t out_addr = get_arg_val<uint32_t>(0);
13
+ const uint32_t u_start = get_arg_val<uint32_t>(1);
14
+ const uint32_t u_count = get_arg_val<uint32_t>(2);
15
+ const InterleavedPow2AddrGen<OUT_DRAM> g = {.bank_base_address = out_addr, .log_base_2_of_page_size = LOG2_PAGE};
16
+ uint32_t slot = 0;
17
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
18
+ cb_wait_front(0, 1);
19
+ const uint32_t src = get_read_ptr(0);
20
+ // two output staging slots (CB 1 holds 2 pages): wait for the write issued two pages ago before reuse
21
+ if (u - u_start >= 2) {
22
+ noc_async_writes_flushed();
23
+ }
24
+ const uint32_t dst = get_write_ptr(1) + slot * (1u << LOG2_PAGE);
25
+ for (uint32_t a = 0; a < 2; ++a) {
26
+ const uint32_t* s0 = reinterpret_cast<const uint32_t*>(src + (2 * a) * SEG);
27
+ const uint32_t* s1 = reinterpret_cast<const uint32_t*>(src + (2 * a + 1) * SEG);
28
+ uint32_t* d = reinterpret_cast<uint32_t*>(dst + a * 2 * SEG);
29
+ #pragma GCC unroll 4
30
+ for (uint32_t k = 0; k < WL / 2; ++k) {
31
+ const uint32_t w0 = s0[k];
32
+ const uint32_t w1 = s1[k];
33
+ d[2 * k] = (w0 & 0xFFFFu) | (w1 << 16);
34
+ d[2 * k + 1] = (w0 >> 16) | (w1 & 0xFFFF0000u);
35
+ }
36
+ }
37
+ cb_pop_front(0, 1);
38
+ noc_async_write(dst, g.get_noc_addr(u), 1u << LOG2_PAGE);
39
+ slot ^= 1;
40
+ }
41
+ noc_async_write_barrier();
42
+ }
code/models/demos/mast3r/tt/kernels/tailf_compute.cpp ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): fused head tail, compute. Per tile: ((((z0 + B) + z1) + z2) + z3),
3
+ // FPU adds with ttnn.add's CB formats (bf16, 16-bit dest), then a WH transpose (exact) so each phase-channel is one tile
4
+ // row. The z_p are nonzero only in their own 4 columns, so every used column gets exactly one rounding, z_p + bias --
5
+ // the value of ttnn's add chain + transpose + bias add.
6
+ #include <cstdint>
7
+ #include "api/compute/common.h"
8
+ #include "api/compute/eltwise_binary.h"
9
+ #include "api/compute/transpose.h"
10
+ #include "api/compute/compute_kernel_hw_startup.h"
11
+ #include "api/compute/cb_api.h"
12
+ #include "api/compute/pack.h"
13
+
14
+ ALWI void add_pass(uint32_t ca, uint32_t cb, uint32_t co, bool pop_b) {
15
+ cb_wait_front(ca, 1);
16
+ cb_wait_front(cb, 1);
17
+ cb_reserve_back(co, 1);
18
+ add_init(ca, cb);
19
+ tile_regs_acquire();
20
+ add_tiles(ca, cb, 0, 0, 0);
21
+ tile_regs_commit();
22
+ tile_regs_wait();
23
+ pack_tile(0, co);
24
+ tile_regs_release();
25
+ cb_push_back(co, 1);
26
+ cb_pop_front(ca, 1);
27
+ if (pop_b) {
28
+ cb_pop_front(cb, 1);
29
+ }
30
+ }
31
+
32
+ void kernel_main() {
33
+ const uint32_t n_tiles = get_arg_val<uint32_t>(0);
34
+ if (n_tiles == 0) {
35
+ return;
36
+ }
37
+ compute_kernel_hw_startup(0, 4, 24);
38
+ cb_wait_front(4, 1);
39
+ for (uint32_t u = 0; u < n_tiles; ++u) {
40
+ add_pass(0, 4, 24, false); // z0 + B (bias CB 4 stays resident)
41
+ add_pass(24, 1, 25, true);
42
+ add_pass(25, 2, 26, true);
43
+ add_pass(26, 3, 27, true);
44
+ cb_wait_front(27, 1);
45
+ cb_reserve_back(16, 1);
46
+ transpose_init(27);
47
+ tile_regs_acquire();
48
+ transpose_tile(27, 0, 0);
49
+ tile_regs_commit();
50
+ tile_regs_wait();
51
+ pack_tile(0, 16);
52
+ tile_regs_release();
53
+ cb_push_back(16, 1);
54
+ cb_pop_front(27, 1);
55
+ }
56
+ cb_pop_front(4, 1);
57
+ }
code/models/demos/mast3r/tt/kernels/tailf_reader.cpp ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): fused head tail, reader. Unit = half a low-res image row (TPH tiles
3
+ // of 32 pixels). Reads tile u of the 4 per-phase 1x1 outputs z_p ([hl*wl, 32] TILE, interleaved) into CB 0..3, and the
4
+ // row-replicated bias tile once into CB 4.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ void kernel_main() {
9
+ constexpr uint32_t TPH = get_compile_time_arg_val(0);
10
+ constexpr uint32_t LOG2_TILE = 11;
11
+ constexpr uint32_t TB = 1u << LOG2_TILE;
12
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
13
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
14
+ if (u_count == 0) {
15
+ return;
16
+ }
17
+ const InterleavedPow2AddrGen<Z_DRAM> g0 = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
18
+ const InterleavedPow2AddrGen<Z_DRAM> g1 = {.bank_base_address = get_common_arg_val<uint32_t>(1), .log_base_2_of_page_size = LOG2_TILE};
19
+ const InterleavedPow2AddrGen<Z_DRAM> g2 = {.bank_base_address = get_common_arg_val<uint32_t>(2), .log_base_2_of_page_size = LOG2_TILE};
20
+ const InterleavedPow2AddrGen<Z_DRAM> g3 = {.bank_base_address = get_common_arg_val<uint32_t>(3), .log_base_2_of_page_size = LOG2_TILE};
21
+ const InterleavedPow2AddrGen<B_DRAM> gb = {.bank_base_address = get_common_arg_val<uint32_t>(4), .log_base_2_of_page_size = LOG2_TILE};
22
+ cb_reserve_back(4, 1);
23
+ noc_async_read(gb.get_noc_addr(0), get_write_ptr(4), TB);
24
+ noc_async_read_barrier();
25
+ cb_push_back(4, 1);
26
+ for (uint32_t u = u_start * TPH; u < (u_start + u_count) * TPH; ++u) {
27
+ cb_reserve_back(0, 1);
28
+ cb_reserve_back(1, 1);
29
+ cb_reserve_back(2, 1);
30
+ cb_reserve_back(3, 1);
31
+ noc_async_read(g0.get_noc_addr(u), get_write_ptr(0), TB);
32
+ noc_async_read(g1.get_noc_addr(u), get_write_ptr(1), TB);
33
+ noc_async_read(g2.get_noc_addr(u), get_write_ptr(2), TB);
34
+ noc_async_read(g3.get_noc_addr(u), get_write_ptr(3), TB);
35
+ noc_async_read_barrier();
36
+ cb_push_back(0, 1);
37
+ cb_push_back(1, 1);
38
+ cb_push_back(2, 1);
39
+ cb_push_back(3, 1);
40
+ }
41
+ }
code/models/demos/mast3r/tt/kernels/tailf_writer.cpp ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op): fused head tail, writer. Unit = half of a low-res image row i
3
+ // (TPH tiles of 32 pixels). The compute kernel hands over the WH-transposed tile: row ch = (2a+b)*4 + c (phase (a,b),
4
+ // channel c), column = pixel jj. out[c][2i+a][2(j0+jj)+b] = T[(2a+b)*4+c][jj]: 16-bit interleave of rows (a,0,c) and
5
+ // (a,1,c) into the staged segments, then for each (c, a) one write of TPH*64 elements into output page (c, i) = hi-res
6
+ // rows 2i, 2i+1 of channel c. Pure data movement: exact.
7
+ #include <stdint.h>
8
+ #include "api/dataflow/dataflow_api.h"
9
+
10
+ void kernel_main() {
11
+ constexpr uint32_t TPH = get_compile_time_arg_val(0);
12
+ constexpr uint32_t WL = get_compile_time_arg_val(1);
13
+ constexpr uint32_t HL = get_compile_time_arg_val(2);
14
+ constexpr uint32_t LOG2_PAGE = get_compile_time_arg_val(3); // log2(4*wl*2)
15
+ constexpr uint32_t SEG = TPH * 64 * 2; // bytes of one (c, a) segment
16
+ constexpr uint32_t SLOT = 8 * SEG;
17
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
18
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
19
+ const uint32_t out_addr = get_arg_val<uint32_t>(2);
20
+ if (u_count == 0) {
21
+ return;
22
+ }
23
+ const InterleavedPow2AddrGen<OUT_DRAM> g = {.bank_base_address = out_addr, .log_base_2_of_page_size = LOG2_PAGE};
24
+ uint32_t slot = 0;
25
+ for (uint32_t hu = u_start; hu < u_start + u_count; ++hu) {
26
+ const uint32_t i = hu >> 1;
27
+ const uint32_t h = hu & 1;
28
+ if (hu - u_start >= 1) {
29
+ noc_async_writes_flushed(); // 2 staging slots: the writes of unit hu-2 are done
30
+ }
31
+ const uint32_t stage = get_write_ptr(8) + slot * SLOT; // CB 8: staging only (never pushed)
32
+ for (uint32_t t = 0; t < TPH; ++t) {
33
+ cb_wait_front(16, 1);
34
+ // transposed tile: row ch < 16 = face 0 row ch (pixels 0..15) + face 1 row ch (pixels 16..31)
35
+ const uint32_t* tsrc = reinterpret_cast<const uint32_t*>(get_read_ptr(16));
36
+ for (uint32_t a = 0; a < 2; ++a) {
37
+ for (uint32_t c = 0; c < 4; ++c) {
38
+ const uint32_t ch0 = (2 * a) * 4 + c;
39
+ const uint32_t ch1 = (2 * a + 1) * 4 + c;
40
+ uint32_t* d = reinterpret_cast<uint32_t*>(stage + (c * 2 + a) * SEG + t * 64 * 2);
41
+ #pragma GCC unroll 2
42
+ for (uint32_t f = 0; f < 2; ++f) {
43
+ const uint32_t* s0 = tsrc + f * 128 + ch0 * 8; // 16-element face row = 8 words
44
+ const uint32_t* s1 = tsrc + f * 128 + ch1 * 8;
45
+ uint32_t* dd = d + f * 16;
46
+ #pragma GCC unroll 8
47
+ for (uint32_t k = 0; k < 8; ++k) {
48
+ const uint32_t w0 = s0[k];
49
+ const uint32_t w1 = s1[k];
50
+ dd[2 * k] = (w0 & 0xFFFFu) | (w1 << 16);
51
+ dd[2 * k + 1] = (w0 >> 16) | (w1 & 0xFFFF0000u);
52
+ }
53
+ }
54
+ }
55
+ }
56
+ cb_pop_front(16, 1);
57
+ }
58
+ for (uint32_t c = 0; c < 4; ++c) {
59
+ const uint64_t pa = g.get_noc_addr(c * HL + i);
60
+ for (uint32_t a = 0; a < 2; ++a) {
61
+ noc_async_write(stage + (c * 2 + a) * SEG, pa + (a * 2 * WL + h * TPH * 64) * 2, SEG);
62
+ }
63
+ }
64
+ slot ^= 1;
65
+ }
66
+ noc_async_write_barrier();
67
+ }
code/models/demos/mast3r/tt/kernels/ups2_compute.cpp ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "tups": bilinear x2 upsample in TILE layout as exact
3
+ // tile matmuls. Output tile (Yr, t', c) = sum over the 2 source rows (vertical weight w in {0.25, 0.75}, folded into the
4
+ // constant tile) and the <= 2 source x-tiles of A_kind (32 out px x 32 in px, entries in {0.25, 0.75, 1} * w) x X_tile.
5
+ // in0 = A (SrcB, <= 4 significant bits), in1 = X (SrcA); HiFi4 + fp32 dest: every product and the <= 4-term sum are
6
+ // exact in fp32, one bf16 round at the pack. Constant tiles in CB 1: index 6 * wi + kind, wi 0 -> 0.25, 1 -> 0.75;
7
+ // kind 0 E_EDGE, 1 E_L, 2 E_R, 3 O_L, 4 O_R, 5 O_EDGE.
8
+ #include <cstdint>
9
+ #include "api/compute/common.h"
10
+ #include "api/compute/matmul.h"
11
+ #include "api/compute/compute_kernel_hw_startup.h"
12
+ #include "api/compute/cb_api.h"
13
+ #include "api/compute/pack.h"
14
+
15
+ constexpr uint32_t CB_X = 0, CB_A = 1, CB_O = 16;
16
+ constexpr uint32_t E_EDGE = 0, E_L = 1, E_R = 2, O_L = 3, O_R = 4, O_EDGE = 5;
17
+
18
+ void kernel_main() {
19
+ constexpr uint32_t NX = get_compile_time_arg_val(0);
20
+ constexpr uint32_t CT = get_compile_time_arg_val(1);
21
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
22
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
23
+ if (u_count == 0) {
24
+ return;
25
+ }
26
+ compute_kernel_hw_startup<SrcOrder::Reverse>(CB_A, CB_X, CB_O);
27
+ matmul_init(CB_A, CB_X);
28
+ cb_wait_front(CB_A, 12);
29
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
30
+ const uint32_t Yr = u / CT;
31
+ const uint32_t w0 = (Yr & 1) == 0 ? 0 : 6; // row 0 weight: 0.25 (even Yr) / 0.75 (odd Yr)
32
+ const uint32_t w1 = (Yr & 1) == 0 ? 6 : 0;
33
+ cb_wait_front(CB_X, 2 * NX);
34
+ for (uint32_t tp = 0; tp < 2 * NX; ++tp) {
35
+ const uint32_t m = tp >> 1;
36
+ cb_reserve_back(CB_O, 1);
37
+ tile_regs_acquire();
38
+ for (uint32_t r = 0; r < 2; ++r) {
39
+ const uint32_t wb = r == 0 ? w0 : w1;
40
+ const uint32_t ro = r * NX;
41
+ if ((tp & 1) == 0) {
42
+ if (m == 0) {
43
+ matmul_tiles(CB_A, CB_X, wb + E_EDGE, ro, 0);
44
+ } else {
45
+ matmul_tiles(CB_A, CB_X, wb + E_L, ro + m - 1, 0);
46
+ matmul_tiles(CB_A, CB_X, wb + E_R, ro + m, 0);
47
+ }
48
+ } else {
49
+ if (m == NX - 1) {
50
+ matmul_tiles(CB_A, CB_X, wb + O_EDGE, ro + m, 0);
51
+ } else {
52
+ matmul_tiles(CB_A, CB_X, wb + O_L, ro + m, 0);
53
+ matmul_tiles(CB_A, CB_X, wb + O_R, ro + m + 1, 0);
54
+ }
55
+ }
56
+ }
57
+ tile_regs_commit();
58
+ tile_regs_wait();
59
+ pack_tile(0, CB_O);
60
+ tile_regs_release();
61
+ cb_push_back(CB_O, 1);
62
+ }
63
+ cb_pop_front(CB_X, 2 * NX);
64
+ }
65
+ cb_pop_front(CB_A, 12);
66
+ }
code/models/demos/mast3r/tt/kernels/ups2_reader.cpp ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "tups": bilinear x2 upsample (half-pixel, clamped, ==
3
+ // ttnn.upsample) of a TILE [H*W, C] pixel-row tensor, reader. Work unit u = (output image row Yr, channel tile c),
4
+ // u = Yr * CT + c. Per unit: the 2 * NX input tiles of the two source image rows (clamped) at channel tile c, in the
5
+ // order [row0 x-tiles..., row1 x-tiles...] into CB 0. Once per core: the 12 constant interpolation tiles into CB 1.
6
+ #include <stdint.h>
7
+ #include "api/dataflow/dataflow_api.h"
8
+
9
+ void kernel_main() {
10
+ constexpr uint32_t NX = get_compile_time_arg_val(0); // input x-tiles per image row (W / 32)
11
+ constexpr uint32_t CT = get_compile_time_arg_val(1); // channel tiles
12
+ constexpr uint32_t H = get_compile_time_arg_val(2); // input image rows
13
+ constexpr uint32_t NA = 12;
14
+ constexpr uint32_t LOG2_TILE = 11;
15
+ constexpr uint32_t TB = 1u << LOG2_TILE;
16
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
17
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
18
+ if (u_count == 0) {
19
+ return;
20
+ }
21
+ const InterleavedPow2AddrGen<IN_DRAM> gx = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
22
+ const InterleavedPow2AddrGen<A_DRAM> ga = {.bank_base_address = get_common_arg_val<uint32_t>(1), .log_base_2_of_page_size = LOG2_TILE};
23
+ cb_reserve_back(1, NA);
24
+ uint32_t pa = get_write_ptr(1);
25
+ for (uint32_t i = 0; i < NA; ++i) {
26
+ noc_async_read(ga.get_noc_addr(i), pa, TB);
27
+ pa += TB;
28
+ }
29
+ noc_async_read_barrier();
30
+ cb_push_back(1, NA);
31
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
32
+ const uint32_t Yr = u / CT;
33
+ const uint32_t c = u - Yr * CT;
34
+ const uint32_t y = Yr >> 1;
35
+ uint32_t y0, y1;
36
+ if ((Yr & 1) == 0) {
37
+ y0 = y > 0 ? y - 1 : 0;
38
+ y1 = y;
39
+ } else {
40
+ y0 = y;
41
+ y1 = y + 1 < H ? y + 1 : H - 1;
42
+ }
43
+ cb_reserve_back(0, 2 * NX);
44
+ uint32_t p = get_write_ptr(0);
45
+ for (uint32_t xt = 0; xt < NX; ++xt) {
46
+ noc_async_read(gx.get_noc_addr((y0 * NX + xt) * CT + c), p, TB);
47
+ p += TB;
48
+ }
49
+ for (uint32_t xt = 0; xt < NX; ++xt) {
50
+ noc_async_read(gx.get_noc_addr((y1 * NX + xt) * CT + c), p, TB);
51
+ p += TB;
52
+ }
53
+ noc_async_read_barrier();
54
+ cb_push_back(0, 2 * NX);
55
+ }
56
+ }
code/models/demos/mast3r/tt/kernels/ups2_writer.cpp ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "tups": writer. Per unit u = (Yr, c): the 2 * NX output
3
+ // tiles of output image row Yr at channel tile c (output pixel-row tile R = Yr * 2 * NX + t').
4
+ #include <stdint.h>
5
+ #include "api/dataflow/dataflow_api.h"
6
+
7
+ void kernel_main() {
8
+ constexpr uint32_t NX = get_compile_time_arg_val(0);
9
+ constexpr uint32_t CT = get_compile_time_arg_val(1);
10
+ constexpr uint32_t LOG2_TILE = 11;
11
+ constexpr uint32_t TB = 1u << LOG2_TILE;
12
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
13
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
14
+ const InterleavedPow2AddrGen<OUT_DRAM> g = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
15
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
16
+ const uint32_t Yr = u / CT;
17
+ const uint32_t c = u - Yr * CT;
18
+ for (uint32_t tp = 0; tp < 2 * NX; ++tp) {
19
+ cb_wait_front(16, 1);
20
+ noc_async_write(get_read_ptr(16), g.get_noc_addr((Yr * 2 * NX + tp) * CT + c), TB);
21
+ noc_async_writes_flushed();
22
+ cb_pop_front(16, 1);
23
+ }
24
+ }
25
+ noc_async_write_barrier();
26
+ }
code/models/demos/mast3r/tt/kernels/ups2h_compute.cpp ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "tups" for W = 16: output tile (Yr, c) = A[w0, h0] x
3
+ // X(y0) + A[w1, h1] x X(y1), where h = y & 1 selects the half of the input tile that holds image row y and w the vertical
4
+ // weight (0.25 / 0.75, folded into the constant tile). Constant tile index 2 * wi + h. HiFi4 + fp32 dest (exact).
5
+ #include <cstdint>
6
+ #include "api/compute/common.h"
7
+ #include "api/compute/matmul.h"
8
+ #include "api/compute/compute_kernel_hw_startup.h"
9
+ #include "api/compute/cb_api.h"
10
+ #include "api/compute/pack.h"
11
+
12
+ constexpr uint32_t CB_X = 0, CB_A = 1, CB_O = 16;
13
+
14
+ void kernel_main() {
15
+ constexpr uint32_t CT = get_compile_time_arg_val(0);
16
+ constexpr uint32_t H = get_compile_time_arg_val(1);
17
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
18
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
19
+ if (u_count == 0) {
20
+ return;
21
+ }
22
+ compute_kernel_hw_startup<SrcOrder::Reverse>(CB_A, CB_X, CB_O);
23
+ matmul_init(CB_A, CB_X);
24
+ cb_wait_front(CB_A, 4);
25
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
26
+ const uint32_t Yr = u / CT;
27
+ const uint32_t y = Yr >> 1;
28
+ uint32_t y0, y1, wa, wb;
29
+ if ((Yr & 1) == 0) {
30
+ y0 = y > 0 ? y - 1 : 0;
31
+ y1 = y;
32
+ wa = 0; // 0.25
33
+ wb = 2; // 0.75
34
+ } else {
35
+ y0 = y;
36
+ y1 = y + 1 < H ? y + 1 : H - 1;
37
+ wa = 2;
38
+ wb = 0;
39
+ }
40
+ cb_wait_front(CB_X, 2);
41
+ cb_reserve_back(CB_O, 1);
42
+ tile_regs_acquire();
43
+ matmul_tiles(CB_A, CB_X, wa + (y0 & 1), 0, 0);
44
+ matmul_tiles(CB_A, CB_X, wb + (y1 & 1), 1, 0);
45
+ tile_regs_commit();
46
+ tile_regs_wait();
47
+ pack_tile(0, CB_O);
48
+ tile_regs_release();
49
+ cb_push_back(CB_O, 1);
50
+ cb_pop_front(CB_X, 2);
51
+ }
52
+ cb_pop_front(CB_A, 4);
53
+ }
code/models/demos/mast3r/tt/kernels/ups2h_reader.cpp ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "tups" for W = 16 (an input tile row holds 2 image rows):
3
+ // reader. Unit u = (output image row Yr, channel tile c) = Yr * CT + c: the input tiles holding the two (clamped) source
4
+ // image rows y0, y1 at channel tile c into CB 0 (2 tiles); once per core the 4 constant tiles into CB 1.
5
+ #include <stdint.h>
6
+ #include "api/dataflow/dataflow_api.h"
7
+
8
+ void kernel_main() {
9
+ constexpr uint32_t CT = get_compile_time_arg_val(0);
10
+ constexpr uint32_t H = get_compile_time_arg_val(1);
11
+ constexpr uint32_t NA = 4;
12
+ constexpr uint32_t LOG2_TILE = 11;
13
+ constexpr uint32_t TB = 1u << LOG2_TILE;
14
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
15
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
16
+ if (u_count == 0) {
17
+ return;
18
+ }
19
+ const InterleavedPow2AddrGen<IN_DRAM> gx = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
20
+ const InterleavedPow2AddrGen<A_DRAM> ga = {.bank_base_address = get_common_arg_val<uint32_t>(1), .log_base_2_of_page_size = LOG2_TILE};
21
+ cb_reserve_back(1, NA);
22
+ uint32_t pa = get_write_ptr(1);
23
+ for (uint32_t i = 0; i < NA; ++i) {
24
+ noc_async_read(ga.get_noc_addr(i), pa, TB);
25
+ pa += TB;
26
+ }
27
+ noc_async_read_barrier();
28
+ cb_push_back(1, NA);
29
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
30
+ const uint32_t Yr = u / CT;
31
+ const uint32_t c = u - Yr * CT;
32
+ const uint32_t y = Yr >> 1;
33
+ uint32_t y0, y1;
34
+ if ((Yr & 1) == 0) {
35
+ y0 = y > 0 ? y - 1 : 0;
36
+ y1 = y;
37
+ } else {
38
+ y0 = y;
39
+ y1 = y + 1 < H ? y + 1 : H - 1;
40
+ }
41
+ cb_reserve_back(0, 2);
42
+ const uint32_t p = get_write_ptr(0);
43
+ noc_async_read(gx.get_noc_addr((y0 >> 1) * CT + c), p, TB);
44
+ noc_async_read(gx.get_noc_addr((y1 >> 1) * CT + c), p + TB, TB);
45
+ noc_async_read_barrier();
46
+ cb_push_back(0, 2);
47
+ }
48
+ }
code/models/demos/mast3r/tt/kernels/ups2h_writer.cpp ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ // mast3r-p150 model-local kernel (ttnn.generic_op), MAST3R_OPT "tups" for W = 16: writer (one output tile per unit, the
3
+ // 32-pixel output image row Yr at channel tile c = tile Yr * CT + c).
4
+ #include <stdint.h>
5
+ #include "api/dataflow/dataflow_api.h"
6
+
7
+ void kernel_main() {
8
+ constexpr uint32_t LOG2_TILE = 11;
9
+ constexpr uint32_t TB = 1u << LOG2_TILE;
10
+ const uint32_t u_start = get_arg_val<uint32_t>(0);
11
+ const uint32_t u_count = get_arg_val<uint32_t>(1);
12
+ const InterleavedPow2AddrGen<OUT_DRAM> g = {.bank_base_address = get_common_arg_val<uint32_t>(0), .log_base_2_of_page_size = LOG2_TILE};
13
+ for (uint32_t u = u_start; u < u_start + u_count; ++u) {
14
+ cb_wait_front(16, 1);
15
+ noc_async_write(get_read_ptr(16), g.get_noc_addr(u), TB);
16
+ noc_async_writes_flushed();
17
+ cb_pop_front(16, 1);
18
+ }
19
+ noc_async_write_barrier();
20
+ }
code/models/demos/mast3r/tt/ttnn_dust3r.py CHANGED
The diff for this file is too large to render. See raw diff