Optimized build (2026-10-03): model call 87 -> 28 ms, pose path 38 ms
Browse filescode/ 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
- GPU_COMPARISON.md +43 -0
- OPT_BASELINE.md +228 -0
- OPT_REPORT.md +0 -0
- README.md +42 -20
- VERIFICATION_2026-10-03.md +85 -0
- code/bench_breakdown.py +147 -0
- code/models/demos/mast3r/postprocess.py +10 -3
- code/models/demos/mast3r/tt/fattn.py +130 -0
- code/models/demos/mast3r/tt/fused.py +215 -0
- code/models/demos/mast3r/tt/heads_rope.py +1138 -0
- code/models/demos/mast3r/tt/kernels/add2_compute.cpp +34 -0
- code/models/demos/mast3r/tt/kernels/add3_compute.cpp +62 -0
- code/models/demos/mast3r/tt/kernels/add3_reader.cpp +27 -0
- code/models/demos/mast3r/tt/kernels/add3_writer.cpp +23 -0
- code/models/demos/mast3r/tt/kernels/add3s_reader.cpp +40 -0
- code/models/demos/mast3r/tt/kernels/add3s_writer.cpp +37 -0
- code/models/demos/mast3r/tt/kernels/fattn_reader.cpp +108 -0
- code/models/demos/mast3r/tt/kernels/fsdpa/compute_common.hpp +2484 -0
- code/models/demos/mast3r/tt/kernels/fsdpa/compute_streaming.hpp +0 -0
- code/models/demos/mast3r/tt/kernels/fsdpa/sdpa.cpp +277 -0
- code/models/demos/mast3r/tt/kernels/heads_concat_reader.cpp +38 -0
- code/models/demos/mast3r/tt/kernels/heads_concat_writer.cpp +43 -0
- code/models/demos/mast3r/tt/kernels/heads_rope_compute.cpp +180 -0
- code/models/demos/mast3r/tt/kernels/heads_rope_reader.cpp +110 -0
- code/models/demos/mast3r/tt/kernels/heads_rope_writer.cpp +57 -0
- code/models/demos/mast3r/tt/kernels/ln_add_compute.cpp +436 -0
- code/models/demos/mast3r/tt/kernels/ln_add_writer.cpp +43 -0
- code/models/demos/mast3r/tt/kernels/ln_add_writer_rb.cpp +64 -0
- code/models/demos/mast3r/tt/kernels/ln_fast_compute.cpp +202 -0
- code/models/demos/mast3r/tt/kernels/ln_reader_split.cpp +209 -0
- code/models/demos/mast3r/tt/kernels/mast3r_gelu_poly.h +116 -0
- code/models/demos/mast3r/tt/kernels/mm_gelu_activation.hpp +123 -0
- code/models/demos/mast3r/tt/kernels/mm_gelu_compute.cpp +657 -0
- code/models/demos/mast3r/tt/kernels/mm_in0_heads_reader.cpp +472 -0
- code/models/demos/mast3r/tt/kernels/phase_il_rm_reader.cpp +45 -0
- code/models/demos/mast3r/tt/kernels/phase_il_rm_writer.cpp +25 -0
- code/models/demos/mast3r/tt/kernels/phase_il_rm_writer_hs.cpp +29 -0
- code/models/demos/mast3r/tt/kernels/strip_gather.cpp +56 -0
- code/models/demos/mast3r/tt/kernels/tail_il_reader.cpp +30 -0
- code/models/demos/mast3r/tt/kernels/tail_il_writer.cpp +42 -0
- code/models/demos/mast3r/tt/kernels/tailf_compute.cpp +57 -0
- code/models/demos/mast3r/tt/kernels/tailf_reader.cpp +41 -0
- code/models/demos/mast3r/tt/kernels/tailf_writer.cpp +67 -0
- code/models/demos/mast3r/tt/kernels/ups2_compute.cpp +66 -0
- code/models/demos/mast3r/tt/kernels/ups2_reader.cpp +56 -0
- code/models/demos/mast3r/tt/kernels/ups2_writer.cpp +26 -0
- code/models/demos/mast3r/tt/kernels/ups2h_compute.cpp +53 -0
- code/models/demos/mast3r/tt/kernels/ups2h_reader.cpp +48 -0
- code/models/demos/mast3r/tt/kernels/ups2h_writer.cpp +20 -0
- 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 (
|
| 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)
|
| 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 |
-
###
|
| 79 |
|
| 80 |
-
|
|
|
|
|
|
|
| 81 |
|---|---:|
|
| 82 |
-
|
|
| 83 |
-
|
|
| 84 |
-
|
|
| 85 |
-
|
|
| 86 |
-
|
|
| 87 |
-
|
|
| 88 |
-
|
|
| 89 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
|
| 91 |
### Caveats
|
| 92 |
|
| 93 |
-
-
|
| 94 |
-
-
|
| 95 |
-
-
|
| 96 |
-
-
|
|
|
|
|
|
|
|
|
|
| 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.
|
| 99 |
-
-
|
| 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 |
-
|
| 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
|
|
|