tt-model push diffusion-planner-p150 (container)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +13 -0
- OPT_BASELINE.md +26 -0
- OPT_REPORT.md +579 -30
- PYTHON.md +20 -19
- README.md +46 -44
- SERVING.md +20 -18
- VERIFICATION_OPT_2026-10-11.md +157 -0
- build_info.json +6 -6
- code/PYTHON.md +20 -19
- code/scripts/bench.py +24 -9
- code/scripts/profile_ops.py +7 -4
- code/tt_diffusion_planner/__init__.py +1 -1
- code/tt_diffusion_planner/host/features.py +37 -10
- code/tt_diffusion_planner/host/normalize.py +63 -1
- code/tt_diffusion_planner/host/pipeline.py +10 -3
- code/tt_diffusion_planner/host/postprocess.py +68 -1
- code/tt_diffusion_planner/tests/test_bundle_host.py +5 -1
- code/tt_diffusion_planner/tests/test_grid_fit_host.py +23 -0
- code/tt_diffusion_planner/tests/test_tt_params_host.py +28 -0
- code/tt_diffusion_planner/tt/attention.py +96 -0
- code/tt_diffusion_planner/tt/config.py +102 -0
- code/tt_diffusion_planner/tt/decoder.py +85 -10
- code/tt_diffusion_planner/tt/encoder.py +94 -13
- code/tt_diffusion_planner/tt/fattn_kernel.py +140 -0
- code/tt_diffusion_planner/tt/inputs.py +6 -3
- code/tt_diffusion_planner/tt/kcat_kernel.py +109 -0
- code/tt_diffusion_planner/tt/kernels/README.md +30 -8
- code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp +222 -0
- code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp +95 -0
- code/tt_diffusion_planner/tt/kernels/fattn_writer.cpp +82 -0
- code/tt_diffusion_planner/tt/kernels/kcat_compute.cpp +74 -0
- code/tt_diffusion_planner/tt/kernels/kcat_reader.cpp +24 -0
- code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp +66 -0
- code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp +282 -0
- code/tt_diffusion_planner/tt/kernels/ln32_reader.cpp +123 -0
- code/tt_diffusion_planner/tt/kernels/ln32_sfpu.h +54 -0
- code/tt_diffusion_planner/tt/kernels/ln32_writer.cpp +84 -0
- code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp +214 -0
- code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp +130 -0
- code/tt_diffusion_planner/tt/kernels/ln32s_writer.cpp +161 -0
- code/tt_diffusion_planner/tt/kernels/smask_compute.cpp +69 -0
- code/tt_diffusion_planner/tt/kernels/smask_reader.cpp +50 -0
- code/tt_diffusion_planner/tt/kernels/smask_writer.cpp +31 -0
- code/tt_diffusion_planner/tt/kernels/smsm_compute.cpp +179 -0
- code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp +83 -0
- code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp +43 -0
- code/tt_diffusion_planner/tt/layers.py +597 -15
- code/tt_diffusion_planner/tt/ln_kernel.py +285 -0
- code/tt_diffusion_planner/tt/model.py +103 -12
- code/tt_diffusion_planner/tt/smask_kernel.py +98 -0
.gitattributes
CHANGED
|
@@ -57,3 +57,16 @@ media/dp_nuscenes_scene-0103_kf12_cam_front_tt_NC.jpg filter=lfs diff=lfs merge=
|
|
| 57 |
media/dp_nuscenes_scene-0757_kf11_bev_tt_NC.png filter=lfs diff=lfs merge=lfs -text
|
| 58 |
media/dp_nuscenes_scene-0916_kf11_bev_tt_NC.png filter=lfs diff=lfs merge=lfs -text
|
| 59 |
media/dp_straight_road_tt_vs_cpu.png filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
media/dp_nuscenes_scene-0757_kf11_bev_tt_NC.png filter=lfs diff=lfs merge=lfs -text
|
| 58 |
media/dp_nuscenes_scene-0916_kf11_bev_tt_NC.png filter=lfs diff=lfs merge=lfs -text
|
| 59 |
media/dp_straight_road_tt_vs_cpu.png filter=lfs diff=lfs merge=lfs -text
|
| 60 |
+
image/blobs/sha256/0505b13139e86b5156c958e759de0b39da1b9d9ae252a64b2f6e7a05a9a3e350 filter=lfs diff=lfs merge=lfs -text
|
| 61 |
+
image/blobs/sha256/1315e7e97a55bd5693b68844d50689456daceaab85d9d103c19e578551169924 filter=lfs diff=lfs merge=lfs -text
|
| 62 |
+
image/blobs/sha256/20475a01a7476e219fb670109972618cc19d3b4b136ef19e1d2ac110271457a2 filter=lfs diff=lfs merge=lfs -text
|
| 63 |
+
image/blobs/sha256/312f7c835935adbd4db8a7fc44867f1fb75883fa686a3c95177fe7f6dc43e34b filter=lfs diff=lfs merge=lfs -text
|
| 64 |
+
image/blobs/sha256/354c8aad2b5a35d4683b0174c5443906ec2a5cf58b6cb3f5540d3b9cacaef0ab filter=lfs diff=lfs merge=lfs -text
|
| 65 |
+
image/blobs/sha256/3678157ea229048016d1006e8432030308540aaa7ddfb0520bdeeba9c833dd93 filter=lfs diff=lfs merge=lfs -text
|
| 66 |
+
image/blobs/sha256/575a85c6e189c4e5adb52123d0a57f0f63352c99d7d3ec20a8a821f27f930709 filter=lfs diff=lfs merge=lfs -text
|
| 67 |
+
image/blobs/sha256/82bcf016508f0384e4aca1660cda09ea6f63d7dec3223a9d02f08cf6be718663 filter=lfs diff=lfs merge=lfs -text
|
| 68 |
+
image/blobs/sha256/9b8844999af8c55ff4923e14d78b61b81f93b1387154462deb0e7f2290abbcfd filter=lfs diff=lfs merge=lfs -text
|
| 69 |
+
image/blobs/sha256/be7e87294626da08f944fc039cac2c322108d9c2dd2335ebf5f63f27c2c6335a filter=lfs diff=lfs merge=lfs -text
|
| 70 |
+
image/blobs/sha256/c2ce3023f2ddc0bc464cfbabb86f6fa277861c23662a8b447909ff5057e1caa2 filter=lfs diff=lfs merge=lfs -text
|
| 71 |
+
image/blobs/sha256/cab9ef9b426e995ed54516bbbff3ed9429b565b12b361fe88f327d1a6aa0a6c9 filter=lfs diff=lfs merge=lfs -text
|
| 72 |
+
image/blobs/sha256/ff074d9d763b545567421043a31b67b5b0a9c8eb4a968e40ecc74476d54d8419 filter=lfs diff=lfs merge=lfs -text
|
OPT_BASELINE.md
CHANGED
|
@@ -20,6 +20,32 @@ power had a median of 76 W and a max of 86 W, at 63-68 °C. The profiler CSV hea
|
|
| 20 |
shared (AMD EPYC-Rome, 8 vCPUs, up to 8 other agents, load average 8-12 during these runs), so host-side stages are
|
| 21 |
quoted as p50 / p99 (min).
|
| 22 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 23 |
## Summary
|
| 24 |
|
| 25 |
| | baseline |
|
|
|
|
| 20 |
shared (AMD EPYC-Rome, 8 vCPUs, up to 8 other agents, load average 8-12 during these runs), so host-side stages are
|
| 21 |
quoted as p50 / p99 (min).
|
| 22 |
|
| 23 |
+
## After optimization (the current release, 2026-10-11)
|
| 24 |
+
|
| 25 |
+
Everything below this section is the baseline of the first release and stays as measured. The optimized release
|
| 26 |
+
(rounds 1-5 of `OPT_REPORT.md`, verified in `VERIFICATION_OPT_2026-10-11.md`, re-measured for the republish on
|
| 27 |
+
`bc966b4` with ttaw 0.23.2; same configuration: ETH dispatch, 1 CQ, 12×10, AICLK 1350 MHz, the stage bench above):
|
| 28 |
+
|
| 29 |
+
| | baseline (`5541833`, 2026-10-08) | current (2026-10-11) |
|
| 30 |
+
|---|---|---|
|
| 31 |
+
| device latency per plan (one trace replay + sync), shipped sample | 102.13 ms p50 | **20.05 ms** p50 (p99 20.08) |
|
| 32 |
+
| back-to-back replays, shipped sample (88 neighbours) | 102.04 ms = 9.80 plans/s | **20.00 ms = 50.0 plans/s** (5.1×) |
|
| 33 |
+
| back-to-back replays by agent bucket (r32 / r64 / r96 / r128 / r192 / full) | 102.04 ms for every scene | 17.44 / 19.25 / 20.00 / 21.11 / 22.30 / 26.65 ms (the last three from `OPT_REPORT.md` round 5) |
|
| 34 |
+
| `model(inputs=arrays)` end to end, shipped sample | 117.90 ms p50 (p99 134.56) | **26.73 ms** p50 (p99 27.48) |
|
| 35 |
+
| `model(inputs=<.npz path>)` | 124.76 ms | 32.27 ms |
|
| 36 |
+
| host pre · pack · host tensors · H2D · D2H · host post | 4.75 · 0.45 · 2.41 · 1.14 · 0.58 · 4.32 ms | 2.45 · 0.52 · 1.57 · 0.76 · 0.22 · 1.39 ms |
|
| 37 |
+
| programs per plan (traced replay profile) | 6,282 | 1,418 at r96 (kernel 19.64 ms, op-to-op gaps 0.89 ms, span 20.54 ms) |
|
| 38 |
+
| traces / trace buffers | 1 / 74.6 MB | 6 (full capacity + 5 agent buckets) / 60.3 MB |
|
| 39 |
+
| `from_pretrained` load, empty / warm JIT cache | 315 s / 8.6 s | 352 s / 12.9 s |
|
| 40 |
+
| served `/predict` `timing_ms.total` (uvicorn on the host) | 123.1 ms median | 37.2 ms median |
|
| 41 |
+
| accuracy: 44 device tests; 99 scenes worst ego max / mean (gates 1.0 / 0.3 m), turn, neighbours (gate 1.5 m) | 44 / 44; 0.313 / 0.143 m, 99 / 99, ≤ 0.086 m | 44 / 44; 0.346 / 0.154 m, 99 / 99, ≤ 0.094 m (two gated precision changes: `OPT_REPORT.md` r2.3, r3.2) |
|
| 42 |
+
|
| 43 |
+
Where the 82 ms went (same-session back-to-back A/B deltas of `OPT_REPORT.md`, largest steps): the encoder's batched channel matmuls on 4-8 cores (round 1,
|
| 44 |
+
−23.4 ms), the precision cost recovered with custom `generic_op` kernels (fused fp32 LayerNorm −15.5 ms, K-concatenated
|
| 45 |
+
split matmul −7.0 ms, fused fp32 attention −5.4 ms, softmax and glue fusions; rounds 2-3), L1 residency (round 4,
|
| 46 |
+
−3.5 ms) and exact agent-bucket compaction (round 5, −6.6 ms on this scene). Logs: `logs/diffusion-planner/republish/`
|
| 47 |
+
(`bench_recheck.json`, `device_suite.log`, `load_*.json`, `served_bench.json`).
|
| 48 |
+
|
| 49 |
## Summary
|
| 50 |
|
| 51 |
| | baseline |
|
OPT_REPORT.md
CHANGED
|
@@ -1,14 +1,71 @@
|
|
| 1 |
# diffusion-planner-p150 optimization report (p150, ETH dispatch, 12×10 grid)
|
| 2 |
|
| 3 |
-
**Status:
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
|
| 13 |
All numbers: `code/scripts/bench.py` (median of 100 warm iterations, batch 1,
|
| 14 |
`code/tt_diffusion_planner/samples/kashiwanoha_dense.npz`), ETH dispatch, 1 CQ, 12×10 grid, the pinned numerics
|
|
@@ -18,19 +75,19 @@ device profiler (one traced replay between signposts). Accuracy gates: `OPT_BASE
|
|
| 18 |
|
| 19 |
## Summary
|
| 20 |
|
| 21 |
-
| | baseline `5541833` (2026-10-08) | **
|
| 22 |
-
|---|---|---|
|
| 23 |
-
| device trace, one blocking plan | 102.13 ms (p99 104.81) | same |
|
| 24 |
-
| back-to-back traces | 102.04 ms (9.80 plans/s) | same |
|
| 25 |
-
| e2e `model(inputs=arrays)` p50 / p99 | 117.90 / 134.56 ms | same |
|
| 26 |
-
| e2e `model(inputs=<.npz path>)` p50 | 124.76 ms | same |
|
| 27 |
-
| host pre-processing / pack / host tensors / H2D / D2H / post-processing | 4.76 / 0.45 / 2.41 / 1.14 / 0.58 / 4.32 ms |
|
| 28 |
-
| device programs per plan (unique programs) | 6,282 (309) | same |
|
| 29 |
-
| kernel sum / op-to-op gaps / span (profile) | 99.27 / 4.39 / 103.66 ms | same |
|
| 30 |
-
| encoder / 11 decoder evaluations (+ solver updates) | 51.6 / 51.4 ms (4.67 ms per evaluation) | same |
|
| 31 |
-
| accuracy gates (PCC / agreement vs the fp32 CPU reference) | 44 / 44 device tests: module PCC ≥ 0.999952, encoding 0.999978, decoder evaluation ≥ 0.9999995; 99 scenes: ego max 0.313 m / mean 0.143 m (gates 1.0 / 0.3 m), turn command 99 / 99, neighbours ≤ 0.086 m (gate 1.5 m) |
|
| 32 |
-
| served `/predict` `timing_ms.total`, median of 50 (uvicorn on the host, the shipped sample) | 123.1 ms, measured on `c0d84f9` (same device code; the release verification) | same |
|
| 33 |
-
| `from_pretrained` load: empty JIT cache / warm cache | 315 s / 8.6 s (`c0d84f9`) | same |
|
| 34 |
|
| 35 |
The host is shared with other agents' jobs, so the host-side rows move with its load (several ms); the device rows
|
| 36 |
repeat to ±0.05 ms. AICLK 1350 MHz (min 1343) during the stage bench. The release re-check on `c0d84f9`
|
|
@@ -41,6 +98,35 @@ repeat to ±0.05 ms. AICLK 1350 MHz (min 1343) during the stage bench. The relea
|
|
| 41 |
| step | commit | what | trace / b2b ms | e2e ms | accuracy | revert switch |
|
| 42 |
|---|---|---|---|---|---|---|
|
| 43 |
| baseline | `5541833` | first correct port: the whole plan in one trace, ETH 1CQ 12×10, the numerics defaults of PORT_LOG decision 11, host pre- and post-processing | 102.13 / 102.04 | 117.90 | all gates (above) | – |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 44 |
|
| 45 |
Later commits up to the first release changed no model numerics and no device code: the baseline documentation
|
| 46 |
(`a6e5bbf`), the ttaw 0.19.0 re-vendor with the `serve.env` pins of the numerics knobs, the quickstart picture and two
|
|
@@ -49,11 +135,33 @@ identical gate values and per-scene numbers (`VERIFICATION_2026-10-08.md`).
|
|
| 49 |
|
| 50 |
## Round 1
|
| 51 |
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
|
| 58 |
### Findings (measured)
|
| 59 |
|
|
@@ -74,19 +182,332 @@ The baseline profile (`OPT_BASELINE.md` "Device profile") is the starting point:
|
|
| 74 |
is equal (the 12th column hardly matters to this graph yet), so ETH stays the default (D14); 2 CQs cost +9.8 ms on ETH,
|
| 75 |
so 1 CQ stays pinned.
|
| 76 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 77 |
### Megakernel / fusion work
|
| 78 |
|
| 79 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 80 |
|
| 81 |
## Known hangs (all rounds)
|
| 82 |
|
| 83 |
| when (UTC) | command | cause | status |
|
| 84 |
|---|---|---|---|
|
| 85 |
-
| – | – | none: no hang, timeout, reset or FAULT marker in any device job of the port, the baseline
|
| 86 |
|
| 87 |
## Rejected / not kept
|
| 88 |
|
| 89 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
`OPT_BASELINE.md` "What the numerics defaults cost"):
|
| 91 |
|
| 92 |
- the first device round's defaults (split only for the mixer inputs and the decoder pre-projection, bf16 SDPA):
|
|
@@ -98,8 +519,131 @@ None yet. Configurations that were measured and are **not** the default, for acc
|
|
| 98 |
- the fastest graph (fused LayerNorm, no split, bf16 SDPA): 43.7 ms, fails the encoder PCC gate (`enc.ego` 0.99856 /
|
| 99 |
0.99261).
|
| 100 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
## Profile at the end (trace replay)
|
| 102 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
The baseline profile (`OPT_BASELINE.md`): 6,282 programs, kernel sum 99.27 ms, span 103.66 ms; by op code Matmul
|
| 104 |
52.3 ms (1,365 programs), BinaryNg 31.5 ms (2,948), Reduce 3.7 ms, Softmax 3.5 ms, Unary 2.8 ms, Typecast 2.4 ms. The
|
| 105 |
ten slowest programs are the neighbour mixer's channel matmuls `[1, 320, 64, 128] @ [128, 128]`, ~970 µs each on 8
|
|
@@ -107,6 +651,11 @@ cores.
|
|
| 107 |
|
| 108 |
## Remaining backlog (gains on the 102.0 ms replay unless marked e2e; estimates, not measurements)
|
| 109 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
From `OPT_BASELINE.md` "Ranked optimization opportunities". The gains are not additive: each item shrinks the base of
|
| 111 |
the next.
|
| 112 |
|
|
|
|
| 1 |
# diffusion-planner-p150 optimization report (p150, ETH dispatch, 12×10 grid)
|
| 2 |
|
| 3 |
+
**Status: optimization rounds 1-5 done and independently verified (`VERIFICATION_OPT_2026-10-11.md`, PASS); the
|
| 4 |
+
optimized release is packaged for re-publishing (2026-10-11, section "Republish (2026-10-11)" below). Round 6
|
| 5 |
+
(map-entity compaction, the cross attention, the MK-E / MK-D megakernels) is in progress on a branch and is NOT in
|
| 6 |
+
this release.**
|
| 7 |
+
Device time per plan, back to back, on the bench scene (`kashiwanoha_dense`, 88 neighbours: the r96 bucket):
|
| 8 |
+
**102.03 (published baseline) -> 70.12 (round 1) -> 40.89 (round 2) -> 30.16 (round 3) -> 26.62 (round 4) -> 20.00 ms
|
| 9 |
+
(round 5)**, 50.0 plans/s; one blocking plan 20.06 ms; e2e `model()` p50 **26.84 ms** (113.1 at the baseline, 81.2
|
| 10 |
+
after round 1, 52.0 after round 2, 40.8 after round 3, 37.4 after round 4). Since round 5 the device time depends on
|
| 11 |
+
the scene (the agent bucket): 17.44 ms (r32, ≤ 31 neighbours: 53 of the 99 gate scenes; 44 in r64, 2 in r96) to 26.65 ms (full capacity,
|
| 12 |
+
> 191 neighbours); over the 99 gate scenes the served device time (`timing_ms.device`) is p50 20.83 ms (29.61 after
|
| 13 |
+
round 4). Accuracy: 44 / 44 device tests with the gates read-only in every step; every round-5 step is bit-identical
|
| 14 |
+
(the 99 per-scene numbers are those of round 3).
|
| 15 |
+
|
| 16 |
+
Round 5 added four kept commits, all bit-identical on the rows that are read back, each with its own knob and a
|
| 17 |
+
same-session alternating A/B (3 pairs, 3 scenes):
|
| 18 |
+
|
| 19 |
+
- the decoder on the first R rows of an agent bucket (`COMPACT=1`, buckets 32 / 64 / 96 / 128 / 192, one trace
|
| 20 |
+
each, picked per request): 26.62 -> 23.25 ms on the bench scene (r96), 21.85 ms at r32;
|
| 21 |
+
- the encoder's neighbour trunk and head on the same R neighbours, the other neighbour tokens zero (`COMPACT=2`):
|
| 22 |
+
23.25 -> 19.99 ms (r96), 17.42 ms at r32;
|
| 23 |
+
- the host pre- and post-processing vectorised, bit-exact (`HOST_FAST`): e2e 30.15 -> 27.31 ms (host pre 3.77 ->
|
| 24 |
+
2.49 ms, post 3.02 -> 1.44 ms);
|
| 25 |
+
- `cs` uploaded as one tile column and the all-zero inputs not re-uploaded (`INPUT_TRIM`): e2e 27.22 -> 26.84 ms,
|
| 26 |
+
device +0.014 ms (one concat program).
|
| 27 |
+
|
| 28 |
+
Round 4 added four kept commits, **all bit-identical** to their predecessor (same raw readback, same 99 per-scene
|
| 29 |
+
numbers), each with its own `DIFFUSION_PLANNER_*` knob and a same-session alternating A/B:
|
| 30 |
+
|
| 31 |
+
- the K = 1024 split operands (mlp fc2, final p4) written L1 block-sharded by the operand build and read in place
|
| 32 |
+
by a sharded-in0 2-D multicast matmul (`KCAT_L1`): −1.63 ms;
|
| 33 |
+
- the fused attention's inputs L1-resident: decoder qkv / cross q outputs, the hoisted cross K / V heads, the
|
| 34 |
+
masks (`ATTN_L1`): −0.41 ms, then −0.39 ms more by keeping the fusion q / kv linears in DRAM (`ATTN_L1=2`: the
|
| 35 |
+
stock linear with an L1 output ran on 6 cores, found in the round-4 profile);
|
| 36 |
+
- the decoder blocks' intermediates in L1 (split-LN outputs and stream, attention outputs, linear outputs:
|
| 37 |
+
`DEC_L1=1`): −1.01 ms; plus the pre-projection and final layer (`DEC_L1=2`): −0.09 ms.
|
| 38 |
+
|
| 39 |
+
Measured and not kept: the operand build's GELU once per tile (`KCAT_ACT_ONCE`, bit-identical, +0.06 ms), the mixer
|
| 40 |
+
intermediates in L1 (`ENC_L1`, +1.67 ms and not bit-identical) and the fusion intermediates in L1 (`FUS_L1`, +0.39
|
| 41 |
+
ms and not bit-identical): stock ops pick other programs for L1 outputs.
|
| 42 |
+
|
| 43 |
+
Round 3 added five kept commits (each with its own `DIFFUSION_PLANNER_*` knob, default on, and a same-session
|
| 44 |
+
alternating A/B); four are **bit-identical** to their predecessor (same raw readback, same 99 per-scene numbers):
|
| 45 |
+
|
| 46 |
+
- the attention score scale + mask + softmax as one `generic_op` (`ATTN_SMSM`): −1.56 ms;
|
| 47 |
+
- the encoder's split linears (pre-projections, ego / neighbour island) in the K-concatenated form (`ENC_KCAT`):
|
| 48 |
+
−1.70 ms, **the round's one precision change**: 99-scene worst ego max 0.327 -> 0.346 m, mean 0.149 -> 0.154 m,
|
| 49 |
+
neighbours 0.104 -> 0.094 m (gates 1.0 / 0.3 / 1.5 m; 47 scenes worse, 45 better);
|
| 50 |
+
- the whole fp32 matmul attention (head split, Q·Kᵀ, scale / mask / softmax, P·V, head merge) as one `generic_op`
|
| 51 |
+
with the scores in L1 (`ATTN_FUSED`): −5.43 ms;
|
| 52 |
+
- the split operand of the next K-concatenated linear written by its producer (split LN, fused attention:
|
| 53 |
+
`KCAT_EMIT`): −0.95 ms;
|
| 54 |
+
- the two transposes around each mixer token-mixing MLP inside the fused LN programs (`LN_TR`): −1.08 ms.
|
| 55 |
+
|
| 56 |
+
Measured and not kept: the mixer fc1 GELU in the matmul epilogue (`LIN_ACT`, −0.13 ms but 99-scene ego mean 0.149 ->
|
| 57 |
+
0.174 m) and the softmax scale by an immediate (`ATTN_SMSM=2`, bit-identical, −0.03 ms).
|
| 58 |
+
|
| 59 |
+
Round 2 (L2 `generic_op` kernels, six of seven steps bit-identical): fused fp32 LayerNorm (`LN_KERNEL=2`, −15.54
|
| 60 |
+
ms), residual adds in the LN (`LN_RESID`, −0.21), K-concatenated decoder split matmuls (`SPLIT_KCAT=2`, −6.98, the
|
| 61 |
+
precision change of round 2: ego max 0.313 -> 0.327 m), score scale + mask (`ATTN_SMASK`, −2.83), GELU in the
|
| 62 |
+
operand build (`KCAT_ACT`, −0.60), split-row decoder LN (`LN_SPLIT`, −0.83), SFPU statistics broadcast
|
| 63 |
+
(`LN_SFPU_BCAST`, −2.24).
|
| 64 |
+
|
| 65 |
+
Round 1 (stock-op program configs, all bit-identical): the encoder's batched linears as 2-D matmuls (`ENC_CH2D`,
|
| 66 |
+
−23.37 ms), P·V on one output tile per core (`ATTN_FAST=1`, −3.35 ms), explicit configs for the decoder's 352-row
|
| 67 |
+
matmuls (`DEC_MMCFG`, −5.18 ms). Two faster attention variants of round 1 stay rejected (not bit-identical, raise
|
| 68 |
+
the 99-scene error).
|
| 69 |
|
| 70 |
All numbers: `code/scripts/bench.py` (median of 100 warm iterations, batch 1,
|
| 71 |
`code/tt_diffusion_planner/samples/kashiwanoha_dense.npz`), ETH dispatch, 1 CQ, 12×10 grid, the pinned numerics
|
|
|
|
| 75 |
|
| 76 |
## Summary
|
| 77 |
|
| 78 |
+
| | baseline `5541833` (2026-10-08) | after round 1 (`dd6f249`, 2026-10-10) | after round 2 (`c42c540`, 2026-10-10) | after round 3 (`0b750e3`, 2026-10-10) | after round 4 (`93428bb`, 2026-10-10) | **after round 5 (`e933166`, 2026-10-10; bench scene = bucket r96)** |
|
| 79 |
+
|---|---|---|---|---|---|---|
|
| 80 |
+
| device trace, one blocking plan | 102.13 ms (p99 104.81) | **70.18 ms** (same-session baseline 102.09) | **40.96 ms** (p99 41.01) | **30.22 ms** (p99 30.42) | **26.68 ms** (p99 26.71) | **20.06 ms** (r32: 17.50, r64: 19.31; full capacity 26.7) |
|
| 81 |
+
| back-to-back traces | 102.04 ms (9.80 plans/s) | **70.12 ms (14.26 plans/s)** (same-session baseline 102.03) | **40.90 ms (24.45 plans/s)** | **30.16 ms (33.16 plans/s)** | **26.62 ms (37.56 plans/s)** | **20.00 ms (50.0 plans/s)**; buckets r32 / r64 / r96 / r128 / r192 / full: 17.44 / 19.25 / 20.00 / 21.11 / 22.30 / 26.65 ms |
|
| 82 |
+
| e2e `model(inputs=arrays)` p50 / p99 | 117.90 / 134.56 ms | **81.2 ms** p50 (same-session baseline 113.1; host load lower than on 2026-10-08) | **51.98 / 56.39 ms** | **40.84 / 41.52 ms** | **37.36 / 38.03 ms** | **26.84 ms** p50 (straight_road r32: 23.36, scene-0103_kf14 r64: 25.72) |
|
| 83 |
+
| e2e `model(inputs=<.npz path>)` p50 | 124.76 ms | 86.7 ms (same-session baseline 118.7) | 57.9 ms | 46.6 ms | 43.3 ms | 32.5 ms |
|
| 84 |
+
| host pre-processing / pack / host tensors / H2D / D2H / post-processing | 4.76 / 0.45 / 2.41 / 1.14 / 0.58 / 4.32 ms | 3.84 / 0.38 / 1.98 / 0.92 / 0.30 / 3.50 ms (unchanged code; quieter host) | 3.91 / 0.39 / 1.99 / 0.94 / 0.35 / 3.54 ms (unchanged code) | 3.85 / 0.39 / 1.96 / 0.92 / 0.33 / 3.49 ms (unchanged code) | 3.88 / 0.39 / 1.94 / 0.92 / 0.33 / 3.47 ms (unchanged code) | 2.49 / 0.52 / 1.58 / 0.76 / 0.25 / 1.44 ms (`HOST_FAST`, `INPUT_TRIM`) |
|
| 85 |
+
| device programs per plan (unique programs) | 6,282 (309) | same | 2,155 (293 program-cache entries); trace buffer 16.8 MB (44.6 before) | 1,405 (195 program-cache entries); trace buffer 12.3 MB | 1,411 (199 program-cache entries); trace buffer 12.3 MB | 1,418 at r96 (480 program-cache entries for the 6 traces); trace buffers 60.3 MB (6 traces) |
|
| 86 |
+
| kernel sum / op-to-op gaps / span (profile) | 99.27 / 4.39 / 103.66 ms | same | 39.99 / 1.47 / 41.46 ms | 29.71 / 0.88 / 30.59 ms | 26.00 / 0.89 / 26.89 ms | 19.64 / 0.89 / 20.54 ms (r96) |
|
| 87 |
+
| encoder / 11 decoder evaluations (+ solver updates) | 51.6 / 51.4 ms (4.67 ms per evaluation) | same | 15.7 / 23.9 ms (2.17 ms per evaluation; kernel sums of the profile) | 12.6 / 17.1 ms (decoder 1.55 ms per evaluation; kernel sums of the profile) | 12.4 / 13.4 ms (decoder 1.22 ms per evaluation; kernel sums of the profile) | 9.4 / 10.0 ms at r96 (decoder blocks 9.49 ms: 0.86 ms per evaluation; kernel sums of the profile) |
|
| 88 |
+
| accuracy gates (PCC / agreement vs the fp32 CPU reference) | 44 / 44 device tests: module PCC ≥ 0.999952, encoding 0.999978, decoder evaluation ≥ 0.9999995; 99 scenes: ego max 0.313 m / mean 0.143 m (gates 1.0 / 0.3 m), turn command 99 / 99, neighbours ≤ 0.086 m (gate 1.5 m) | identical (bit-identical readback; 44 / 44, all 99 per-scene numbers equal) | 44 / 44; 99 scenes: ego max 0.327 m / mean 0.149 m, turn 99 / 99, neighbours ≤ 0.104 m (the r2.3 precision change; every other step bit-identical) | 44 / 44; 99 scenes: ego max 0.346 m / mean 0.154 m, turn 99 / 99, neighbours ≤ 0.094 m (the r3.2 precision change; every other step bit-identical) | 44 / 44; 99 scenes identical to round 3 (every round-4 step bit-identical) | 44 / 44; 99 scenes identical to round 3 (every round-5 step bit-identical on the read rows) |
|
| 89 |
+
| served `/predict` `timing_ms.total`, median of 50 (uvicorn on the host, the shipped sample) | 123.1 ms, measured on `c0d84f9` (same device code; the release verification) | same | not re-measured (no release in this round) | not re-measured | not re-measured | **37.2 ms** (min 36.7; decode 8.6, `model()` 26.8; republish re-measure on `bc966b4`, 2026-10-11) |
|
| 90 |
+
| `from_pretrained` load: empty JIT cache / warm cache | 315 s / 8.6 s (`c0d84f9`) | same | – / 7.9 s (warm, bench) | – / 7.1 s (warm, bench) | – / 7.8 s (warm, bench) | 352 s / 12.9 s (6 traces; republish re-measure, 2026-10-11; 11.4 s in the bench process) |
|
| 91 |
|
| 92 |
The host is shared with other agents' jobs, so the host-side rows move with its load (several ms); the device rows
|
| 93 |
repeat to ±0.05 ms. AICLK 1350 MHz (min 1343) during the stage bench. The release re-check on `c0d84f9`
|
|
|
|
| 98 |
| step | commit | what | trace / b2b ms | e2e ms | accuracy | revert switch |
|
| 99 |
|---|---|---|---|---|---|---|
|
| 100 |
| baseline | `5541833` | first correct port: the whole plan in one trace, ETH 1CQ 12×10, the numerics defaults of PORT_LOG decision 11, host pre- and post-processing | 102.13 / 102.04 | 117.90 | all gates (above) | – |
|
| 101 |
+
| r1.0 | `0f557b9` | re-vendor ttaw 0.20.0 -> 0.23.0 (no planner module changed) | – | – | 44 / 44, 99 scenes identical | – |
|
| 102 |
+
| r1.1 | `9a885b3` | encoder linears on batched `[1, E, T, C]` activations as one 2-D matmul: free `[1, 1, E·64, 128]` view for the MixerBlock channel MLPs (auto 2-D multicast), a 1-D in1-multicast config with `fuse_batch=True` for the pre-projections (T = 6 / 20 / 40) | 102.09 -> 78.72 / 102.03 -> **78.66** (−23.37) | 113.1 -> 89.9 | bit-identical readback; 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_ENC_CH2D=0` |
|
| 103 |
+
| r1.2 | `ef6a763` | fp32 matmul attention (`tt/attention.py`): P·V on `MatmulMultiCoreReuseProgramConfig` with one output tile per core and the whole K in one block (88 cores instead of 11; sweep: 73 -> 21 µs self, 71 -> 32 µs cross, 81 -> 42 µs fusion) | 78.72 -> 75.37 / 78.66 -> **75.31** (−3.35) | 89.8 -> 86.5 | bit-identical; 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_ATTN_FAST=0` |
|
| 104 |
+
| r1.3 | `dd6f249` | decoder 352-row matmuls (every split pass): explicit `MatmulMultiCoreReuseMultiCastProgramConfig` per (K, N, pass), the fastest bit-identical candidate of a device sweep (`tt/layers.py` `DEC_MM_CONFIGS`; e.g. mlp fc2 30 -> 12 µs) | 75.37 -> 70.18 / 75.30 -> **70.12** (−5.18) | 86.5 -> 81.2 | bit-identical; 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_DEC_MMCFG=0` |
|
| 105 |
+
| r2.1 | `2fc9af8` | fused fp32 LayerNorm (`LN_KERNEL=2`): the 9-program `layer_norm_fp32` decomposition as one `generic_op` (`tt/ln_kernel.py`, `kernels/ln32_*.cpp`), the same SFPU LLK calls in the same order (accurate `ttnn.mean` fold + `sfpu_reduce`, binary_ng fp32 SFPU ops, `rsqrt` Default) | 70.18 -> 54.64 / 70.12 -> **54.58** (−15.54) | 81.6 -> 65.6 | bit-identical readback; 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_LN_KERNEL=0` |
|
| 106 |
+
| r2.2 | `6ef00e6` | residual adds fused into the LN program (`LN_RESID`): `h = x + r (* gate)` then `LN(h)`, writing both (mixers: x + token-mix, x + channel-MLP; decoder: the four stream adds of a block and the final one) | 54.64 -> 54.43 / 54.58 -> **54.37** (−0.21) | 65.5 -> 65.2 | bit-identical; 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_LN_RESID=0` |
|
| 107 |
+
| r2.3 | `d7f84e1` | decoder split matmuls as ONE K-concatenated fp32 matmul (`SPLIT_KCAT=2`): `[x_hi \| x_hi \| x_lo \| 1 \| 0..] @ [w_hi; w_lo; w_hi; b_hi, b_lo; 0..]`, operand built by one `generic_op` (`kernels/kcat_*.cpp`), K' padded for K = 1024, per-shape 2-D multicast configs from a sweep | 54.44 -> 47.45 / 54.37 -> **47.39** (−6.98) | 65.4 -> 58.3 | **not bit-identical** (one fp32 accumulation, exact bias): final_x0 rel 1.1e-3; 44 / 44 green; 99 scenes ego max 0.313 -> 0.327 m, mean 0.143 -> 0.149 m, neighbours 0.086 -> 0.104 m (57 scenes worse, 41 better) | `DIFFUSION_PLANNER_SPLIT_KCAT=0` |
|
| 108 |
+
| r2.4 | `2da6100` | attention score scale + mask in one pass (`ATTN_SMASK`, `kernels/smask_*.cpp`); reproduces the TF32 truncation of the stock fp32 + bf16 add (finding below) | 47.45 -> 44.63 / 47.39 -> **44.56** (−2.83) | 58.3 -> 55.6 | bit-identical to r2.3 (readback, 99 / 99 scenes); 44 / 44 | `DIFFUSION_PLANNER_ATTN_SMASK=0` |
|
| 109 |
+
| r2.5 | `4699447` | the GELU between two K-concatenated linears applied inside the next operand build (`KCAT_ACT`; 77 unary programs per plan) | 44.62 -> 44.02 / 44.56 -> **43.96** (−0.60) | 55.5 -> 54.9 | bit-identical to r2.3; 44 / 44 | `DIFFUSION_PLANNER_KCAT_ACT=0` |
|
| 110 |
+
| r2.6 | `98ca6c7` | decoder LayerNorms with each tile row over 8 cores (`LN_SPLIT`, `kernels/ln32s_*.cpp`: gather to a root core, ordered fold, broadcast; semaphores): 88 cores instead of 11 | 44.02 -> 43.19 / 43.96 -> **43.13** (−0.83) | 54.8 -> 54.3 | bit-identical to r2.3; 44 / 44 | `DIFFUSION_PLANNER_LN_SPLIT=0` |
|
| 111 |
+
| r2.7 | `c42c540` | LN statistics broadcast by the SFPU row reduce itself (`LN_SFPU_BCAST`, `kernels/ln32_sfpu.h`: the stock `sfpu_reduce` row-sum arithmetic with the result stored to every column) instead of a RISC-V 1024-store column fill per statistic | 43.19 -> 40.95 / 43.13 -> **40.89** (−2.24) | 53.9 -> 51.9 | bit-identical to r2.3; 44 / 44 | `DIFFUSION_PLANNER_LN_SFPU_BCAST=0` |
|
| 112 |
+
| r3.1 | `c03adc3` | attention score scale + mask + softmax as one `generic_op` (`ATTN_SMSM=1`, `kernels/smsm_*.cpp`): the smask SFPU sequence into an L1 row, then the stock `ttnn.softmax(numeric_stable=True)` compute with the same `kernel_lib` calls | 40.96 -> 39.39 / 40.89 -> **39.32** (−1.56) | 51.9 -> 50.3 | bit-identical to r2.7 (readback, 99 / 99 scenes); 44 / 44 | `DIFFUSION_PLANNER_ATTN_SMSM=0` |
|
| 113 |
+
| r3.2 | `c6edfaf` | encoder split linears (`enc.pre.*`, `enc.island.*`, any K: weight blocks padded to whole tiles) as one K-concatenated fp32 matmul (`ENC_KCAT`), 1-D `fuse_batch` config for the `[1, E, T, K']` operands | 39.38 -> 37.68 / 39.32 -> **37.62** (−1.70) | 50.3 -> 48.4 | **not bit-identical** (one fp32 accumulation, exact bias; final_x0 rel 2.9e-3); 44 / 44 green; 99 scenes ego max 0.327 -> 0.346 m (scene-0103_kf14), mean 0.149 -> 0.154 m, neighbours 0.104 -> 0.094 m (47 scenes worse, 45 better, mean change −0.2 mm) | `DIFFUSION_PLANNER_ENC_KCAT=0` |
|
| 114 |
+
| r3.3 | `8d402d1` | (not kept, knob default 0) mixer fc1 GELU in the matmul epilogue (`LIN_ACT`): an explicit copy of the stock auto config (bit-identical without the activation) + `fused_activation` | 39.38 -> 39.25 / 39.32 -> 39.19 (−0.13) | 50.3 -> 50.1 | not bit-identical; 44 / 44 green; ego max 0.327 -> 0.392 m, mean 0.149 -> 0.174 m (54 worse, 38 better) | default 0 |
|
| 115 |
+
| r3.4 | `4acde70` | the whole fp32 matmul attention as one `generic_op` (`ATTN_FUSED`, `kernels/fattn_*.cpp`): per (head, query tile row) Q·Kᵀ into DEST (`matmul_block`, in1 transposed), the smsm phases on the row in L1, P·V accumulated in DEST in key order; Q / K / V read in place (no head split), output written as the merged heads (no head merge) | 37.68 -> 32.25 / 37.62 -> **32.19** (−5.43) | 48.9 -> 43.2 | bit-identical to r3.2 (readback, 99 / 99 scenes); 44 / 44 | `DIFFUSION_PLANNER_ATTN_FUSED=0` |
|
| 116 |
+
| r3.5 | `ad03767` | the decoder's split-row LNs and the fused attention write the split operand `[x_hi \| x_hi \| x_lo \| 1]` of the next K-concatenated linear (`KCAT_EMIT`; `layers.KcatOperand`): 198 `kcat_operand` programs per plan less | 32.25 -> 31.30 / 32.19 -> **31.24** (−0.95) | 43.1 -> 42.4 | bit-identical to r3.4; 44 / 44 | `DIFFUSION_PLANNER_KCAT_EMIT=0` |
|
| 117 |
+
| r3.6 | `0b750e3` | the two transposes around each mixer token-mixing MLP inside the fused LN programs (`LN_TR`): n1 writes its output per-entity transposed, n2 reads the token-mixing output transposed (`transpose_tile`, the stock transpose LLK with fp32 unpack-to-dest: exact); 48 transpose programs less | 31.30 -> 30.22 / 31.24 -> **30.16** (−1.08) | 42.1 -> 41.0 | bit-identical to r3.5; 44 / 44 | `DIFFUSION_PLANNER_LN_TR=0` |
|
| 118 |
+
| r4.1 | `7e662d3` | the K = 1024 split operands X' `[352, 3328]` (mlp fc2) / `[352, 3168]` (final p4, K' padded to 11 slices) written **L1 block-sharded** by the operand build (its writer's TensorAccessor resolves the shards) and read in place by a 2-D multicast matmul with in0 sharded (`KCAT_L1`, `layers.KCAT_L1_CONFIGS`): no 4.7 MB DRAM round trip (kernel check: operand + matmul 62.6 -> 41.3 µs, 62.3 -> 41.8 µs) | 30.22 -> 28.59 / 30.16 -> **28.53** (−1.63) | 40.8 -> 39.3 | bit-identical to r3.6 (readback, 99 / 99 scenes); 44 / 44 | `DIFFUSION_PLANNER_KCAT_L1=0` |
|
| 119 |
+
| r4.2 | `310bd38` | (not kept, knob default 0) the deferred GELU of the operand build on one DEST copy + `copy_dest_values` (`KCAT_ACT_ONCE`) | 28.62 -> 28.66 / 28.53 -> 28.59 (+0.06) | (host-noisy) | bit-identical | default 0 |
|
| 120 |
+
| r4.3 | `f6cc333` | the fused attention's inputs L1-resident (`ATTN_L1=1`): decoder qkv / cross q outputs and fusion q / kv outputs in L1, the hoisted cross K / V heads moved to L1 once per plan, the self / cross / fusion masks in L1 (kernel check: self 51.0 -> 35.1 µs, cross 61.0 -> 54.2, fusion 106.2 -> 88.4) | 28.59 -> 28.18 / 28.53 -> **28.12** (−0.41) | 39.3 -> 38.9 | bit-identical to r4.1; 44 / 44 | `DIFFUSION_PLANNER_ATTN_L1=0` |
|
| 121 |
+
| r4.4 | `e5fb151` | the decoder blocks' intermediates in L1 interleaved (`DEC_L1=1`): split-LN outputs and stream `h`, fused attention outputs, out / mlp / cross-out linear outputs (kernel check: split LN 24.1 -> 20.6 µs, K = 256 matmuls −2 to −4 µs, attention −1.7 µs) | 28.18 -> 27.16 / 28.12 -> **27.11** (−1.01) | 38.8 -> 37.9 | bit-identical to r4.3; 44 / 44 | `DIFFUSION_PLANNER_DEC_L1=0` |
|
| 122 |
+
| r4.5 | `60eaef1` | (not kept, default 0) the mixer blocks' intermediates in L1 (`ENC_L1`) | 27.16 -> 28.84 / 27.11 -> 28.77 (+1.67) | 37.7 -> 39.5 | **not** bit-identical (final_x0 rel 8e-3; ego max 0.379 m): the stock matmul / unary programs change with an L1 output; 44 / 44 | default 0 |
|
| 123 |
+
| r4.6 | `dab96fb` | `DEC_L1=2`: the decoder pre-projection and final layer intermediates in L1 too | 27.17 -> 27.07 / 27.11 -> **27.02** (−0.09) | 37.9 -> 37.9 | bit-identical to r4.4; 44 / 44 | `DIFFUSION_PLANNER_DEC_L1=1` |
|
| 124 |
+
| r4.7 | `c3b6a5a` | (not kept, default 0) the fusion blocks' intermediates in L1 (`FUS_L1`) | 27.07 -> 27.47 / 27.02 -> 27.40 (+0.39) | 37.8 -> 38.2 | **not** bit-identical (ego max 0.364 m; `ttnn.layer_norm` with an L1 output); 44 / 44 | default 0 |
|
| 125 |
+
| r4.8 | `93428bb` | `ATTN_L1=2`: the fusion q / kv linears back to DRAM outputs (the stock linear with an L1 output ran on 6 cores, 54 µs, with a separate bias add: the fusion stage 1.41 -> 1.74 ms in the first round-4 profile); masks and decoder inputs stay in L1 | 27.08 -> 26.69 / 27.02 -> **26.62** (−0.39) | 37.8 -> 37.7 | bit-identical to r4.6; 44 / 44 | `DIFFUSION_PLANNER_ATTN_L1=1` |
|
| 126 |
+
| r5.1 | `9fe29a8` | exact decoder compaction (`COMPACT=1`): one trace per agent bucket (`AGENT_BUCKETS` 32, 64, 96, 128, 192) besides the full 352-row plan; the decoder runs on the first R rows (the ego, every valid self-attention key, every emitted neighbour), `y0` / `cs` / the key row sliced in the trace; per-R agent rows and cross mask; the matmul configs keep their K blocking (`per_core_M` fitted to the grid); the bucket picked per request (`TtDiffusionPlanner.variant_for`) | 26.68 -> 23.31 / 26.62 -> **23.25** (−3.37; r32 21.85, r64 22.91) | 38.3 -> 33.8 | bit-identical on the needed rows vs the full plan (15 scenes, `compact_check.py`); 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_COMPACT=0` |
|
| 127 |
+
| r5.2 | `268ff32` | `COMPACT=2`: the encoder's neighbour trunk and head on the same first R neighbours, the other neighbour tokens a zero block (invalid entities: `token_valid` zeroes them anyway); the bucket also covers every valid neighbour token | 23.30 -> 20.05 / 23.25 -> **19.99** (−3.26; r32 17.42, r64 19.23) | 33.2 -> 30.1 | bit-identical on the needed rows (15 scenes); 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_COMPACT=1` |
|
| 128 |
+
| r5.3 | `3ce8ea9` | vectorised host pre- / post-processing (`HOST_FAST`): normalisation of the non-small rows only, neighbour features on the 6 kept rows, the trajectory loops as array operations in the same float64 / float32 order, denormalisation and predicted paths on the emitted rows only | device unchanged (19.99) | 30.15 -> **27.31** | bit-exact on 134 inputs (`host_equiv.py`, both directions of the knob); 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_HOST_FAST=0` |
|
| 129 |
+
| r5.4 | `e933166` | trimmed uploads (`INPUT_TRIM`): `cs` as one tile column `[352, 32]` widened in the trace by a tile-aligned concat; an all +0.0 input (`y0`, `static_x`) not re-uploaded while its buffer holds this model's last zeros | 20.05 -> 20.06 / 19.99 -> 20.00 (+0.014) | 27.22 -> **26.84** | readback bit-identical (bench dumps, 3 scenes x 6 runs); 44 / 44; 99 scenes identical | `DIFFUSION_PLANNER_INPUT_TRIM=0` |
|
| 130 |
|
| 131 |
Later commits up to the first release changed no model numerics and no device code: the baseline documentation
|
| 132 |
(`a6e5bbf`), the ttaw 0.19.0 re-vendor with the `serve.env` pins of the numerics knobs, the quickstart picture and two
|
|
|
|
| 135 |
|
| 136 |
## Round 1
|
| 137 |
|
| 138 |
+
Date 2026-10-10, plan `OPT_PLAN.md` §7 round 1 (items 1, 2a, 4a). The device was shared with a METEOR debugging job
|
| 139 |
+
that hung the chip three times during the round (05:22, 05:45 and 07:35 UTC, FAULT markers written by its own devrun
|
| 140 |
+
jobs). None of this round's jobs ran on a faulted chip: every job script checks for the marker after taking the lock,
|
| 141 |
+
and the queued jobs aborted with exit 75 and were re-queued after the orchestrator cleared the marker. No job of this
|
| 142 |
+
round hung.
|
| 143 |
+
|
| 144 |
+
Method, per step (`logs/diffusion-planner/opt_r1/`; scripts in `scripts/`):
|
| 145 |
+
|
| 146 |
+
- `ab.sh <tag> <KNOB> <off> <on>`: the device suite with the new value (`TTAW_GATES_READONLY=1`, both device test
|
| 147 |
+
files, `DIFFUSION_PLANNER_E2E_REPORT` per-scene JSON), then 3 alternating pairs of `bench.py` (`--iters 30`, b2b 3 ×
|
| 148 |
+
50) in one devrun window. Every bench run writes its raw packed readback (`--dump`), and `cmp.py` compares on / off
|
| 149 |
+
bit for bit and the per-scene report with the baseline's `e2e_device.json`. The other knobs are pinned so each A/B
|
| 150 |
+
isolates one change; jobs run from a snapshot of the code (`snap.sh`), so later edits cannot leak into a queued job.
|
| 151 |
+
- `mm_sweep.py` / `attn_sweep.py`: device micro-benches (20-call traces) of the decoder matmul shapes (three split
|
| 152 |
+
passes each) and the attention pieces under the auto config and candidate program configs, outputs compared bit for
|
| 153 |
+
bit (`mm_sweep.json`, `attn_sweep.log`).
|
| 154 |
+
|
| 155 |
+
| A/B (same session, 3 pairs unless noted, kashiwanoha_dense) | off: b2b / trace / e2e ms | on: b2b / trace / e2e ms | Δ b2b | outputs | kept |
|
| 156 |
+
|---|---|---|---:|---|---|
|
| 157 |
+
| `ENC_CH2D` 0 / 1 (ATTN_FAST=0, DEC_MMCFG=0) | 102.03 / 102.09 / 113.1 | 78.66 / 78.72 / 89.9 | −23.37 | bit-identical | yes |
|
| 158 |
+
| `ATTN_FAST` 0 / 1 (ENC_CH2D=1, DEC_MMCFG=0) | 78.66 / 78.72 / 89.8 | 75.31 / 75.37 / 86.5 | −3.35 | bit-identical | yes |
|
| 159 |
+
| `DEC_MMCFG` 0 / 1 (ENC_CH2D=1, ATTN_FAST=1) | 75.30 / 75.37 / 86.5 | 70.12 / 70.18 / 81.2 | −5.18 | bit-identical | yes |
|
| 160 |
+
| `ATTN_FAST` 0 / 3: + scale as an SFPU pre-activation of the mask add (ENC_CH2D=1, DEC_MMCFG=0; the first form of the knob, then a bool) | 78.66 / 78.72 / 89.6 | 72.81 / 72.88 / 83.8 | −5.85 (−2.5 over mode 1) | final_x0 rel 4.1e-3; 99 scenes: ego max 0.313 -> 0.385 m, mean 0.143 -> 0.173 m; 44 / 44 green | no |
|
| 161 |
+
| `ATTN_FAST` 1 / 2: + scale on Q before Q·Kᵀ (2 pairs) | 70.12 / 70.18 / 81.2 | 68.15 / 68.21 / 79.5 | −1.97 | final_x0 rel 3.0e-3; ego max 0.313 -> 0.379 m, mean 0.143 -> 0.172 m, neighbours 0.086 -> 0.098 m; 44 / 44 green | no |
|
| 162 |
+
|
| 163 |
+
AICLK 1350 MHz (min 1343) in every run. The device rows repeat to ±0.01 ms within a session; e2e moves with the host
|
| 164 |
+
load by ±0.5 ms.
|
| 165 |
|
| 166 |
### Findings (measured)
|
| 167 |
|
|
|
|
| 182 |
is equal (the 12th column hardly matters to this graph yet), so ETH stays the default (D14); 2 CQs cost +9.8 ms on ETH,
|
| 183 |
so 1 CQ stays pinned.
|
| 184 |
|
| 185 |
+
Round 1 findings:
|
| 186 |
+
|
| 187 |
+
- **The 4-D batched linear is the stock trap.** `ttnn.linear` on `[1, E, T, C] @ [C, N]` auto-configures with M per
|
| 188 |
+
batch element (`fuse_batch=False`, `matmul_program_config.cpp` `create_simple_matmul_program_config`), so the 6 ×
|
| 189 |
+
320 entity blocks ran on 8 cores. A free reshape (T a multiple of 32) or a program config with `fuse_batch=True`
|
| 190 |
+
gives the whole grid the work: same products, same fp32 DEST order, bit-identical output. 23.4 ms, the largest
|
| 191 |
+
single gain of the plan so far.
|
| 192 |
+
- **K-blocking decides bit-identity, the core grid does not.** In both sweeps every candidate with the same
|
| 193 |
+
`in0_block_w` as the auto config (the whole K in one block for P·V; 8 for the decoder linears) was bit-identical;
|
| 194 |
+
smaller `in0_block_w` changed the P·V output by up to 7e-4 (the partial sums spill between K blocks).
|
| 195 |
+
- **The 99-scene plan error reacts to any rounding change.** Both attention variants that only move one fp32 rounding
|
| 196 |
+
(where the 1/√32 is applied) raised the worst ego error by ~20 % (0.31 -> 0.38 m), with every gate still green. This
|
| 197 |
+
is the chaotic sensitivity of PORT_LOG known issues; bit-identical rewrites were preferred for that reason.
|
| 198 |
+
- **`input_tensor_a_activations` is not a free fusion for fp32.** The scale pre-activation on the mask add is not
|
| 199 |
+
bit-identical to `multiply` + `add` (rel 4e-3 on `final_x0`), although both look like an fp32 SFPU multiply then
|
| 200 |
+
the same add.
|
| 201 |
+
|
| 202 |
### Megakernel / fusion work
|
| 203 |
|
| 204 |
+
No megakernel yet; the candidates (D19) are in the backlog (item 6): the persistent DiT-evaluation kernel first. Round
|
| 205 |
+
2's kernels are its building blocks (the fp32 LN with the residual phase, the split-row LN with semaphore gather /
|
| 206 |
+
broadcast, the K-concatenated split operand, the score scale + mask). Round 3 added the first block-level fusion: the
|
| 207 |
+
whole fp32 attention (Q·Kᵀ, scale / mask / softmax, P·V, head split / merge) as one bit-identical program
|
| 208 |
+
(`ATTN_FUSED`), the attention phase of MK-D, plus producer-side operand emission (`KCAT_EMIT`) and the mixer
|
| 209 |
+
transposes folded into the LN (`LN_TR`), which are the MK-E / MK-D data-layout pieces. The megakernel attempt keeps its
|
| 210 |
+
own phase.
|
| 211 |
+
|
| 212 |
+
## Round 2
|
| 213 |
+
|
| 214 |
+
Date 2026-10-10, plan `OPT_PLAN.md` §7 round 2 (L2 `generic_op` kernels: items 3, 2, 4b, 5). The device was shared
|
| 215 |
+
with METEOR hang-debugging jobs (two FAULT markers from their jobs, 09:19 and 10:05 UTC) and other agents' optimization
|
| 216 |
+
jobs; none of this round's jobs hung, timed out or ran on a faulted chip (every job checks the marker after taking the
|
| 217 |
+
lock). Every custom kernel went through the hang protocol before it touched the plan: smallest shape under
|
| 218 |
+
`TT_METAL_WATCHER=2`, eager twice, then a traced replay, then the real shapes (`logs/diffusion-planner/opt_r2/`).
|
| 219 |
+
|
| 220 |
+
Method, per step (`logs/diffusion-planner/opt_r2/`; scripts in `scripts/`): a device micro-check of the kernel against
|
| 221 |
+
the stock programs it replaces (`ln_check.py`, `kcat_check.py`, `smask_check.py`: bit comparison, eager repeat,
|
| 222 |
+
trace replay, µs per call over a 20-call trace), then `ab.sh <tag> <KNOB> <off> <on>` as in round 1 (device suite
|
| 223 |
+
with the gates read-only, 3 alternating `bench.py` pairs with raw readback dumps, `cmp.py` bit comparison and the
|
| 224 |
+
99-scene report). Every job runs from a snapshot of the code (`snap.sh`; scripts copied into it).
|
| 225 |
+
|
| 226 |
+
| A/B (same session, 3 pairs, kashiwanoha_dense) | off: b2b / trace / e2e ms | on: b2b / trace / e2e ms | Δ b2b | outputs | kept |
|
| 227 |
+
|---|---|---|---:|---|---|
|
| 228 |
+
| `LN_KERNEL` 0 / 2 (LN_RESID=0, SPLIT_KCAT=0) | 70.12 / 70.18 / 81.6 | 54.58 / 54.64 / 65.6 | −15.54 | bit-identical | yes |
|
| 229 |
+
| `LN_RESID` 0 / 1 (LN_KERNEL=2, SPLIT_KCAT=0) | 54.58 / 54.64 / 65.5 | 54.37 / 54.43 / 65.2 | −0.21 | bit-identical | yes |
|
| 230 |
+
| `SPLIT_KCAT` 0 / 2 (LN_KERNEL=2, LN_RESID=1) | 54.37 / 54.44 / 65.4 | 47.39 / 47.45 / 58.3 | −6.98 | final_x0 rel 1.1e-3; 99 scenes ego max 0.313 -> 0.327 m, mean 0.143 -> 0.149 m, nb 0.086 -> 0.104 m; 44 / 44 green | yes |
|
| 231 |
+
| `ATTN_SMASK` 0 / 1 (KCAT_ACT=0, LN_SPLIT=0) | 47.39 / 47.45 / 58.3 | 44.56 / 44.63 / 55.6 | −2.83 | bit-identical | yes |
|
| 232 |
+
| `KCAT_ACT` 0 / 1 (ATTN_SMASK=1, LN_SPLIT=0) | 44.56 / 44.62 / 55.5 | 43.96 / 44.02 / 54.9 | −0.60 | bit-identical | yes |
|
| 233 |
+
| `LN_SPLIT` 0 / 1 (ATTN_SMASK=1, KCAT_ACT=1) | 43.96 / 44.02 / 54.8 | 43.13 / 43.19 / 54.3 | −0.83 | bit-identical | yes |
|
| 234 |
+
| `LN_SFPU_BCAST` 0 / 1 (all of the above on) | 43.13 / 43.19 / 53.9 | 40.89 / 40.95 / 51.9 | −2.24 | bit-identical | yes |
|
| 235 |
+
|
| 236 |
+
AICLK 1350 MHz (min 1343) in every run; the device rows repeat to ±0.01 ms within a session.
|
| 237 |
+
|
| 238 |
+
Kernel micro-benchmarks (µs per call, 20-call trace; stock programs -> the kernel; all bit-identical except the
|
| 239 |
+
K-concatenated matmul, which is compared with its stock-op operand build):
|
| 240 |
+
|
| 241 |
+
| kernel | shape | stock | kernel |
|
| 242 |
+
|---|---|---:|---:|
|
| 243 |
+
| fp32 LN (`LN_KERNEL`, one core per tile row) | decoder [352, 256] + affine / neighbour mixer [320·64, 128] / lane [140·64, 128] | 47.8 / 510.8 / 238.3 | 27.6 / 126.5 / 73.1 |
|
| 244 |
+
| + residual and gate (`LN_RESID`) | decoder / neighbour mixer | 61.7 / 680.5 | 38.5 / 191.9 |
|
| 245 |
+
| split-row LN (`LN_SPLIT`, 8 cores per row) | decoder, plain / + residual and gate | 47.9 / 61.7 | 26.8 / 29.4 |
|
| 246 |
+
| split linear: 3 passes (`DEC_MMCFG`) vs K-concatenated (operand kernel + one matmul) | qkv 256->768 / out 256->256 / fc1 256->1024 / fc2 1024->256 | 62.1 / 50.0 / 70.4 / 70.7 | 27.2 / 22.4 / 31.2 / 62.5 |
|
| 247 |
+
| score scale + mask (`ATTN_SMASK`) | 8 x 352 x 352 / 8 x 352 x 576 / 8 x 576 x 576 | 74.9 / 73.8 / 118.9 | 27.0 / 46.1 / 64.7 |
|
| 248 |
+
|
| 249 |
+
### Findings (measured)
|
| 250 |
+
|
| 251 |
+
- **Bit-identical fused kernels are possible for the precision decompositions** when the kernel issues the same LLK
|
| 252 |
+
calls in the same order as the stock programs and keeps every intermediate fp32 in L1 / DST: the fp32 LayerNorm
|
| 253 |
+
(accurate `ttnn.mean` = `add_binary_tile` fold + `sfpu_reduce` + `mul_unary_tile(1/W)`; binary_ng fp32 SFPU ops;
|
| 254 |
+
`rsqrt_tile<Default>`), the residual / gate adds, the GELUs and the attention scale reproduce the published numerics
|
| 255 |
+
bit for bit. That made 5 of the 6 kept steps exact rewrites (same readback, same 99 per-scene numbers).
|
| 256 |
+
- **The stock fp32 + bf16 add is not an fp32 add.** `ttnn.add(scores_fp32, mask_bf16)` (the key mask of every fp32
|
| 257 |
+
matmul attention of the plan) runs on binary_ng's FPU path; its SrcA truncates the fp32 scores to TF32 (stock ==
|
| 258 |
+
exact & 0xFFFFE000 on every finite element, `smask_diag.log`; with an fp32 mask the add is exact). The published
|
| 259 |
+
graph therefore feeds TF32-truncated scores to the softmax. `ATTN_SMASK` reproduces that (an SFPU AND on the raw
|
| 260 |
+
bits); an exact variant (mask in fp32, or no truncation) would be a precision change to gate separately.
|
| 261 |
+
- **One K-concatenated fp32 matmul replaces the three split passes** (and their typecast / subtract / add programs):
|
| 262 |
+
the products are the same (bf16-exact hi parts, TF32-truncated lo parts), only the accumulation order changes and
|
| 263 |
+
the bias becomes exact (two K rows instead of the packer epilogue, which rounds the fp32 output: TransFusion T4).
|
| 264 |
+
Per linear the rms error vs the fp64 product goes 3.2e-4 -> 3.1e-4; over the 99 scenes the worst scene moves
|
| 265 |
+
+4 % (0.313 -> 0.327 m), the others both ways (57 worse, 41 better, mean change +0.3 mm on the ego max). Kept: the
|
| 266 |
+
gates have 3x (max) / 2x (mean) margin, and it is the largest non-exact gain of the round (−7.0 ms).
|
| 267 |
+
- **The K block must divide K'.** 3·Kt + 1 is prime for K = 1024 (97 tiles): the auto config then uses a 1-tile K
|
| 268 |
+
block (90 µs); padding K' to 104 = 8 x 13 with zero tiles and a 13-tile block gives 41 µs.
|
| 269 |
+
- **The mixer LNs are DRAM-bound**, not compute-bound: the neighbour LN with its residual moves 42 MB (x, residual, h,
|
| 270 |
+
y; fp32) in ~190 µs. Removing that traffic needs the entity block resident in L1 (MK-E, item 7).
|
| 271 |
+
- **The decoder LN is latency-bound** on 11 cores (one tile row each): 8 tiles per core of SFPU work in sequence. The
|
| 272 |
+
split-row form (8 cores per row, gather / fold / broadcast with semaphores) is bit-identical and 24 % faster with
|
| 273 |
+
the residual (38.5 -> 29.4 µs); its own floor is the root's two ordered 8-tile folds plus two broadcast rounds.
|
| 274 |
+
- **A RISC-V fill loop is not free.** Broadcasting one fp32 statistic tile over its columns with 1,024 volatile L1
|
| 275 |
+
stores on a data-movement core sat on the critical path of every LN row twice. Storing the SFPU reduce's result
|
| 276 |
+
(already replicated across the 8 lanes by its butterfly) to all four column groups instead costs three more SFPU
|
| 277 |
+
stores: −2.24 ms over the plan (decoder split LN 29.4 -> 21.6 µs, neighbour mixer LN 191.9 -> 158.3 µs).
|
| 278 |
+
|
| 279 |
+
### Profile (after r2.3, `d7f84e1`, one traced replay; `profile_r2/`)
|
| 280 |
+
|
| 281 |
+
2,304 programs (from 6,282), kernel sum 46.52 ms, span 48.06 ms (b2b 47.39 ms). By op code: Matmul 16.28 ms (759
|
| 282 |
+
programs), GenericOp (the new kernels) 13.34 ms (540), BinaryNg 7.80 ms (353), Softmax 3.48 ms (72), Unary 2.03 ms
|
| 283 |
+
(191), Transpose 1.36 ms (84). Largest single items: the decoder LN programs on 11 cores (5.6 ms incl. the residual
|
| 284 |
+
form; r2.6 then spread them over 88 cores), the mixer LNs (3.5 ms on 120 cores, DRAM-bound), the score scale + mask
|
| 285 |
+
programs (4.8 ms decoder + 0.7 ms fusion; r2.4 removed half of that), the softmax (3.5 ms), the K = 1024 K-concatenated
|
| 286 |
+
matmul (2.5 ms) and its operand build (1.6 ms, writes a 4.7 MB operand), the attention matmuls (~3.4 ms), the mixer
|
| 287 |
+
transposes (1.3 ms) and the encoder pre-projection split glue (1.2 ms).
|
| 288 |
+
|
| 289 |
+
## Round 3
|
| 290 |
+
|
| 291 |
+
Date 2026-10-10, plan `OPT_PLAN.md` §0.2 (the ranked list after round 2: softmax fusion, encoder split linears, then
|
| 292 |
+
the attention and operand traffic the round-3 profile pointed at). The device was shared with METEOR debugging jobs
|
| 293 |
+
and other agents' optimization jobs; none of this round's jobs hung, timed out or ran on a faulted chip (every job
|
| 294 |
+
checks the FAULT marker after taking the lock). Every new kernel or kernel mode went through the hang protocol before
|
| 295 |
+
it touched the plan: smallest shapes under `TT_METAL_WATCHER=2`, eager twice, then the real shapes with a traced
|
| 296 |
+
replay (`logs/diffusion-planner/opt_r3/*_tiny_watcher.log`, `*_real.log`).
|
| 297 |
+
|
| 298 |
+
Method, per step (`logs/diffusion-planner/opt_r3/`; scripts in `scripts/`): a device micro-check of the kernel against
|
| 299 |
+
the stock programs it replaces (`smsm_check.py`, `fattn_check.py`, `kemit_check.py`, `lntr_check.py`,
|
| 300 |
+
`linact_check.py`: bit comparison, eager repeat, trace replay, µs per call over a 20-call trace), then `ab.sh <tag>
|
| 301 |
+
<KNOB> <off> <on>` (device suite with the gates read-only and the per-scene report, 3 alternating `bench.py` pairs with
|
| 302 |
+
raw readback dumps, `cmp.py` bit comparison and 99-scene comparison). Each A/B ran on the previous step's defaults
|
| 303 |
+
(the knob of a not-yet-committed step set by the chain script); every job ran from a snapshot of the code.
|
| 304 |
+
|
| 305 |
+
| A/B (same session, 3 pairs, kashiwanoha_dense) | off: b2b / trace / e2e ms | on: b2b / trace / e2e ms | Δ b2b | outputs | kept |
|
| 306 |
+
|---|---|---|---:|---|---|
|
| 307 |
+
| `ATTN_SMSM` 0 / 1 (round-2 defaults) | 40.89 / 40.96 / 51.9 | 39.32 / 39.39 / 50.3 | −1.56 | bit-identical | yes |
|
| 308 |
+
| `ENC_KCAT` 0 / 1 | 39.32 / 39.38 / 50.3 | 37.62 / 37.68 / 48.4 | −1.70 | final_x0 rel 2.9e-3; 99 scenes ego max 0.327 -> 0.346 m, mean 0.149 -> 0.154 m, nb 0.104 -> 0.094 m; 44 / 44 green | yes |
|
| 309 |
+
| `LIN_ACT` 0 / 1 (ENC_KCAT=0) | 39.32 / 39.38 / 50.3 | 39.19 / 39.25 / 50.1 | −0.13 | final_x0 rel 3.1e-3; ego max 0.327 -> 0.392 m, mean 0.149 -> 0.174 m; 44 / 44 green | no |
|
| 310 |
+
| `ATTN_SMSM` 1 / 2 (ENC_KCAT=1) | 37.62 / 37.68 / 48.6 | 37.59 / 37.65 / 48.9 | −0.03 | bit-identical | no (noise level) |
|
| 311 |
+
| `ATTN_FUSED` 0 / 1 | 37.62 / 37.68 / 48.9 | 32.19 / 32.25 / 43.2 | −5.43 | bit-identical | yes |
|
| 312 |
+
| `KCAT_EMIT` 0 / 1 | 32.19 / 32.25 / 43.1 | 31.24 / 31.30 / 42.4 | −0.95 | bit-identical | yes |
|
| 313 |
+
| `LN_TR` 0 / 1 (KCAT_EMIT=1) | 31.24 / 31.30 / 42.1 | 30.16 / 30.22 / 41.0 | −1.08 | bit-identical | yes |
|
| 314 |
+
|
| 315 |
+
AICLK 1350 MHz (min 1343) in every run; the device rows repeat to ±0.01 ms within a session. Final check of the
|
| 316 |
+
shipped defaults (`0b750e3`, no knob set; `logs/diffusion-planner/opt_r3/final/`): 44 / 44, readback and the 99
|
| 317 |
+
per-scene numbers identical to the `LN_TR` on side; bench b2b 30.16 ms, trace p50 30.22 ms, e2e p50 40.84 ms.
|
| 318 |
+
|
| 319 |
+
Kernel micro-benchmarks (µs per call, 20-call trace; every row bit-identical to the stock chain it replaces):
|
| 320 |
+
|
| 321 |
+
| kernel | shape | stock chain | kernel |
|
| 322 |
+
|---|---|---:|---:|
|
| 323 |
+
| scale + mask + softmax (`ATTN_SMSM`) vs smask + `ttnn.softmax` | 8 x 352 x 352 / 8 x 352 x 576 / 8 x 576 x 576 | 62.0 / 101.6 / 149.9 | 50.6 / 73.2 / 106.1 |
|
| 324 |
+
| fused attention (`ATTN_FUSED`) vs head split + Q·Kᵀ + smsm + P·V + head merge | decoder self (qkv in place) / cross (hoisted K, V heads) / fusion (q, kv in place) | 110.6 / 144.6 / 214.6 | 51.0 / 63.0 / 106.1 |
|
| 325 |
+
| fused attention emitting the out-linear operand (`KCAT_EMIT`) vs fused attention + `kcat_operand` | decoder self | 57.9 | 51.7 |
|
| 326 |
+
| split-row LN emitting the operand vs split-row LN + `kcat_operand` | [352, 256] | 26.2 | 22.1 |
|
| 327 |
+
| LN writing its output per-entity transposed (`LN_TR`) vs LN + `ttnn.transpose` | neighbour [320, 64, 128] / lane [140, 64, 128] | 149.5 / 80.6 | 102.6 / 61.8 |
|
| 328 |
+
| LN (+ residual) reading the residual transposed vs `ttnn.transpose` + LN | neighbour / lane | 200.1 / 102.6 | 147.2 / 81.1 |
|
| 329 |
+
| mixer fc1 + GELU in the epilogue (`LIN_ACT`, not bit-identical) vs linear + GELU program | [40960, 64] @ [64, 64] / [20480, 128] @ [128, 128] | 107.8 / 106.4 | 95.8 / 108.7 |
|
| 330 |
+
|
| 331 |
+
### Findings (measured)
|
| 332 |
+
|
| 333 |
+
- **The attention was DRAM-bound, not compute-bound.** The fused attention costs about what the scale / mask /
|
| 334 |
+
softmax kernel alone cost (51.0 vs 50.6 µs self-attention), although it adds Q·Kᵀ and P·V: the time of the five
|
| 335 |
+
stock programs was the round trips of the `[1, 8, Sq, Sk]` fp32 scores and probabilities (4-8 MB each) and the
|
| 336 |
+
head split / merge copies. With the scores in L1 the per-core critical path is the stock softmax arithmetic on one
|
| 337 |
+
tile row; the decoder attention went from ~255 to ~114 µs per block.
|
| 338 |
+
- **Stock LLK sequences compose bit-identically inside one kernel.** The `kernel_lib` softmax calls (reduce,
|
| 339 |
+
bcast-sub + exp, reduce + precise reciprocal, bcast-mul) work on plain CB ids in a `generic_op` and give the stock
|
| 340 |
+
softmax bit for bit; `matmul_block` with the in1 transpose flag reproduces the stock Q·Kᵀ (K = 32 in one block) and
|
| 341 |
+
a key-ordered DEST accumulation reproduces the stock P·V (the whole K in one block); the score value never needs to
|
| 342 |
+
leave DEST between Q·Kᵀ and the scale / mask SFPU ops. The global `UnpackToDestEn` (set when any CB unpacks to DEST)
|
| 343 |
+
did not perturb the FPU phases.
|
| 344 |
+
- **The fp32 SFPU multiply by an immediate equals the multiply by a filled tile** (`mul_unary_tile` vs
|
| 345 |
+
`mul_binary_tile`, fp32 DEST): bit-identical, but only 0.03 ms of the plan; kept as the fused attention's form.
|
| 346 |
+
- **Producers can emit the consumer's operand.** Writing `[x_hi | x_hi | x_lo | 1]` from the split LN / fused attention
|
| 347 |
+
(the `kcat_operand` LLK calls on the packed fp32 tile) costs 1-2 µs per producer and removes a 6.8 µs program plus
|
| 348 |
+
its 0.6 µs gap: −0.95 ms over 198 programs.
|
| 349 |
+
- **The stock fp32 `ttnn.transpose` is exact** (unpack-to-dest + `transpose_dest`), so the mixer's transposes fold
|
| 350 |
+
into the LN kernels with `transpose_tile` bit for bit; the scattered tile writes cost nothing measurable (the LN is
|
| 351 |
+
DRAM-bound either way): −1.08 ms.
|
| 352 |
+
- **A more accurate GELU still moves the plan.** `LIN_ACT` applies the GELU before the bf16 rounding (per-linear rms
|
| 353 |
+
error 2.06e-3 -> 1.66e-3 vs fp64), yet the 99-scene ego mean rises 0.149 -> 0.174 m: the chaotic scene sensitivity
|
| 354 |
+
of PORT_LOG known issues again; at −0.13 ms it is not worth that.
|
| 355 |
+
- **The encoder K-concatenation pays off more than estimated** (−1.70 vs −1 to −2 ms): beyond the 3-pass glue, the
|
| 356 |
+
pre-projection matmuls on padded `[1, E, T, K]` rows read the operand once instead of three times.
|
| 357 |
+
|
| 358 |
+
## Round 4
|
| 359 |
+
|
| 360 |
+
Date 2026-10-10, plan `OPT_PLAN.md` §0.3 (item 1, the K = 1024 split operand traffic, then what the measurements of
|
| 361 |
+
that step pointed at: the decoder's DRAM round trips). The device was shared with METEOR and other bundles' jobs;
|
| 362 |
+
no job of this round hung, timed out or ran on a faulted chip (each job checks the FAULT marker after taking the
|
| 363 |
+
lock). New kernel modes went through the hang protocol (smallest shapes under `TT_METAL_WATCHER=2`, eager twice,
|
| 364 |
+
then the real shapes with a trace replay: `logs/diffusion-planner/opt_r4/*_tiny_watcher.log`, `*_real.log`).
|
| 365 |
+
|
| 366 |
+
Method as in round 3 (`logs/diffusion-planner/opt_r4/`, scripts in `scripts/`): a device micro-check per kernel
|
| 367 |
+
mode (`kl1_check.py`, `kl1_sweep.py`, `ko_check.py`, `fl1_check.py`, `dl1_check.py`: bit comparison with the shipped
|
| 368 |
+
path, eager repeat, trace replay, µs per call over a 20-call trace), then `ab.sh <tag> <KNOB> <off> <on>` (device
|
| 369 |
+
suite with the gates read-only and the per-scene report, 3 alternating `bench.py` pairs with raw readback dumps,
|
| 370 |
+
`cmp.py`). Each A/B ran on the previous step's defaults, every job from a snapshot of the code.
|
| 371 |
+
|
| 372 |
+
| A/B (same session, 3 pairs, kashiwanoha_dense) | off: b2b / trace / e2e ms | on: b2b / trace / e2e ms | Δ b2b | outputs | kept |
|
| 373 |
+
|---|---|---|---:|---|---|
|
| 374 |
+
| `KCAT_L1` 0 / 1 (round-3 defaults) | 30.16 / 30.22 / 40.78 | 28.53 / 28.59 / 39.30 | −1.63 | bit-identical | yes |
|
| 375 |
+
| `KCAT_ACT_ONCE` 0 / 1 | 28.53 / 28.62 / 44.4 | 28.59 / 28.66 / 43.6 | +0.06 | bit-identical | no |
|
| 376 |
+
| `ATTN_L1` 0 / 1 | 28.53 / 28.59 / 39.33 | 28.12 / 28.18 / 38.93 | −0.41 | bit-identical | yes |
|
| 377 |
+
| `DEC_L1` 0 / 1 | 28.12 / 28.18 / 38.83 | 27.11 / 27.16 / 37.92 | −1.01 | bit-identical | yes |
|
| 378 |
+
| `ENC_L1` 0 / 1 | 27.11 / 27.16 / 37.70 | 28.77 / 28.84 / 39.49 | +1.67 | final_x0 rel 8e-3, ego max 0.346 -> 0.379 m | no |
|
| 379 |
+
| `DEC_L1` 1 / 2 | 27.11 / 27.17 / 37.91 | 27.02 / 27.07 / 37.90 | −0.09 | bit-identical | yes |
|
| 380 |
+
| `FUS_L1` 0 / 1 | 27.02 / 27.07 / 37.76 | 27.40 / 27.47 / 38.16 | +0.39 | final_x0 rel 5e-4, ego max 0.346 -> 0.364 m | no |
|
| 381 |
+
| `ATTN_L1` 1 / 2 | 27.02 / 27.08 / 37.82 | 26.62 / 26.69 / 37.71 | −0.39 | bit-identical | yes |
|
| 382 |
+
|
| 383 |
+
AICLK 1350 MHz (min 1343) in every run; the device rows repeat to ±0.01 ms within a session (the `KCAT_ACT_ONCE`
|
| 384 |
+
e2e column ran on a busier host). Final check of the shipped defaults (`93428bb`, no knob set;
|
| 385 |
+
`logs/diffusion-planner/opt_r4/final/`): 44 / 44, readback and the 99 per-scene numbers identical to round 3; bench
|
| 386 |
+
b2b 26.62 ms, trace p50 26.68 ms, e2e p50 37.36 ms.
|
| 387 |
+
|
| 388 |
+
Kernel micro-benchmarks (µs per call, 20-call trace; every row bit-identical to the shipped path):
|
| 389 |
+
|
| 390 |
+
| check | shape | DRAM (shipped) | L1 |
|
| 391 |
+
|---|---|---:|---:|
|
| 392 |
+
| K = 1024 operand build + matmul (`KCAT_L1`; best of 7 / 6 sharded configs) | `[352, 1024]` -> X' `[352, 3328]` @ `[3328, 256]`, 2d M2 N1 kb13 | 62.6 (23.4 + 41.4) | 41.3 (19.0 + 24.2) |
|
| 393 |
+
| same | final p4 N = 324: K' 104 -> 99 tiles (11 slices), 2d M2 N1 kb9 | 62.3 | 41.8 |
|
| 394 |
+
| K = 256 operand L1-sharded (`kl1_sweep.py`: K' padded to 26-32 to match the N blocks) | N = 256 / 768 / 1024, matmul only | 15.3 / 20.4 / 24.2 | 14.2 / 29.7 / 34.7 (best) |
|
| 395 |
+
| operand build, GELU once (`KCAT_ACT_ONCE`) | K = 1024 / 512, gelu_tanh | 23.2 / 14.8 | 23.5 / 13.2 |
|
| 396 |
+
| fused attention, inputs (+ mask) in L1 (`ATTN_L1`) | self 8 x 352 x 352 / cross 8 x 352 x 576 / fusion 8 x 576 x 576 | 51.0 / 61.0 / 106.2 | 35.1 / 54.2 / 88.4 |
|
| 397 |
+
| split-row LN + residual + gate + emitted operand, inputs and outputs in L1 (`DEC_L1`) | `[352, 256]` | 24.1 | 20.6 |
|
| 398 |
+
| K = 256 K-concatenated matmuls, in0 and output in L1 interleaved (`DEC_L1`) | N = 256 / 768 / 1024 | 14.9 / 19.7 / 24.1 | 12.7 / 15.6 / 20.0 |
|
| 399 |
+
| fused self attention writing the operand to L1 (`DEC_L1`) | 8 x 352 x 352 | 37.6 | 36.0 |
|
| 400 |
+
|
| 401 |
+
### Findings (measured)
|
| 402 |
+
|
| 403 |
+
- **The decoder was bound by DRAM round trips of small tensors, not by compute.** The K = 1024 operand (4.7 MB)
|
| 404 |
+
costs as much to write and re-read as the matmul costs to compute: written to L1 shards and read in place, the
|
| 405 |
+
matmul alone drops 41.4 -> 24.2 µs. The same holds, at smaller scale, for every decoder intermediate (0.4-1.4 MB):
|
| 406 |
+
L1-interleaved inputs and outputs take 2-4 µs off each split LN, matmul and attention program.
|
| 407 |
+
- **K / V re-reads.** The fused attention's 11 query-row cores of a head each read all K / V tiles; from L1 banks
|
| 408 |
+
instead of DRAM the self attention drops 51 -> 35 µs. No multicast kernel was needed for most of the estimated gain.
|
| 409 |
+
- **L1 is not faster for large tensors.** The mixer's 5-10 MB intermediates in L1 interleaved made the trunk +1.67 ms
|
| 410 |
+
slower (120 cores reading every bank: NoC contention instead of 8 DRAM channels).
|
| 411 |
+
- **Stock ops change program (and numerics) with an L1 output.** `ttnn.linear` with an L1 output ran the fusion kv
|
| 412 |
+
projection on 6 cores with a separate bias add (54 vs 16 µs; bit-identical here), and the mixer / fusion stock
|
| 413 |
+
matmuls, unary GELU and `ttnn.layer_norm` with L1 outputs were not bit-identical. L1 residency is safe and exact
|
| 414 |
+
for the bundle's own `generic_op` kernels and for matmuls with explicit program configs; stock ops need a profile
|
| 415 |
+
after the change (the first round-4 profile caught the 6-core linear).
|
| 416 |
+
- **The operand build is write-bound at K = 1024**: running the GELU once instead of twice saves 1.6 µs at K = 512
|
| 417 |
+
but nothing at K = 1024 (not kept).
|
| 418 |
+
- **Uneven block shards work**: X' `[352, ·]` over 6 shard rows of 64 (the last one half full) is read correctly by
|
| 419 |
+
the sharded-in0 matmul and written correctly through the TensorAccessor of the generic_op writer.
|
| 420 |
+
|
| 421 |
+
## Round 5
|
| 422 |
+
|
| 423 |
+
Date 2026-10-10, plan `OPT_PLAN.md` §0.4 (item 1, exact compaction, decoder buckets first; item 2, the host items).
|
| 424 |
+
The device was shared with METEOR and other bundles' jobs (several of this round's jobs waited for the lock behind
|
| 425 |
+
a CenterPoint A/B); no job of this round hung, timed out or ran on a faulted chip. One gate job failed fast with a
|
| 426 |
+
host-side `TT_FATAL` (a slice end past the trimmed `cs` width: `ab_trim/gates_fail_slice.log`), fixed before any
|
| 427 |
+
A/B. The one new kernel shape (the split-row fp32 LN at W = 1024, Wt = 32 members per row, which the r32 / r64 / r96
|
| 428 |
+
buckets use for the final-layer LN) went through the hang protocol first: `TT_METAL_WATCHER=2`, rows 32 and 96,
|
| 429 |
+
eager twice and a 2-call trace, bit-identical to the one-core-per-row kernel (`opt_r5/ln1024_bringup.log`).
|
| 430 |
+
|
| 431 |
+
Method (`logs/diffusion-planner/opt_r5/`, scripts in `scripts/`): `compact_check.py` (one process with the bucket
|
| 432 |
+
traces: every scene through the full `plan` and through its bucket, final_x0 on the needed rows, logits and the 11 ego
|
| 433 |
+
iterates compared bit for bit, then b2b of every variant), `host_equiv.py` (old and new host package side by side on
|
| 434 |
+
134 inputs: the `prepare` dataclasses, `plan_inputs` and six `make_output` calls per input with random and
|
| 435 |
+
force-stop trajectories and two parameter sets, compared byte for byte), then per step the device suite with the
|
| 436 |
+
gates read-only and the per-scene report, and 3 alternating `bench.py` pairs on 3 scenes (`ab_*/`).
|
| 437 |
+
|
| 438 |
+
| A/B (same session, 3 pairs; kashiwanoha_dense / straight_road / scene-0103_kf14) | off: b2b ms | on: b2b ms | off -> on e2e p50 (kashiwanoha) | outputs | kept |
|
| 439 |
+
|---|---|---|---|---|---|
|
| 440 |
+
| `COMPACT` 0 / 1 (round-4 defaults) | 26.62 / 26.62 / 26.62 | 23.25 (r96) / 21.85 (r32) / 22.91 (r64) | 38.28 -> 33.79 | bit-identical on the read rows | yes |
|
| 441 |
+
| `COMPACT` 1 / 2 | 23.25 / 21.85 / 22.91 | 19.99 / 17.42 / 19.23 | 33.21 -> 30.14 | bit-identical on the read rows | yes |
|
| 442 |
+
| `HOST_FAST` 0 / 1 | 19.99 / 17.42 / 19.23 | 19.99 / 17.42 / 19.23 | 30.15 -> 27.31 (straight_road 27.05 -> 23.87, scene-0103 29.20 -> 26.22) | bit-exact host arrays | yes |
|
| 443 |
+
| `INPUT_TRIM` 0 / 1 | 19.99 / 17.42 / 19.23 | 20.00 / 17.44 / 19.25 | 27.22 -> 26.84 (host_in 1.95 -> 1.58, h2d 0.94 -> 0.76, pack 0.38 -> 0.52) | readback bit-identical | yes |
|
| 444 |
+
|
| 445 |
+
AICLK 1350 MHz (min 1343) in every run; the device rows repeat to ±0.003 ms within a session. The last A/B's `on`
|
| 446 |
+
runs are the shipped defaults (`e933166`): 44 / 44 (`ab_trim/gates.log`), the 99 per-scene numbers identical to
|
| 447 |
+
rounds 3 / 4.
|
| 448 |
+
|
| 449 |
+
Per-bucket device time (`compact2_check.log`, b2b, before `INPUT_TRIM`'s +0.014 ms):
|
| 450 |
+
|
| 451 |
+
| variant | rows (decoder agents / neighbour entities) | b2b ms | 99 gate scenes in it |
|
| 452 |
+
|---|---|---:|---:|
|
| 453 |
+
| `plan_r32` | 32 | 17.42 | 53 |
|
| 454 |
+
| `plan_r64` | 64 | 19.24 | 44 |
|
| 455 |
+
| `plan_r96` | 96 | 19.99 | 2 |
|
| 456 |
+
| `plan_r128` / `plan_r192` | 128 / 192 | 21.11 / 22.30 | 0 |
|
| 457 |
+
| `plan` | 352 / 320 | 26.65 | 0 |
|
| 458 |
+
|
| 459 |
+
### Findings (measured)
|
| 460 |
+
|
| 461 |
+
- **Compaction is exact here and needs no new numerics.** Valid agents and valid neighbour tokens are a prefix in
|
| 462 |
+
all 99 scenes, but the bucket does not rely on it: it covers the last needed row. The rows past R
|
| 463 |
+
are masked self-attention keys (exp gives exact zeros, added in the same tile order) and invalid neighbour tokens
|
| 464 |
+
(multiplied by `token_valid` = 0), so dropping them changes no read value. Every derived matmul config keeps the
|
| 465 |
+
swept K blocking; only `per_core_M` is fitted, which leaves each output element's accumulation unchanged.
|
| 466 |
+
- **The decoder is latency-bound, not row-bound.** From 352 to 96 rows the decoder blocks fall only 12.0 -> 9.5 ms
|
| 467 |
+
(the profile below): its 616 programs per plan run at 8-43 µs each almost regardless of the row count (split-row LN
|
| 468 |
+
16.7 µs on 24 cores, self attention 15.9 µs, cross attention 43.3 µs on 24 cores, K = 256 matmuls 7.8-14.6 µs).
|
| 469 |
+
Smaller buckets mostly remove the encoder's neighbour work (3.9 -> 1.5 ms for the trunk at r96) and the K = 1024
|
| 470 |
+
operand traffic. What remains is per-program latency: the MK-D megakernel's target.
|
| 471 |
+
- **Host time halves with plain numpy.** The node-faithful loops (trajectory smoothing, per-tensor normalisation
|
| 472 |
+
with boolean-mask assignment, all 320 neighbours denormalised and converted) cost 6.8 ms; the same arithmetic
|
| 473 |
+
vectorised in the same float64 / float32 order costs 3.9 ms, byte for byte equal. `math.hypot` with three arguments
|
| 474 |
+
stays a Python loop over 80 points (`np.hypot` rounds differently in the last bit).
|
| 475 |
+
- **Upload trimming is small.** The tilization of the 19 inputs is ~1.6 ms of `host_in`; the all-zero test of every
|
| 476 |
+
input costs 0.14 ms of `pack`, so the net gain is 0.37 ms e2e.
|
| 477 |
|
| 478 |
## Known hangs (all rounds)
|
| 479 |
|
| 480 |
| when (UTC) | command | cause | status |
|
| 481 |
|---|---|---|---|
|
| 482 |
+
| – | – | none: no hang, timeout, reset or FAULT marker in any device job of the port, the baseline, the release docs or rounds 1-5 (the FAULT markers of 2026-10-10 came from METEOR jobs of another agent; this bundle's queued jobs waited and resumed after the orchestrator cleared them). Round 3: one profile job was cancelled from this session before it ran (a `pkill` of its wrapper script); no FAULT, no device effect | – |
|
| 483 |
|
| 484 |
## Rejected / not kept
|
| 485 |
|
| 486 |
+
Round 4: `KCAT_ACT_ONCE` (bit-identical, +0.06 ms), `ENC_L1` (+1.67 ms, not bit-identical), `FUS_L1` (+0.39 ms, not
|
| 487 |
+
bit-identical), `ATTN_L1=1` (the fusion q / kv linears with L1 outputs on 6 cores; superseded by `ATTN_L1=2`), the
|
| 488 |
+
K = 256 split operands L1 block-sharded (`kl1_sweep.py`: −1.1 µs at N = 256 with K' padded 25 -> 32, slower at N =
|
| 489 |
+
768 / 1024; not wired). All stay selectable (knobs) except the last.
|
| 490 |
+
|
| 491 |
+
Round 3: `LIN_ACT` (the mixer fc1 GELU in the matmul epilogue: −0.13 ms, 99-scene ego mean 0.149 -> 0.174 m, max
|
| 492 |
+
0.327 -> 0.392 m; default 0, selectable) and `ATTN_SMSM=2` (the scale by an immediate in the smsm kernel:
|
| 493 |
+
bit-identical, −0.03 ms; superseded by `ATTN_FUSED`, which uses that form).
|
| 494 |
+
|
| 495 |
+
Round 2: every A/B'd step was kept. Measured on the way and not used: the first `ATTN_SMASK` kernel (an exact fp32
|
| 496 |
+
add of the mask) was not bit-identical, because it drops the TF32 truncation the stock fp32 + bf16 add applies
|
| 497 |
+
(`smask_real.log`: every unmasked score differs); it was replaced by the truncating form before any A/B, and an
|
| 498 |
+
exact-score variant remains a separate precision experiment. `LN_KERNEL=1` (per-tile unpacker / SFPU inits) is 2-4 %
|
| 499 |
+
slower than `LN_KERNEL=2` at kernel level (`ln_real.json` vs `ln_lean.json`) and stays selectable.
|
| 500 |
+
|
| 501 |
+
|
| 502 |
+
Round 1 (measured, gates green, not kept because they are not bit-identical and raise the 99-scene error):
|
| 503 |
+
|
| 504 |
+
- `ATTN_FAST=3`, the attention scale as an SFPU pre-activation of the mask add (one program instead of two over the
|
| 505 |
+
`[1, 8, Sq, Sk]` scores; sweep 75 -> 42 µs per call): −2.5 ms over mode 1, ego max 0.385 m / mean 0.173 m;
|
| 506 |
+
- `ATTN_FAST=2`, the scale on Q before Q·Kᵀ: −1.97 ms, ego max 0.379 m / mean 0.172 m, neighbours 0.098 m.
|
| 507 |
+
|
| 508 |
+
Both stay selectable for timing. The round-2 fused attention kernel replaces these pieces anyway.
|
| 509 |
+
|
| 510 |
+
Earlier, configurations that were measured and are **not** the default, for accuracy (PORT_LOG decision 11,
|
| 511 |
`OPT_BASELINE.md` "What the numerics defaults cost"):
|
| 512 |
|
| 513 |
- the first device round's defaults (split only for the mixer inputs and the decoder pre-projection, bf16 SDPA):
|
|
|
|
| 519 |
- the fastest graph (fused LayerNorm, no split, bf16 SDPA): 43.7 ms, fails the encoder PCC gate (`enc.ego` 0.99856 /
|
| 520 |
0.99261).
|
| 521 |
|
| 522 |
+
## Republish (2026-10-11)
|
| 523 |
+
|
| 524 |
+
Republish prep of the optimized release (no optimization step; evidence under `logs/diffusion-planner/republish/`):
|
| 525 |
+
|
| 526 |
+
| item | result |
|
| 527 |
+
|---|---|
|
| 528 |
+
| ttaw | re-vendored 0.23.0 -> 0.23.2 (common `60b6dd7`: Resize2d tail and `reshape_patch_present`, neither used by this model), `c91ce91`; host suite 108 passed / 45 skipped |
|
| 529 |
+
| tt-metal patches | the shared tree carries `patches/tt-metal-reshape-rm-sys1419.patch` since 2026-10-10 (every round-2..5 device number was measured with it); the bundle now ships it and the image build asserts its marker |
|
| 530 |
+
| device suite (gates read-only) on `bc966b4` | 44 passed; again 44 passed under `TT_METAL_TRACE_ALLOC_TRACKING=1`; 99 per-scene numbers identical to round 5 (`opt_r5/e2e_compact2.json`, 0 of 99 differ) |
|
| 531 |
+
| stage bench (100 iterations) | kashiwanoha_dense r96: b2b **20.00 ms**, one plan 20.05 ms, e2e p50 26.73 ms (p99 27.48), `.npz` path 32.27 ms; straight_road r32: 17.44 / 17.50 / 23.13 ms; scene-0103_kf14 r64: 19.25 / 19.30 / 25.65 ms; AICLK 1350 MHz (min 1343) |
|
| 532 |
+
| **fix `7e003c4`: `dispatch="worker"` crashed at capture** | found by this re-measure: the swept 2-D multicast configs of rounds 1-5 assume the 12-column ETH grid (decoder qkv: 12 output blocks on x), so the 11×10 WORKER grid, the A/B opt-in and the fallback without the ETH patch, failed with `TT_FATAL: Num output blocks along x (12) must be smaller than or equal to the number of columns in compute grid (11)`. Now each config is checked against `compute_with_storage_grid_size()` and falls back to the auto config (`mcast_fits`, host test `test_grid_fit_host.py`). ETH: unchanged configs, device suite 44 passed, 99 numbers identical. WORKER: `test_e2e_device.py` 2 passed with the same 99 numbers as ETH; b2b 20.61 ms (r96) / 18.01 ms (r32), 3 % slower than ETH |
|
| 533 |
+
| load / served | `from_pretrained` 352 s with an empty JIT cache (6 traces), 12.9 s warm; served `/predict` 37.2 ms median (`model()` 26.8 ms, JSON decode 8.6 ms), client round trip 41.6 ms, READY 13.0 s on the host; the bundle's `smoke_test.py` passes against the host server |
|
| 534 |
+
| accuracy re-checks | ORT oracle: 92 nuScenes instants worst ego max 0.346 m / mean 0.154 m, turn 92 / 92; scene-0061 sequence 0.112 / 0.055 m, 33 / 33; open loop ADE / FDE 4.75 / 11.87 m (CPU 4.74 / 11.87), blinker 88 / 92; neighbours of the shipped sample: per-agent max 8.6 cm median, 0.23 m p90, 3.72 m worst |
|
| 535 |
+
| media | the 13 renders of `media/` re-rendered from the optimized p150 outputs |
|
| 536 |
+
|
| 537 |
## Profile at the end (trace replay)
|
| 538 |
|
| 539 |
+
**After round 5** (`e933166`, the shipped defaults, the bench scene's bucket `plan_r96`;
|
| 540 |
+
`logs/diffusion-planner/opt_r5/profile_r5final/`: `analyze.md`, `breakdown.md`, `ops_perf_results.csv.zst`): **1,418
|
| 541 |
+
programs**, kernel sum **19.64 ms**, op-to-op gaps 0.89 ms, span 20.54 ms at 1.35 GHz (b2b 20.00 ms; replay ==
|
| 542 |
+
eager). Stage groups: decoder blocks 9.49 ms (11 evaluations at 96 rows; 12.01 at 352), mixer blocks 5.66 ms (8.12),
|
| 543 |
+
fusion 1.36 ms, mixer pre-projections 1.28 ms (2.08), decoder pre-projection 0.66 ms (1.03), solver 0.31 ms, the rest
|
| 544 |
+
0.9 ms.
|
| 545 |
+
|
| 546 |
+
Largest remaining items (kernel time, share of the 20.5 ms span):
|
| 547 |
+
|
| 548 |
+
| item | ms | share | what limits it / next step |
|
| 549 |
+
|---|---:|---:|---|
|
| 550 |
+
| decoder split-row LNs with residual (154 programs on 24 cores, 16.7 µs) + final LNs | 2.58 + 0.34 | 14.2 % | gather / fold / broadcast latency, row-count independent; MK-D |
|
| 551 |
+
| map mixer trunks: lanes (140 entities) 2.41, line strings (60) 1.19, route (25) 0.61 | 4.21 | 20.5 % | full capacity in every scene (kashiwanoha: 123 / 60 / 17 valid; p50 scenes 100 / 30 / 6): map-entity compaction (est. −1.5 to −2 ms at p50 scenes, ~−0.4 ms on the bench scene) |
|
| 552 |
+
| mixer LNs (36 programs, 120 cores, 59 µs) | 2.13 | 10.4 % | DRAM-bound; MK-E |
|
| 553 |
+
| decoder cross attention (33 x 43.3 µs on 24 cores: 8 heads x 3 query tiles, 18 key tiles each) | 1.43 | 7.0 % | one core per (head, query tile): a key-split form with an ordered gather (bit-identical like the split LN) could use 4x the cores (est. −0.8 ms) |
|
| 554 |
+
| decoder K = 1024 split linears: matmuls 1.45 + operand builds 0.83 | 2.28 | 11.1 % | latency of 24-33-core programs at 96 rows; MK-D |
|
| 555 |
+
| decoder K = 256 split linears (qkv, out, mlp fc1, cross q / out) | 2.26 | 11.0 % | 7.8-14.6 µs per program; MK-D |
|
| 556 |
+
| mixer fc1 GELU programs | 0.71 | 3.5 % | bit-identical fusion needs a matmul epilogue that rounds to bf16 before the GELU |
|
| 557 |
+
| fusion attention (6 x 94 µs) | 0.57 | 2.8 % | – |
|
| 558 |
+
|
| 559 |
+
**After round 4** (`93428bb`, the shipped defaults; `logs/diffusion-planner/opt_r4/profile_r4final/`: `analyze.md`,
|
| 560 |
+
`breakdown.md`, `ops_perf_results.csv.zst`): **1,411 programs**, kernel sum **26.00 ms**, op-to-op gaps 0.89 ms,
|
| 561 |
+
span 26.89 ms at 1.35 GHz (b2b 26.62 ms). Stage groups: decoder blocks 12.01 ms (11 evaluations; 15.61 after round
|
| 562 |
+
3), mixer blocks 8.12 ms, mixer pre-projections 2.08 ms, fusion 1.35 ms, decoder pre-projection 1.03 ms, solver 0.39
|
| 563 |
+
ms, the rest 1.1 ms. (The first round-4 profile, `profile_r4final_c3b6a5a/`, showed the fusion at 1.74 ms: the
|
| 564 |
+
6-core fusion linears fixed by `ATTN_L1=2`.)
|
| 565 |
+
|
| 566 |
+
Largest remaining items (kernel time, share of the 26.9 ms span):
|
| 567 |
+
|
| 568 |
+
| item | ms | share | what limits it / next step |
|
| 569 |
+
|---|---:|---:|---|
|
| 570 |
+
| decoder split-row LNs (4 per block) + fused self attention (88-core GenericOp group) | 4.02 | 14.9 % | gather / ordered fold / broadcast latency (~17-20 µs per LN); MK-D |
|
| 571 |
+
| mixer LNs (36 programs, 120 cores) | 3.11 | 11.6 % | DRAM-bound; L1 was slower (`ENC_L1`); MK-E |
|
| 572 |
+
| decoder K = 256 K-concatenated matmuls (qkv, out, mlp fc1, cross q / out) | 2.82 | 10.5 % | 10-18 µs on 48-72 cores; fp32 HiFi4 compute |
|
| 573 |
+
| decoder K = 1024 split linears: matmuls 1.61 + operand builds 1.27 | 2.88 | 10.7 % | the operand build is write-bound (3 tiles per input tile, gelu_tanh); in-kernel split matmul (MK-D phase) |
|
| 574 |
+
| mixer channel / token matmuls | ~2.6 | 9.7 % | bf16 / fp32 linears at 60-77 µs on 107-117 cores; MK-E |
|
| 575 |
+
| decoder cross attention | 1.11 | 4.1 % | 33.6 µs per call |
|
| 576 |
+
| mixer fc1 GELU programs | 1.04 | 3.9 % | bit-identical fusion needs a custom matmul epilogue (round to bf16 before the GELU) |
|
| 577 |
+
| 1024-wide final-layer LN (11 cores) | 0.66 | 2.4 % | a multi-tile split form (est. −0.3 ms: the 32-tile fold stays serial on the root) |
|
| 578 |
+
| pre-projections (operand builds 0.52, island glue 0.28, matmuls) | 2.08 | 7.7 % | – |
|
| 579 |
+
|
| 580 |
+
**After round 3** (`0b750e3`, the shipped defaults; `logs/diffusion-planner/opt_r3/profile_r3final/`: `analyze.md`,
|
| 581 |
+
`breakdown.md`, `ops_perf_results.csv.zst`): **1,405 programs**, kernel sum **29.71 ms**, op-to-op gaps 0.88 ms,
|
| 582 |
+
span 30.59 ms at 1.35 GHz (b2b 30.16 ms). Stage groups: decoder blocks 15.61 ms (11 evaluations), mixer blocks
|
| 583 |
+
8.15 ms, mixer pre-projections 2.11 ms, fusion 1.41 ms, decoder pre-projection 1.05 ms, solver 0.39 ms, the rest
|
| 584 |
+
1.0 ms. By op code: GenericOp 14.18 ms (438 programs), Matmul 11.90 ms (567), Unary 1.44 ms (114, mostly the mixer
|
| 585 |
+
fc1 GELU), BinaryNg 1.08 ms (137), Transpose 0.21 ms (36, from 84).
|
| 586 |
+
|
| 587 |
+
Largest remaining items (kernel time, share of the 30.6 ms span):
|
| 588 |
+
|
| 589 |
+
| item | ms | share | what limits it / next step |
|
| 590 |
+
|---|---:|---:|---|
|
| 591 |
+
| decoder K = 1024 split linears (mlp fc2 x 2, final p4): matmul `[352, 3328] @ [3328, N]` on 48-66 cores + the operand build after the tanh GELU | 2.95 + 1.61 | 14.9 % | in0 is 4.7 MB of fp32 written and read per call; L1-sharded operand, bf16 `w_lo` (precision change), or the in-kernel split matmul (MK-D building block) |
|
| 592 |
+
| decoder split-row LayerNorms (4 per block, 88 cores) + the 1024-wide final LN (11 cores) | 2.75 + 0.40 + 0.67 | 12.5 % | latency-bound (gather / ordered fold / broadcast, ~20 µs per LN); MK-D |
|
| 593 |
+
| decoder fused attention (self 49, cross 59 µs) + fusion (103 µs) | 3.58 + 0.62 | 13.7 % | the stock softmax arithmetic on one tile row per core (88 of 120 cores busy); K / V re-read by 11 cores per head |
|
| 594 |
+
| decoder K = 256 split linears (qkv, out, mlp fc1, cross q / out) | 3.50 | 11.4 % | small matmuls (12-22 µs on 48-72 cores, configs from the round-2 sweep) |
|
| 595 |
+
| mixer LNs (24 programs, 120 cores) | 3.13 | 10.2 % | DRAM-bound (x, residual, h, y fp32); MK-E (entity block resident in L1) |
|
| 596 |
+
| mixer channel / token matmuls | ~2.6 | 8.5 % | bf16 / fp32 linears at 60-77 µs on 107-117 cores |
|
| 597 |
+
| mixer fc1 GELU programs | 1.05 | 3.4 % | a bit-identical fusion needs a custom matmul epilogue (round to bf16, then GELU); `LIN_ACT` (not bit-identical) rejected |
|
| 598 |
+
| pre-projections (12 operand builds, island glue, matmuls) | 2.11 | 6.9 % | – |
|
| 599 |
+
|
| 600 |
+
**After round 2** (`c42c540`, the shipped defaults; `logs/diffusion-planner/opt_r2/profile_r2final/`: `analyze.md`,
|
| 601 |
+
`breakdown.md`, the per-program CSVs as `.zst`): **2,155 programs**, kernel sum **39.99 ms**, op-to-op gaps 1.47 ms,
|
| 602 |
+
span 41.46 ms at 1.35 GHz (b2b 40.90 ms). Stage groups: decoder blocks 22.45 ms (11 evaluations), mixer blocks
|
| 603 |
+
9.21 ms, mixer pre-projections 3.59 ms, fusion 2.31 ms, decoder pre-projection 1.05 ms, solver 0.38 ms, the rest
|
| 604 |
+
1.0 ms.
|
| 605 |
+
|
| 606 |
+
| op code | after round 1 ms | after round 2 ms (programs) |
|
| 607 |
+
|---|---:|---:|
|
| 608 |
+
| Matmul | 18.1 | 16.28 (759) |
|
| 609 |
+
| GenericOp (the round-2 kernels) | – | 12.91 (612) |
|
| 610 |
+
| Softmax | 3.5 | 3.48 (72) |
|
| 611 |
+
| BinaryNg | 31.5 | 2.29 (209) |
|
| 612 |
+
| Unary / Typecast | 2.8 / 2.4 | 1.44 / 0.42 |
|
| 613 |
+
| Transpose (mixer token mixing) | 1.3 | 1.36 |
|
| 614 |
+
| Reduce (LN means) | 3.7 | 0.09 |
|
| 615 |
+
|
| 616 |
+
Largest remaining single items: the decoder softmax (2.97 ms, 66 x 45 µs on 88 cores), the mixer LN programs (2.88 ms
|
| 617 |
+
on 120 cores, DRAM-bound), the K = 1024 K-concatenated matmuls (2.94 ms) and their operand builds (1.62 ms), the
|
| 618 |
+
decoder LN programs (2.38 ms on 88 cores + 0.67 ms for the 1024-wide final-layer LN, which stays on 11 cores: 352
|
| 619 |
+
rows x 32 tiles do not fit the split form), the score scale + mask passes (2.29 ms decoder + 0.37 ms fusion), the
|
| 620 |
+
attention matmuls (3.40 ms), the encoder pre-projection split glue (1.23 ms BinaryNg) and the mixer transposes
|
| 621 |
+
(1.16 ms).
|
| 622 |
+
|
| 623 |
+
|
| 624 |
+
**After round 1** (`dd6f249`, one traced replay under the device profiler, `logs/diffusion-planner/opt_r1/
|
| 625 |
+
scripts/profile.sh`; numbers from the device-side per-program report `profile_r1/cpp_device_perf_report.csv.zst`,
|
| 626 |
+
because the Tracy post-processing ran out of disk, see below): **6,282 programs** (unchanged: round 1 changed program
|
| 627 |
+
configs, not the op graph), **kernel sum 65.06 ms** (was 99.27), span 70.96 ms at 1.35 GHz (b2b 70.12 ms).
|
| 628 |
+
|
| 629 |
+
| op code | baseline ms | after round 1 ms |
|
| 630 |
+
|---|---:|---:|
|
| 631 |
+
| BinaryNg (fp32 LN decomposition, split-matmul glue, attention scale / mask, residual adds, adaLN gates, solver) | 31.5 | **31.5** (48 % of the kernel sum) |
|
| 632 |
+
| Matmul | 52.3 | **18.1** |
|
| 633 |
+
| Reduce (LN means) | 3.7 | 3.7 |
|
| 634 |
+
| Softmax | 3.5 | 3.5 |
|
| 635 |
+
| Unary / Typecast | 2.8 / 2.4 | 2.8 / 2.4 |
|
| 636 |
+
| Transpose (mixer token mixing) | – | 1.3 |
|
| 637 |
+
|
| 638 |
+
The element-wise decompositions that the gates need are now the larger half of the plan: that is what the round-2
|
| 639 |
+
L2 kernels (fused split matmul, fused fp32 LayerNorm, fp32 attention kernel; `OPT_PLAN.md` items 2-4) remove.
|
| 640 |
+
|
| 641 |
+
Disk incident: the Tracy run of the 6,282-program replay wrote 17 GB of raw logs (`tracy_ops_times.csv` 9.0 GB,
|
| 642 |
+
`profile_log_device.csv` 7.4 GB) and filled the root disk while copying them into the report folder (09:06 UTC, a
|
| 643 |
+
few minutes at 100 %). The partial copy was deleted at once and the raw logs compressed with zstd (17 GB -> 1.4 GB,
|
| 644 |
+
`generated/profiler/diffusion-planner_opt_r1/.logs/`). Future profiles of this model: check free space first (> 25
|
| 645 |
+
GB), and profile the replay only (the eager signposted pass doubles the logs).
|
| 646 |
+
|
| 647 |
The baseline profile (`OPT_BASELINE.md`): 6,282 programs, kernel sum 99.27 ms, span 103.66 ms; by op code Matmul
|
| 648 |
52.3 ms (1,365 programs), BinaryNg 31.5 ms (2,948), Reduce 3.7 ms, Softmax 3.5 ms, Unary 2.8 ms, Typecast 2.4 ms. The
|
| 649 |
ten slowest programs are the neighbour mixer's channel matmuls `[1, 320, 64, 128] @ [128, 128]`, ~970 µs each on 8
|
|
|
|
| 651 |
|
| 652 |
## Remaining backlog (gains on the 102.0 ms replay unless marked e2e; estimates, not measurements)
|
| 653 |
|
| 654 |
+
Updated ranking on the 20.0 ms base (bucket r96) after round 5: `OPT_PLAN.md` §0.5. After round 4 it was §0.4. After round 3 it was §0.3 (the K = 1024 operand traffic, the decoder
|
| 655 |
+
LN latency, then compaction, MK-E and MK-D; the table "After round 3" above). After round 2 it was §0.2 (all done in
|
| 656 |
+
round 3 except the K = 1024 operand).
|
| 657 |
+
The list below is the original one; items 1-5 are done or partly done (rounds 1-2).
|
| 658 |
+
|
| 659 |
From `OPT_BASELINE.md` "Ranked optimization opportunities". The gains are not additive: each item shrinks the base of
|
| 660 |
the next.
|
| 661 |
|
PYTHON.md
CHANGED
|
@@ -68,7 +68,7 @@ the three v5.0 ONNX files and `diffusion_planner.param.json`, with an offline fa
|
|
| 68 |
file and the weights' `major_version == 5` are checked), opens the chip (ETH dispatch, 12×10, 1 CQ; the other open
|
| 69 |
parameters are `DEVICE_DEFAULTS` in `tt_diffusion_planner/device.py`, overridable with `DIFFUSION_PLANNER_*`), reads the
|
| 70 |
ONNX initializers as data, uploads the weights and constants (48.1 MB), builds the graph, then compiles and captures
|
| 71 |
-
the metal
|
| 72 |
(`model.info["device"]["fallback"]` names it). Any other keyword argument is a `TypeError`.
|
| 73 |
|
| 74 |
The numerics are not arguments: the published configuration is the default of the `DIFFUSION_PLANNER_LN_FP32`,
|
|
@@ -78,9 +78,9 @@ accuracy figures of the card until the gates are re-run.
|
|
| 78 |
|
| 79 |
## Warm-up
|
| 80 |
|
| 81 |
-
`from_pretrained` returns a warm model: it builds the graph, runs the plan once eagerly (this first run compiles every kernel into the JIT cache), then captures the whole plan as
|
| 82 |
|
| 83 |
-
Measured on the shipped sample (`model.info["warmup_ms"]`, 2026-10-
|
| 84 |
|
| 85 |
## Call: `model(...)`
|
| 86 |
|
|
@@ -143,7 +143,7 @@ plans, the caller keeps the same state (SERVING.md 3.5 has the details):
|
|
| 143 |
|
| 144 |
## Lifetime and information
|
| 145 |
|
| 146 |
-
- `model.close()` releases the
|
| 147 |
idempotent. `with` calls it for you; an unclosed model is closed when Python exits.
|
| 148 |
- `model.info`: weights (repo, tag, revision, path), device (dispatch, grid, CQs, fallback), variant, warm variants,
|
| 149 |
warm-up times, runtime parameter defaults, the input schema, the numerics options and precision policy in effect, the
|
|
@@ -152,28 +152,29 @@ plans, the caller keeps the same state (SERVING.md 3.5 has the details):
|
|
| 152 |
|
| 153 |
## Speed
|
| 154 |
|
| 155 |
-
Warm calls, batch 1, ETH dispatch, 1 CQ, 12×10, the pinned numerics (`code/scripts/bench.py`, 100 iterations
|
| 156 |
|
| 157 |
-
| stage | shipped sample `kashiwanoha_dense` |
|
| 158 |
|---|---:|
|
| 159 |
-
| `.npz` decode + schema check (path inputs only) |
|
| 160 |
-
| host pre-processing (the node's normalization, the encoder's host features, masks, `x_T`) |
|
| 161 |
-
| packing the
|
| 162 |
-
| **device trace, one blocking plan** | **
|
| 163 |
-
| D2H (one packed read
|
| 164 |
-
| host post-processing (trajectory, predicted paths, turn decision) |
|
| 165 |
-
| **`model(inputs=arrays)` end to end** | **
|
| 166 |
-
| `model(inputs=<.npz path>)` |
|
| 167 |
-
| back-to-back replays (device time per plan) |
|
| 168 |
-
|
| 169 |
-
The device time
|
| 170 |
|
| 171 |
## Limits
|
| 172 |
|
| 173 |
- Batch 1 on the chip; one model per process.
|
| 174 |
- Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history
|
| 175 |
-
and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the
|
| 176 |
-
the
|
|
|
|
| 177 |
- The node's guidance services (start / stop / centerline guidance) are not available (the node's default is off).
|
| 178 |
- Accuracy is agreement with the fp32 CPU reference of the same network (README "Demo & Performances"); the planner's
|
| 179 |
driving quality is the weights' (trained by TIER IV on data that is not public).
|
|
|
|
| 68 |
file and the weights' `major_version == 5` are checked), opens the chip (ETH dispatch, 12×10, 1 CQ; the other open
|
| 69 |
parameters are `DEVICE_DEFAULTS` in `tt_diffusion_planner/device.py`, overridable with `DIFFUSION_PLANNER_*`), reads the
|
| 70 |
ONNX initializers as data, uploads the weights and constants (48.1 MB), builds the graph, then compiles and captures
|
| 71 |
+
the metal traces. If ETH dispatch cannot open (tt-metal without the patch), it warns and falls back to WORKER dispatch
|
| 72 |
(`model.info["device"]["fallback"]` names it). Any other keyword argument is a `TypeError`.
|
| 73 |
|
| 74 |
The numerics are not arguments: the published configuration is the default of the `DIFFUSION_PLANNER_LN_FP32`,
|
|
|
|
| 78 |
|
| 79 |
## Warm-up
|
| 80 |
|
| 81 |
+
`from_pretrained` returns a warm model: it builds the graph, runs the plan once eagerly (this first run compiles every kernel into the JIT cache), then captures the whole plan as metal traces (`warmup_variants="default"`: the full-capacity variant `plan` and one variant per agent bucket, `plan_r32`, `plan_r64`, `plan_r96`, `plan_r128`, `plan_r192`; each call replays the smallest that holds the scene, exact) with program-cache misses forbidden, so no later call compiles anything. `model.warmup()` is idempotent; `warmup_variants="none"` defers the capture to `model.warmup()`.
|
| 82 |
|
| 83 |
+
Measured on the shipped sample (`model.info["warmup_ms"]`, 2026-10-11): the load takes 352 s with an empty JIT cache and 12.9 s with a warm one (build 0.76 s: the ONNX initializers read and 118.9 MB of weights and constants uploaded; warm-up and capture of the 6 traces 7.7 s; the rest is the device open). The first call then takes 37 ms and the second 32 ms (an `.npz` path; the stage bench's steady state: 27 ms p50 for decoded arrays, 32 ms for an `.npz` path). The 6 traces hold 60.3 MB of DRAM (`trace_region_size` 192 MiB).
|
| 84 |
|
| 85 |
## Call: `model(...)`
|
| 86 |
|
|
|
|
| 143 |
|
| 144 |
## Lifetime and information
|
| 145 |
|
| 146 |
+
- `model.close()` releases the traces and the persistent device tensors and closes the chip if the model opened it;
|
| 147 |
idempotent. `with` calls it for you; an unclosed model is closed when Python exits.
|
| 148 |
- `model.info`: weights (repo, tag, revision, path), device (dispatch, grid, CQs, fallback), variant, warm variants,
|
| 149 |
warm-up times, runtime parameter defaults, the input schema, the numerics options and precision policy in effect, the
|
|
|
|
| 152 |
|
| 153 |
## Speed
|
| 154 |
|
| 155 |
+
Warm calls, batch 1, ETH dispatch, 1 CQ, 12×10, the pinned numerics (`code/scripts/bench.py`, 100 iterations, 2026-10-11, the optimized release, on a shared host; p50, with p99 in brackets):
|
| 156 |
|
| 157 |
+
| stage | shipped sample `kashiwanoha_dense` (88 neighbours: bucket r96) |
|
| 158 |
|---|---:|
|
| 159 |
+
| `.npz` decode + schema check (path inputs only) | 5.76 (7.33) ms |
|
| 160 |
+
| host pre-processing (the node's normalization, the encoder's host features, masks, `x_T`) | 2.45 (3.16) ms |
|
| 161 |
+
| packing the persistent trace inputs · ttnn host tensors · H2D | 0.52 · 1.57 · 0.76 ms |
|
| 162 |
+
| **device trace, one blocking plan** | **20.05** (20.08) ms |
|
| 163 |
+
| D2H (one packed read) | 0.22 ms |
|
| 164 |
+
| host post-processing (trajectory, predicted paths, turn decision) | 1.39 (1.84) ms |
|
| 165 |
+
| **`model(inputs=arrays)` end to end** | **26.73** (27.48) ms |
|
| 166 |
+
| `model(inputs=<.npz path>)` | 32.27 (33.61) ms |
|
| 167 |
+
| back-to-back replays (device time per plan) | 20.00 ms = 50.0 plans/s |
|
| 168 |
+
|
| 169 |
+
The device time depends on the agent bucket the scene needs (exact compaction: the smallest of 32 / 64 / 96 / 128 / 192 decoder rows that holds the ego and every valid neighbour row, else the full capacity). Back to back: r32 (≤ 31 neighbours, e.g. `straight_road`) 17.44 ms, r64 (a nuScenes instant with 42 neighbours) 19.25 ms, r96 20.00 ms, r128 21.11 ms, r192 22.30 ms, full capacity (> 191) 26.65 ms (the last three: `OPT_REPORT.md` round 5). The map entities are computed at full capacity in every plan. The first release took 102.04 ms for every scene (`OPT_BASELINE.md`). 2 CQs were not re-measured on this release (at the first release they did not help a synchronous request: `OPT_BASELINE.md`). Where the time goes and what comes next: `OPT_REPORT.md`.
|
| 170 |
|
| 171 |
## Limits
|
| 172 |
|
| 173 |
- Batch 1 on the chip; one model per process.
|
| 174 |
- Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history
|
| 175 |
+
and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the traces; the neighbour trunk
|
| 176 |
+
and the decoder run on the smallest agent bucket that holds the scene (exact), so the device time depends on the
|
| 177 |
+
neighbour count (17.4-26.7 ms); the map entities are computed at full capacity.
|
| 178 |
- The node's guidance services (start / stop / centerline guidance) are not available (the node's default is off).
|
| 179 |
- Accuracy is agreement with the fp32 CPU reference of the same network (README "Demo & Performances"); the planner's
|
| 180 |
driving quality is the weights' (trained by TIER IV on data that is not public).
|
README.md
CHANGED
|
@@ -1,22 +1,22 @@
|
|
| 1 |
---
|
| 2 |
tags:
|
| 3 |
-
- autonomous-driving
|
| 4 |
-
- autoware
|
| 5 |
- blackhole
|
| 6 |
-
- diffusion
|
| 7 |
-
- diffusion-planner
|
| 8 |
-
- motion-prediction
|
| 9 |
- p150
|
| 10 |
-
- planning
|
| 11 |
-
- tenstorrent
|
| 12 |
-
- trajectory-generation
|
| 13 |
- tt-dit-server
|
| 14 |
-
- tt-metal
|
| 15 |
- tt-model-cache
|
| 16 |
-
- tt-model-catalog
|
| 17 |
- tt-model-container
|
| 18 |
-
-
|
| 19 |
- ttnn
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
pipeline_tag: robotics
|
| 21 |
license: apache-2.0
|
| 22 |
license_link: https://huggingface.co/AutowareFoundation/diffusion_planner
|
|
@@ -26,16 +26,16 @@ base_model:
|
|
| 26 |
|
| 27 |
# diffusion-planner-p150
|
| 28 |
|
| 29 |
-
Diffusion Planner v5.0 (Autoware diffusion_planner): the network Autoware deploys in `autoware_diffusion_planner`, ported to one Tenstorrent Blackhole p150 with tt-nn. The whole plan runs on the chip as one metal trace: the scene encoder, the 11 DiT decoder evaluations of the DPM-Solver++(2M) loop with their solver updates, and the turn-indicator head. The Autoware planner tensors in (ego and neighbour histories, lanes, route, polygons, line strings, goal, ego shape, turn-indicator history); an 8 s ego trajectory, the predicted 8 s paths of the neighbours and a turn-indicator command out, with the node's exact pre- and post-processing.
|
| 30 |
Weights: [AutowareFoundation/diffusion_planner `v5.0`](https://huggingface.co/AutowareFoundation/diffusion_planner/tree/423efde67f5414734da43a7ad856c17ceb8b51aa) · Paper: [arXiv:2501.15564](https://arxiv.org/abs/2501.15564) · Autoware package: [autoware_diffusion_planner](https://github.com/autowarefoundation/autoware_universe/tree/9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd/planning/autoware_diffusion_planner) · Training code: [tier4/Diffusion-Planner (TIER IV training fork)](https://github.com/tier4/Diffusion-Planner) · Port: [`code/`](https://huggingface.co/changh95/diffusion-planner-p150/tree/main/code)
|
| 31 |
|
| 32 |
-
Runs on **p150** (mesh `P150`). Configuration: dispatch on the ETH cores, 1 command queue, 12×10 compute grid. Numerics (the default and the `serve.env` pins): fp32 residual streams and solver state, HiFi4 with fp32 accumulation, split hi / lo matmuls for the mixer inputs and every decoder linear, fp32 LayerNorm in the mixers and the decoder
|
| 33 |
|
| 34 |
Packaged and published with [tt-model-manager](https://github.com/tenstorrent/tt-model-manager) 0.1.0 (manifest schema 5.1).
|
| 35 |
|
| 36 |
## Quickstart (Python)
|
| 37 |
|
| 38 |
-
Prerequisite: a tt-metal / ttnn environment at tt-metal [`44d66500520`](https://github.com/tenstorrent/tt-metal/commit/44d66500520fda9f2c7060c0f6b41ec48f7ab37e) with [`patches/tt-metal-eth-dispatch.patch`](patches/tt-metal-eth-dispatch.patch) applied. ttnn is not on PyPI.
|
| 39 |
|
| 40 |
```bash
|
| 41 |
hf download changh95/diffusion-planner-p150 --exclude "image/*" --local-dir diffusion-planner-p150 && cd diffusion-planner-p150
|
|
@@ -56,8 +56,8 @@ print(out.poses[:5])
|
|
| 56 |
print(out.turn_indicator["command_name"], out.predicted_agents.shape)
|
| 57 |
```
|
| 58 |
|
| 59 |
-
- `from_pretrained` downloads the three v5.0 ONNX files and `diffusion_planner.param.json` (58.9 MB) of [`AutowareFoundation/diffusion_planner`](https://huggingface.co/AutowareFoundation/diffusion_planner) at the pinned commit `423efde67f5` (tag `v5.0`) to your HF cache (no token needed), checks their sha256, opens the chip, builds the graph and captures the metal
|
| 60 |
-
- The
|
| 61 |
- The `with` block releases the trace and closes the chip. Without `with`, call `model.close()`.
|
| 62 |
|
| 63 |
| | |
|
|
@@ -70,6 +70,7 @@ print(out.turn_indicator["command_name"], out.predicted_agents.shape)
|
|
| 70 |
- The API gives the same output as the HTTP server `/predict`: both share the decoders, the device trace and the host post-processing (checked on the device by `test_api_equals_server`).
|
| 71 |
- The API is stateless: one call is one plan. What the Autoware node keeps between plans stays with the caller: the turn-indicator hold window (1.0 s; the node's manager ships as `tt_diffusion_planner.host.postprocess.TurnIndicatorManager`), the initial solver state `sampled_trajectories` (zeros = the node's default temperature 0; noise for a temperature > 0; the previous plan for the RTC prefix) and the agent and ego histories. Details: [`code/PYTHON.md`](code/PYTHON.md) "What the caller keeps".
|
| 72 |
- One model uses one chip; calls from several threads are serialised.
|
|
|
|
| 73 |
- Full reference: [`code/PYTHON.md`](code/PYTHON.md). Runnable example: [`examples/quickstart.py`](examples/quickstart.py) (also writes `quickstart_bev.png`, the input tensors and the plan from above).
|
| 74 |
|
| 75 |
## Serving (HTTP)
|
|
@@ -95,11 +96,11 @@ The response for the shipped sample (served on the p150; trajectory cut to 3 of
|
|
| 95 |
"model": "diffusion-planner-p150",
|
| 96 |
"frame_id": "base_link",
|
| 97 |
"meta": {"predicted_agent_columns": ["x", "y", "yaw", "cos", "sin"], "predicted_agent_rows": [0, 1, 2, "... 85 more"], "force_stop": false, "time_from_start_s": [0.1, 0.2, 0.3, "... 77 more"], "valid_counts": {"ego": 1, "neighbor": 88, "static": 0, "lane": 123, "route": 17, "polygon": 0, "line_string": 60, "goal": 1, "ego_shape": 1, "turn": 1}},
|
| 98 |
-
"timing_ms": {"preprocess":
|
| 99 |
"num_poses": 80,
|
| 100 |
"columns": ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"],
|
| 101 |
-
"trajectory": [[0.
|
| 102 |
-
"turn_indicator": {"command": 1, "command_name": "DISABLE", "keep_selected": true, "held": false, "logits": [-15.
|
| 103 |
"predicted_agents": {"format": "npz", "key": "predicted_agents", "dtype": "float32", "shape": [88, 80, 5], "data": "<base64 npz>"}
|
| 104 |
}
|
| 105 |
```
|
|
@@ -132,40 +133,41 @@ nuScenes renders: rendered from the nuScenes dataset (v1.0-mini, CAN bus expansi
|
|
| 132 |
|
| 133 |
## Demo & Performances
|
| 134 |
|
| 135 |
-
Warm, batch 1
|
| 136 |
|
| 137 |
| Metric | Performance |
|
| 138 |
|---|---:|
|
| 139 |
-
| Agreement with the fp32 CPU reference, shipped sample `kashiwanoha_dense` (8 s ego plan, 88 neighbours) | ego max **2.
|
| 140 |
-
| Agreement, shipped sample `straight_road` | ego max
|
| 141 |
-
| Agreement, all 99 gated scenes (2 samples, 5 research scenes, 92 nuScenes-mini instants) | worst ego max **0.
|
| 142 |
-
| Agreement with an independent oracle: the research pipeline's ONNX Runtime outputs (raw x0), 92 nuScenes instants + the 33-plan scene-0061 sequence | worst ego max 0.
|
| 143 |
-
| Module PCC vs the fp32 reference (encoder categories, encoding, teacher-forced decoder evaluation; gate 0.999) | ≥ 0.
|
| 144 |
| Open-loop vs the nuScenes log, 92 instants (a sanity check against one human driver, not a planning metric) | p150 ADE / FDE at 8 s 4.75 / 11.87 m (CPU reference 4.74 / 11.87; constant velocity 4.99 / 13.06); turn command = logged blinker 88 / 92 (CPU 88 / 92) |
|
| 145 |
-
| Python `model()` call, shipped sample (host pre-processing, H2D, trace, D2H, host post-processing) | **
|
| 146 |
-
| Served `/predict` `timing_ms.total` (uvicorn on the host, the shipped sample) | **
|
| 147 |
-
| Served client round trip, loopback (base64 `.npz` request, 0.15 MB) |
|
| 148 |
-
| Device trace, one blocking plan (encoder + 11 DiT evaluations + 10 solver updates + turn head) | **
|
| 149 |
-
| Back-to-back trace replays | **
|
| 150 |
-
| Host pre-processing · pack · host tensors · H2D · D2H · host post-processing |
|
| 151 |
-
| `from_pretrained` load: empty JIT cache / warm cache |
|
| 152 |
|
| 153 |
-
All numbers in this table were measured with dispatch on the ETH cores, 1 command queue and a 12×10 compute grid on one p150, with the pinned numerics of `serve.env`. Accuracy is agreement with the fp32 CPU reference of the same Autoware network (same weights, same pre- and post-processing); no dataset-level accuracy is claimed: the paper's benchmark is nuPlan closed loop (an account-gated dataset and simulator, not run here), and the deployed v5.0 weights were trained by TIER IV on data that is not public, so no public benchmark is in-domain. Details: [`VERIFICATION_2026-10-08.md`](VERIFICATION_2026-10-08.md), [`
|
| 154 |
|
| 155 |
No GPU comparison: no GPU was available on the host where this port was built and measured, so this card makes no GPU speed claim. The reference rows are the port's own fp32 CPU reference on the same host (a correctness baseline, not a speed target). Autoware's CHANGELOG quotes 5.13 ms mean (300 runs) for an older single-step engine of this planner on an RTX PRO 6000 Blackwell with TensorRT (precision not stated); it is not like-for-like with this v5.0 multi-step port. p150 power was not measured, so no efficiency comparison is made.
|
| 156 |
|
| 157 |
## Caveats
|
| 158 |
|
| 159 |
-
-
|
|
|
|
| 160 |
- Deployment status in Autoware: `autoware_diffusion_planner` is an alternative to the default rule-based planning stack, selected with `planning_setting:=diffusion_planner` (package README); it is aimed at Autoware's proposed new planning framework. This bundle is not a ROS 2 node (Python API and HTTP) and not a certified Autoware component; do not use it for safety-critical driving decisions or closed-loop vehicle control.
|
| 161 |
- Stateless API: the node's state between plans (the turn-indicator hold window, the RTC prefix and temperature of the initial solver state, the agent buffers and the ego history) is the client's ("What the caller keeps" in [`code/PYTHON.md`](code/PYTHON.md)). The node's guidance services (start / stop / centerline guidance) are off, as in the node's default.
|
| 162 |
-
- Precision policy of this release: fp32 residual streams and solver state; HiFi4 with fp32 accumulation for every matmul; the ego / neighbour pre-projection as a pad-relative fp32 island; split hi / lo matmuls (bf16 hi + fp32 lo parts, ~1e-5 relative, because a device fp32 matmul rounds its operands like TF32) for the mixer inputs and every decoder linear; an fp32 LayerNorm
|
| 163 |
- Validation scope: the p150 output agrees with the fp32 CPU reference on 99 scenes (above). The two shipped samples and the five research scenes are synthetic-but-faithful scenes built with a Python port of the node's tensor construction; the 92 nuScenes instants are converted from nuScenes v1.0-mini, a domain the model never saw (Singapore and Boston, right-hand traffic in Boston, oracle tracks from 2 Hz annotations, no traffic-light states, no speed limits, stop areas instead of stop lines). On them the plans are plausible but conservative (moving plans about 17 % shorter than the logged drive); the open-loop numbers above are a sanity check against one human driver, not a planning metric.
|
| 164 |
-
- Neighbour predictions: the gate is the median over agents of each agent's max displacement, so single agents can differ more. On the shipped sample `kashiwanoha_dense` (
|
| 165 |
- Documented deviations from the node: the output stays in `base_link` (the node transforms it to `map` with the ego pose); `yaw` is what `tf2::getYaw` reads from the node's quaternion of the unnormalised cos / sin rotation (it differs from atan2(sin, cos) when |(cos, sin)| ≠ 1, as in the node); the speed masks follow the node's TensorRT path (`> FLT_EPSILON`); the turn indicator is decided without the hold window; `delay` is accepted and ignored (the node's multi-step mode never reads it).
|
| 166 |
- Does not scale to multiple p150 in a mesh configuration. The build uses a 12×10 compute grid of Tensix cores: the dispatch functions move from one Tensix column to the ETH cores (`patches/tt-metal-eth-dispatch.patch`), so this build assumes that you do not need chip-to-chip ethernet communication.
|
| 167 |
-
- `dispatch="worker"` (server: `DIFFUSION_PLANNER_DISPATCH=worker`) is an A/B opt-in. On a p150 it gives an 11×10 grid (
|
| 168 |
-
- Batch 1, one plan per request; requests are serialised on the chip. The shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history and 80 future steps) and 10 DPM-Solver steps are compiled into the
|
| 169 |
- Not an OpenAI-compatible API; `GET /v1/models` is a stub so the tt-model ready card does not 404.
|
| 170 |
- p150 power was not measured, so no efficiency comparison is made.
|
| 171 |
|
|
@@ -173,7 +175,7 @@ No GPU comparison: no GPU was available on the host where this port was built an
|
|
| 173 |
|
| 174 |
- Weights: [AutowareFoundation/diffusion_planner](https://huggingface.co/AutowareFoundation/diffusion_planner) at tag `v5.0` (commit `423efde67f5414734da43a7ad856c17ceb8b51aa`), Apache-2.0 per its model card. Not redistributed here: the package only points to them. The upstream card states that TIER IV trained the models on TIER IV synthetic and real driving data; the dataset composition is not publicly documented.
|
| 175 |
- Pre- and post-processing ported from autoware_universe [`planning/autoware_diffusion_planner`](https://github.com/autowarefoundation/autoware_universe/tree/9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd/planning/autoware_diffusion_planner) (Apache-2.0).
|
| 176 |
-
- Port and serving code (`code/`): Apache-2.0. `patches/tt-metal-eth-dispatch.patch`
|
| 177 |
- Sample data: `code/tt_diffusion_planner/samples/kashiwanoha_dense.npz` is derived from the Apache-2.0 Lanelet2 map [AutowareFoundation/map-carla-kashiwanoha](https://huggingface.co/datasets/AutowareFoundation/map-carla-kashiwanoha) (0.2.0) with a scripted ego, route and agents; `straight_road.npz` is a procedural scene generated by this repo. Both Apache-2.0, each with its stored CPU-reference output (`*.reference.json`). Only these redistributable samples ship; the nuScenes-derived planning instants of the accuracy tables are not in this repository.
|
| 178 |
- Demo media (`media/`, sources and changes in [`media/ATTRIBUTION.md`](media/ATTRIBUTION.md)):
|
| 179 |
- nuScenes renders (`media/*_NC.*`), **non-commercial, CC BY-NC-SA 4.0**: Rendered from the nuScenes dataset (v1.0-mini, CAN bus expansion and map expansion v1.3), © Motional AD Inc., CC BY-NC-SA 4.0 and the nuScenes Terms of Use (https://www.nuscenes.org/terms-of-use). Non-commercial use only; adaptations under the same license. Motional does not endorse this work. Cite: H. Caesar et al., *nuScenes: A Multimodal Dataset for Autonomous Driving*, CVPR 2020.
|
|
@@ -185,11 +187,11 @@ These are the exact sources the container image was built from:
|
|
| 185 |
|
| 186 |
| component | built from |
|
| 187 |
| --- | --- |
|
| 188 |
-
| tt-metal | [`44d66500520fda9f2c7060c0f6b41ec48f7ab37e`](https://github.com/tenstorrent/tt-metal/commit/44d66500520fda9f2c7060c0f6b41ec48f7ab37e) + [`patches/tt-metal-eth-dispatch.patch`](patches/tt-metal-eth-dispatch.patch) (sha256 `08d0ddf6…45cc`, 4 files
|
| 189 |
| weights | [`AutowareFoundation/diffusion_planner@423efde67f5414734da43a7ad856c17ceb8b51aa`](https://huggingface.co/AutowareFoundation/diffusion_planner/tree/423efde67f5414734da43a7ad856c17ceb8b51aa) (tag `v5.0`), files `diffusion_planner_encoder.onnx, diffusion_planner_decoder.onnx, diffusion_planner_turn_indicator.onnx, diffusion_planner.param.json` (sha256 `2856886a…ca49`, `eb30c0c0…57ca`, `07acfb58…a732`, `ee3145b6…a268`, checked at load) |
|
| 190 |
| Autoware reference | autoware_universe [`9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd`](https://github.com/autowarefoundation/autoware_universe/tree/9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd) (`planning/autoware_diffusion_planner`, package 0.53.0, `multi_step` mode) |
|
| 191 |
-
| shared package | `ttaw` 0.
|
| 192 |
-
| `code/` digest (image) | `
|
| 193 |
-
| image | `tt-model/diffusion-planner-p150:
|
| 194 |
| base images | build stage `ghcr.io/tenstorrent/tt-metal/tt-metalium/ubuntu-22.04-dev-amd64:latest` @ `sha256:df9d279c7f85c17c6fad982d196802682d669cca1b7ced9cbaad8181339cd5fc`; runtime stage `docker.io/library/ubuntu:22.04` @ `sha256:5ec03bb3441e8b0bf3b4f9cd4629a1ae763010dc3035bb8da3ae6cf026486401` (tt-model's `FROM` tags float; these are the digests this build resolved, see [`build_info.json`](build_info.json)) |
|
| 195 |
-
| built | 2026-10-
|
|
|
|
| 1 |
---
|
| 2 |
tags:
|
|
|
|
|
|
|
| 3 |
- blackhole
|
|
|
|
|
|
|
|
|
|
| 4 |
- p150
|
|
|
|
|
|
|
|
|
|
| 5 |
- tt-dit-server
|
|
|
|
| 6 |
- tt-model-cache
|
|
|
|
| 7 |
- tt-model-container
|
| 8 |
+
- tenstorrent
|
| 9 |
- ttnn
|
| 10 |
+
- tt-metal
|
| 11 |
+
- tt-nn
|
| 12 |
+
- autoware
|
| 13 |
+
- autonomous-driving
|
| 14 |
+
- planning
|
| 15 |
+
- trajectory-generation
|
| 16 |
+
- motion-prediction
|
| 17 |
+
- diffusion
|
| 18 |
+
- diffusion-planner
|
| 19 |
+
- tt-model-catalog
|
| 20 |
pipeline_tag: robotics
|
| 21 |
license: apache-2.0
|
| 22 |
license_link: https://huggingface.co/AutowareFoundation/diffusion_planner
|
|
|
|
| 26 |
|
| 27 |
# diffusion-planner-p150
|
| 28 |
|
| 29 |
+
Diffusion Planner v5.0 (Autoware diffusion_planner): the network Autoware deploys in `autoware_diffusion_planner`, ported to one Tenstorrent Blackhole p150 with tt-nn. The whole plan runs on the chip as one metal trace (one per agent bucket, picked per request): the scene encoder, the 11 DiT decoder evaluations of the DPM-Solver++(2M) loop with their solver updates, and the turn-indicator head. The Autoware planner tensors in (ego and neighbour histories, lanes, route, polygons, line strings, goal, ego shape, turn-indicator history); an 8 s ego trajectory, the predicted 8 s paths of the neighbours and a turn-indicator command out, with the node's exact pre- and post-processing.
|
| 30 |
Weights: [AutowareFoundation/diffusion_planner `v5.0`](https://huggingface.co/AutowareFoundation/diffusion_planner/tree/423efde67f5414734da43a7ad856c17ceb8b51aa) · Paper: [arXiv:2501.15564](https://arxiv.org/abs/2501.15564) · Autoware package: [autoware_diffusion_planner](https://github.com/autowarefoundation/autoware_universe/tree/9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd/planning/autoware_diffusion_planner) · Training code: [tier4/Diffusion-Planner (TIER IV training fork)](https://github.com/tier4/Diffusion-Planner) · Port: [`code/`](https://huggingface.co/changh95/diffusion-planner-p150/tree/main/code)
|
| 31 |
|
| 32 |
+
Runs on **p150** (mesh `P150`). Configuration: dispatch on the ETH cores, 1 command queue, 12×10 compute grid. Numerics (the default and the `serve.env` pins): fp32 residual streams and solver state, HiFi4 with fp32 accumulation, split hi / lo matmuls for the mixer inputs and every decoder linear (each as one K-concatenated fp32 matmul), fp32 LayerNorm in the mixers and the decoder and fp32 matmul attention in the fusion encoder and the decoder (each as one fused custom kernel): the configuration the end-to-end accuracy gates need (Caveats). Optimized release: 5.1x faster on the device than the first release (102.0 -> 20.0 ms per plan on the shipped sample), same gates. All numbers on this card were measured in this configuration.
|
| 33 |
|
| 34 |
Packaged and published with [tt-model-manager](https://github.com/tenstorrent/tt-model-manager) 0.1.0 (manifest schema 5.1).
|
| 35 |
|
| 36 |
## Quickstart (Python)
|
| 37 |
|
| 38 |
+
Prerequisite: a tt-metal / ttnn environment at tt-metal [`44d66500520`](https://github.com/tenstorrent/tt-metal/commit/44d66500520fda9f2c7060c0f6b41ec48f7ab37e) with [`patches/tt-metal-eth-dispatch.patch`](patches/tt-metal-eth-dispatch.patch) applied. The image also carries [`patches/tt-metal-reshape-rm-sys1419.patch`](patches/tt-metal-reshape-rm-sys1419.patch), a guard against an intermittent hang of row-major reshapes under ETH dispatch on Blackhole; the numbers on this card were measured with both. ttnn is not on PyPI.
|
| 39 |
|
| 40 |
```bash
|
| 41 |
hf download changh95/diffusion-planner-p150 --exclude "image/*" --local-dir diffusion-planner-p150 && cd diffusion-planner-p150
|
|
|
|
| 56 |
print(out.turn_indicator["command_name"], out.predicted_agents.shape)
|
| 57 |
```
|
| 58 |
|
| 59 |
+
- `from_pretrained` downloads the three v5.0 ONNX files and `diffusion_planner.param.json` (58.9 MB) of [`AutowareFoundation/diffusion_planner`](https://huggingface.co/AutowareFoundation/diffusion_planner) at the pinned commit `423efde67f5` (tag `v5.0`) to your HF cache (no token needed), checks their sha256, opens the chip, builds the graph and captures the metal traces (6: the full capacity and 5 agent buckets). The first load compiles the kernels (352 s with an empty JIT cache, firmware and every kernel compiled); later loads take about 13 s (device open, reading and uploading the weights, warm-up and the capture of the 6 traces).
|
| 60 |
+
- The traces are captured during the load, so no call compiles anything (first call 37 ms, later calls 32 ms with an `.npz` path).
|
| 61 |
- The `with` block releases the trace and closes the chip. Without `with`, call `model.close()`.
|
| 62 |
|
| 63 |
| | |
|
|
|
|
| 70 |
- The API gives the same output as the HTTP server `/predict`: both share the decoders, the device trace and the host post-processing (checked on the device by `test_api_equals_server`).
|
| 71 |
- The API is stateless: one call is one plan. What the Autoware node keeps between plans stays with the caller: the turn-indicator hold window (1.0 s; the node's manager ships as `tt_diffusion_planner.host.postprocess.TurnIndicatorManager`), the initial solver state `sampled_trajectories` (zeros = the node's default temperature 0; noise for a temperature > 0; the previous plan for the RTC prefix) and the agent and ego histories. Details: [`code/PYTHON.md`](code/PYTHON.md) "What the caller keeps".
|
| 72 |
- One model uses one chip; calls from several threads are serialised.
|
| 73 |
+
- Each call replays the trace of the smallest agent bucket (32 / 64 / 96 / 128 / 192 rows, else the full 321) that holds the ego and every valid neighbour row (the node fills the rows from the front); the result is exact (bit-identical to the full-capacity trace), only the device time depends on the scene.
|
| 74 |
- Full reference: [`code/PYTHON.md`](code/PYTHON.md). Runnable example: [`examples/quickstart.py`](examples/quickstart.py) (also writes `quickstart_bev.png`, the input tensors and the plan from above).
|
| 75 |
|
| 76 |
## Serving (HTTP)
|
|
|
|
| 96 |
"model": "diffusion-planner-p150",
|
| 97 |
"frame_id": "base_link",
|
| 98 |
"meta": {"predicted_agent_columns": ["x", "y", "yaw", "cos", "sin"], "predicted_agent_rows": [0, 1, 2, "... 85 more"], "force_stop": false, "time_from_start_s": [0.1, 0.2, 0.3, "... 77 more"], "valid_counts": {"ego": 1, "neighbor": 88, "static": 0, "lane": 123, "route": 17, "polygon": 0, "line_string": 60, "goal": 1, "ego_shape": 1, "turn": 1}},
|
| 99 |
+
"timing_ms": {"preprocess": 2.4, "device": 22.8, "postprocess": 1.2, "total": 36.9, "decode": 8.4, "model_call": 26.7},
|
| 100 |
"num_poses": 80,
|
| 101 |
"columns": ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"],
|
| 102 |
+
"trajectory": [[0.3644, 0.01, -0.0169, 0.9964, -0.0169, 3.8033, -0.0392], [0.7626, -0.0021, -0.0412, 0.9968, -0.0412, 3.7994, -0.5496], [1.1517, -0.0281, -0.0674, 0.9942, -0.0672, 3.7444, -0.4683], "... 77 more rows"],
|
| 103 |
+
"turn_indicator": {"command": 1, "command_name": "DISABLE", "keep_selected": true, "held": false, "logits": [-15.960197448730469, -5.119836807250977, -5.23490047454834, -0.7636525630950928, 5.502235412597656], "probabilities": [1.655269699085693e-09, 8.448460721410811e-05, 7.530191942350939e-05, 0.00658634165301919, 0.9931545257568359]},
|
| 104 |
"predicted_agents": {"format": "npz", "key": "predicted_agents", "dtype": "float32", "shape": [88, 80, 5], "data": "<base64 npz>"}
|
| 105 |
}
|
| 106 |
```
|
|
|
|
| 133 |
|
| 134 |
## Demo & Performances
|
| 135 |
|
| 136 |
+
Warm, batch 1, measured 2026-10-11 on the optimized release (bundle `bc966b4`; the packaged code adds only the WORKER-grid fallback `7e003c4`, bit-identical on ETH, and the version string). Latency: the stage bench of [`OPT_REPORT.md`](OPT_REPORT.md) (`code/scripts/bench.py`, 100 iterations per stage) on the shipped sample `kashiwanoha_dense.npz` (88 neighbours, 123 lanes, 17 route lanes, 60 line strings: agent bucket r96). The device time depends on the agent bucket the scene needs (exact compaction, see Caveats): 17.44 ms for ≤ 31 neighbours up to 26.65 ms at the full 320. The served rows come from uvicorn on the host (the app the container runs, with its serve pins) and a loopback client, 50 requests of the shipped sample. The host is shared with other jobs, so the host stages move with its load by a few ms; the device rows repeat to ±0.02 ms. Accuracy: the p150 output against the fp32 CPU reference of the same network on 99 scenes: the 2 shipped samples, 5 research scenes and 92 nuScenes v1.0-mini planning instants (the frozen end-to-end gates), and against the research pipeline's ONNX Runtime outputs as an independent oracle. The first release's numbers are in [`OPT_BASELINE.md`](OPT_BASELINE.md).
|
| 137 |
|
| 138 |
| Metric | Performance |
|
| 139 |
|---|---:|
|
| 140 |
+
| Agreement with the fp32 CPU reference, shipped sample `kashiwanoha_dense` (8 s ego plan, 88 neighbours) | ego max **2.0 cm** / mean 0.8 cm; turn command identical; neighbours: median per-agent max 8.6 cm |
|
| 141 |
+
| Agreement, shipped sample `straight_road` | ego max 2.3 cm / mean 0.9 cm; turn command identical; neighbours 6.0 cm |
|
| 142 |
+
| Agreement, all 99 gated scenes (2 samples, 5 research scenes, 92 nuScenes-mini instants) | worst ego max **0.346 m** / mean **0.154 m** (gates 1.0 / 0.3 m; nuscenes/scene-0103_kf14); ego mean: median 1.3 cm, 95th percentile 6.3 cm; turn command identical **99 / 99**; neighbours: median per-agent max ≤ 0.094 m (gate 1.5 m). First release: 0.313 / 0.143 m (two gated precision changes since, Caveats) |
|
| 143 |
+
| Agreement with an independent oracle: the research pipeline's ONNX Runtime outputs (raw x0), 92 nuScenes instants + the 33-plan scene-0061 sequence | worst ego max 0.346 m / mean 0.154 m, turn command identical 92 / 92; sequence: worst ego max 0.112 m / mean 0.055 m, turn 33 / 33; every instant within the gates |
|
| 144 |
+
| Module PCC vs the fp32 reference (encoder categories, encoding, teacher-forced decoder evaluation; gate 0.999) | ≥ 0.999948 / 0.999976 / ≥ 0.9999995 |
|
| 145 |
| Open-loop vs the nuScenes log, 92 instants (a sanity check against one human driver, not a planning metric) | p150 ADE / FDE at 8 s 4.75 / 11.87 m (CPU reference 4.74 / 11.87; constant velocity 4.99 / 13.06); turn command = logged blinker 88 / 92 (CPU 88 / 92) |
|
| 146 |
+
| Python `model()` call, shipped sample (host pre-processing, H2D, trace, D2H, host post-processing) | **26.7 ms p50** (p99 27.5) · 37.4 plans/s (first release 117.9 ms) |
|
| 147 |
+
| Served `/predict` `timing_ms.total` (uvicorn on the host, the shipped sample) | **37.2 ms median** (min 36.7; of which JSON decode 8.6, `model()` 26.8; first release 123.1 ms) |
|
| 148 |
+
| Served client round trip, loopback (base64 `.npz` request, 0.15 MB) | 41.6 ms median (first release 127.5 ms) |
|
| 149 |
+
| Device trace, one blocking plan (encoder + 11 DiT evaluations + 10 solver updates + turn head) | **20.05 ms** (bucket r96; straight_road r32 17.50, nuScenes scene-0103_kf14 r64 19.30; first release 102.13 ms for every scene) |
|
| 150 |
+
| Back-to-back trace replays, by agent bucket (neighbours) | **20.00 ms per plan** · 50.0 plans/s on the shipped sample (r96, 64-95); r32 (≤ 31) 17.44 · r64 19.25 · r128 21.11 · r192 22.30 · full (> 191) 26.65 ms; first release 102.04 ms for every scene |
|
| 151 |
+
| Host pre-processing · pack · host tensors · H2D · D2H · host post-processing | 2.45 · 0.52 · 1.57 · 0.76 · 0.22 · 1.39 ms |
|
| 152 |
+
| `from_pretrained` load: empty JIT cache / warm cache (6 traces) | 352 s / 12.9 s |
|
| 153 |
|
| 154 |
+
All numbers in this table were measured with dispatch on the ETH cores, 1 command queue and a 12×10 compute grid on one p150, with the pinned numerics of `serve.env`. Accuracy is agreement with the fp32 CPU reference of the same Autoware network (same weights, same pre- and post-processing); no dataset-level accuracy is claimed: the paper's benchmark is nuPlan closed loop (an account-gated dataset and simulator, not run here), and the deployed v5.0 weights were trained by TIER IV on data that is not public, so no public benchmark is in-domain. The r128, r192 and full-capacity rows are the same-session values of [`OPT_REPORT.md`](OPT_REPORT.md) round 5, reproduced by [`VERIFICATION_OPT_2026-10-11.md`](VERIFICATION_OPT_2026-10-11.md). Details: [`VERIFICATION_OPT_2026-10-11.md`](VERIFICATION_OPT_2026-10-11.md) (the optimization rounds), [`VERIFICATION_2026-10-08.md`](VERIFICATION_2026-10-08.md) (the first release), [`OPT_REPORT.md`](OPT_REPORT.md), [`OPT_BASELINE.md`](OPT_BASELINE.md).
|
| 155 |
|
| 156 |
No GPU comparison: no GPU was available on the host where this port was built and measured, so this card makes no GPU speed claim. The reference rows are the port's own fp32 CPU reference on the same host (a correctness baseline, not a speed target). Autoware's CHANGELOG quotes 5.13 ms mean (300 runs) for an older single-step engine of this planner on an RTX PRO 6000 Blackwell with TensorRT (precision not stated); it is not like-for-like with this v5.0 multi-step port. p150 power was not measured, so no efficiency comparison is made.
|
| 157 |
|
| 158 |
## Caveats
|
| 159 |
|
| 160 |
+
- Optimized release (5 optimization rounds, [`OPT_REPORT.md`](OPT_REPORT.md)): 102.0 -> 20.0 ms per plan on the device for the shipped sample. What changed: explicit multi-core program configs for the encoder and decoder matmuls; custom `generic_op` kernels for the fp32 LayerNorm, the K-concatenated split matmul operands, the whole fp32 attention (scores in L1) and the score scale / mask / softmax; L1-resident intermediates in the decoder; exact compaction by agent buckets (one trace per bucket); vectorised host pre- and post-processing. The plan is still a sequence of ~1,400 device programs (1,418 at r96, kernel-bound: kernel 19.64 ms, op-to-op gaps 0.89 ms); no megakernel ships in this release (a persistent mixer / DiT megakernel is work in progress and not measured yet). What remains and what comes next: [`OPT_REPORT.md`](OPT_REPORT.md).
|
| 161 |
+
- Two of the optimizations change the numerics slightly, with the frozen gates green: the decoder split linears (round 2) and the encoder split linears (round 3) each accumulate their hi / lo products in one fp32 matmul instead of three. The 99-scene worst case moved from 0.313 / 0.143 m (ego max / mean) to 0.346 / 0.154 m, the gates are 1.0 / 0.3 m; scene by scene the error moves both ways. A held-out check on 50 further scenes (the 33-plan scene-0061 sequence, full-capacity and noisy-x_T scenes, every agent-bucket boundary) passed every gate ([`VERIFICATION_OPT_2026-10-11.md`](VERIFICATION_OPT_2026-10-11.md)); its worst ego max is 0.507 m, on a synthetic scene with 192 neighbours. Every other step is bit-identical to its predecessor.
|
| 162 |
- Deployment status in Autoware: `autoware_diffusion_planner` is an alternative to the default rule-based planning stack, selected with `planning_setting:=diffusion_planner` (package README); it is aimed at Autoware's proposed new planning framework. This bundle is not a ROS 2 node (Python API and HTTP) and not a certified Autoware component; do not use it for safety-critical driving decisions or closed-loop vehicle control.
|
| 163 |
- Stateless API: the node's state between plans (the turn-indicator hold window, the RTC prefix and temperature of the initial solver state, the agent buffers and the ego history) is the client's ("What the caller keeps" in [`code/PYTHON.md`](code/PYTHON.md)). The node's guidance services (start / stop / centerline guidance) are off, as in the node's default.
|
| 164 |
+
- Precision policy of this release: fp32 residual streams and solver state; HiFi4 with fp32 accumulation for every matmul; the ego / neighbour pre-projection as a pad-relative fp32 island; split hi / lo matmuls (bf16 hi + fp32 lo parts, ~1e-5 relative, because a device fp32 matmul rounds its operands like TF32; computed as one fp32 matmul over the K-concatenated operand `[x_hi | x_hi | x_lo | 1]`) for the mixer inputs and every decoder linear; an fp32 LayerNorm (a fused custom kernel with fp32 statistics) in the mixers and the decoder (the fused `ttnn.layer_norm` loses the per-entity signal on the mixers' offset-dominated rows); fp32 matmul attention in the fusion encoder and the decoder (a fused kernel with the scores in L1); the turn head in fp32; the other weights and the hidden MLP activations of the mixers and the fusion encoder in bf16. The plan's sensitivity to these choices is chaotic per scene: with the first device round's numerics, one nuScenes instant (scene-0103_kf14) moved 0.35 m on average while the module PCCs differed only in the 5th decimal. So every numerics change is re-checked on all 99 scenes.
|
| 165 |
- Validation scope: the p150 output agrees with the fp32 CPU reference on 99 scenes (above). The two shipped samples and the five research scenes are synthetic-but-faithful scenes built with a Python port of the node's tensor construction; the 92 nuScenes instants are converted from nuScenes v1.0-mini, a domain the model never saw (Singapore and Boston, right-hand traffic in Boston, oracle tracks from 2 Hz annotations, no traffic-light states, no speed limits, stop areas instead of stop lines). On them the plans are plausible but conservative (moving plans about 17 % shorter than the logged drive); the open-loop numbers above are a sanity check against one human driver, not a planning metric.
|
| 166 |
+
- Neighbour predictions: the gate is the median over agents of each agent's max displacement, so single agents can differ more. On the shipped sample `kashiwanoha_dense` (p150 output vs the stored CPU reference, 2026-10-11) the per-agent max displacement is 8.6 cm median and 0.23 m at the 90th percentile; the worst of the 88 agents differs by 3.72 m. The ego plan is gated on its max and mean; neighbour paths only on that median.
|
| 167 |
- Documented deviations from the node: the output stays in `base_link` (the node transforms it to `map` with the ego pose); `yaw` is what `tf2::getYaw` reads from the node's quaternion of the unnormalised cos / sin rotation (it differs from atan2(sin, cos) when |(cos, sin)| ≠ 1, as in the node); the speed masks follow the node's TensorRT path (`> FLT_EPSILON`); the turn indicator is decided without the hold window; `delay` is accepted and ignored (the node's multi-step mode never reads it).
|
| 168 |
- Does not scale to multiple p150 in a mesh configuration. The build uses a 12×10 compute grid of Tensix cores: the dispatch functions move from one Tensix column to the ETH cores (`patches/tt-metal-eth-dispatch.patch`), so this build assumes that you do not need chip-to-chip ethernet communication.
|
| 169 |
+
- `dispatch="worker"` (server: `DIFFUSION_PLANNER_DISPATCH=worker`) is an A/B opt-in. On a p150 it gives an 11×10 grid (20.61 ms per replay on the shipped sample, 3 % slower than ETH; the 99 gated scenes give the same agreement numbers as ETH); if ETH dispatch is not available (tt-metal without the patch), the model falls back to it with a warning. The numbers on this card do not apply to that mode.
|
| 170 |
+
- Batch 1, one plan per request; requests are serialised on the chip. The shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history and 80 future steps) and 10 DPM-Solver steps are compiled into the traces. The neighbour trunk of the encoder and the decoder run on the smallest agent bucket that holds the scene (one trace per bucket, exact: the outputs equal the full-capacity trace bit for bit), so the device time depends on the neighbour count (17.4-26.7 ms); the map entities (lanes, route, polygons, line strings) are still computed at full capacity in every plan.
|
| 171 |
- Not an OpenAI-compatible API; `GET /v1/models` is a stub so the tt-model ready card does not 404.
|
| 172 |
- p150 power was not measured, so no efficiency comparison is made.
|
| 173 |
|
|
|
|
| 175 |
|
| 176 |
- Weights: [AutowareFoundation/diffusion_planner](https://huggingface.co/AutowareFoundation/diffusion_planner) at tag `v5.0` (commit `423efde67f5414734da43a7ad856c17ceb8b51aa`), Apache-2.0 per its model card. Not redistributed here: the package only points to them. The upstream card states that TIER IV trained the models on TIER IV synthetic and real driving data; the dataset composition is not publicly documented.
|
| 177 |
- Pre- and post-processing ported from autoware_universe [`planning/autoware_diffusion_planner`](https://github.com/autowarefoundation/autoware_universe/tree/9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd/planning/autoware_diffusion_planner) (Apache-2.0).
|
| 178 |
+
- Port and serving code (`code/`): Apache-2.0. `patches/tt-metal-eth-dispatch.patch` and `patches/tt-metal-reshape-rm-sys1419.patch` modify tt-metal (Apache-2.0).
|
| 179 |
- Sample data: `code/tt_diffusion_planner/samples/kashiwanoha_dense.npz` is derived from the Apache-2.0 Lanelet2 map [AutowareFoundation/map-carla-kashiwanoha](https://huggingface.co/datasets/AutowareFoundation/map-carla-kashiwanoha) (0.2.0) with a scripted ego, route and agents; `straight_road.npz` is a procedural scene generated by this repo. Both Apache-2.0, each with its stored CPU-reference output (`*.reference.json`). Only these redistributable samples ship; the nuScenes-derived planning instants of the accuracy tables are not in this repository.
|
| 180 |
- Demo media (`media/`, sources and changes in [`media/ATTRIBUTION.md`](media/ATTRIBUTION.md)):
|
| 181 |
- nuScenes renders (`media/*_NC.*`), **non-commercial, CC BY-NC-SA 4.0**: Rendered from the nuScenes dataset (v1.0-mini, CAN bus expansion and map expansion v1.3), © Motional AD Inc., CC BY-NC-SA 4.0 and the nuScenes Terms of Use (https://www.nuscenes.org/terms-of-use). Non-commercial use only; adaptations under the same license. Motional does not endorse this work. Cite: H. Caesar et al., *nuScenes: A Multimodal Dataset for Autonomous Driving*, CVPR 2020.
|
|
|
|
| 187 |
|
| 188 |
| component | built from |
|
| 189 |
| --- | --- |
|
| 190 |
+
| tt-metal | [`44d66500520fda9f2c7060c0f6b41ec48f7ab37e`](https://github.com/tenstorrent/tt-metal/commit/44d66500520fda9f2c7060c0f6b41ec48f7ab37e) + [`patches/tt-metal-eth-dispatch.patch`](patches/tt-metal-eth-dispatch.patch) (sha256 `08d0ddf6…45cc`, 4 files) + [`patches/tt-metal-reshape-rm-sys1419.patch`](patches/tt-metal-reshape-rm-sys1419.patch) (sha256 `6be07ce5…9cfa`, 1 file) (dirty tree: the image includes both patches; the image build asserts both markers) |
|
| 191 |
| weights | [`AutowareFoundation/diffusion_planner@423efde67f5414734da43a7ad856c17ceb8b51aa`](https://huggingface.co/AutowareFoundation/diffusion_planner/tree/423efde67f5414734da43a7ad856c17ceb8b51aa) (tag `v5.0`), files `diffusion_planner_encoder.onnx, diffusion_planner_decoder.onnx, diffusion_planner_turn_indicator.onnx, diffusion_planner.param.json` (sha256 `2856886a…ca49`, `eb30c0c0…57ca`, `07acfb58…a732`, `ee3145b6…a268`, checked at load) |
|
| 192 |
| Autoware reference | autoware_universe [`9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd`](https://github.com/autowarefoundation/autoware_universe/tree/9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd) (`planning/autoware_diffusion_planner`, package 0.53.0, `multi_step` mode) |
|
| 193 |
+
| shared package | `ttaw` 0.23.2, vendored as `code/tt_diffusion_planner/ttaw` from the Autoware ports' shared `common` repository at commit `60b6dd7` (`code/tt_diffusion_planner/ttaw/VENDORED.json`: version, commit and per-file sha256) |
|
| 194 |
+
| `code/` digest (image) | `5da22b97bbf89b03` (sha256, first 16 hex digits; `built.code_sha256` of `tt_kernel_manifest.json`) |
|
| 195 |
+
| image | `tt-model/diffusion-planner-p150:dde78ac2f0be` (`sha256:dde78ac2f0be5b7e637ddceba1a7c30fd832c2a50dd3e728acbf187f86354cf4`) |
|
| 196 |
| base images | build stage `ghcr.io/tenstorrent/tt-metal/tt-metalium/ubuntu-22.04-dev-amd64:latest` @ `sha256:df9d279c7f85c17c6fad982d196802682d669cca1b7ced9cbaad8181339cd5fc`; runtime stage `docker.io/library/ubuntu:22.04` @ `sha256:5ec03bb3441e8b0bf3b4f9cd4629a1ae763010dc3035bb8da3ae6cf026486401` (tt-model's `FROM` tags float; these are the digests this build resolved, see [`build_info.json`](build_info.json)) |
|
| 197 |
+
| built | 2026-10-11T03:53:34+00:00 by tt-model 0.1.0 |
|
SERVING.md
CHANGED
|
@@ -9,13 +9,13 @@ image. The device open, the decoders, the Python-API contract and the HTTP contr
|
|
| 9 |
|
| 10 |
| item | value |
|
| 11 |
|---|---|
|
| 12 |
-
| tt-metal tree | `/home/ubuntu/experiments/tt-models/tt-metal` (main `44d66500520`, `v0.80.0-dev20261006-78-g44d6650052`, + `patches/tt-metal-eth-dispatch.patch`; torch 2.11.0+cpu) |
|
| 13 |
| weights | `AutowareFoundation/diffusion_planner` @ `423efde67f5414734da43a7ad856c17ceb8b51aa` (tag `v5.0`; files `diffusion_planner_encoder.onnx`, `diffusion_planner_decoder.onnx`, `diffusion_planner_turn_indicator.onnx`, `diffusion_planner.param.json`; 58.9 MB; Apache-2.0; public, ungated; sha256 of every file checked at load) |
|
| 14 |
| Autoware reference | `planning/autoware_diffusion_planner` @ autoware_universe `9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd` (package 0.53.0), the node's `multi_step` mode |
|
| 15 |
-
| shared package | `ttaw` 0.
|
| 16 |
| app | `tt_diffusion_planner.server.app:app` (uvicorn; the ASGI lifespan does weights -> device -> graph -> trace capture) |
|
| 17 |
-
| device recipe | `ttnn.open_device(device_id, dispatch_core_config=DispatchCoreConfig(ETH), l1_small_size=32768, trace_region_size=192 MiB, num_command_queues=1)` (`DEVICE_DEFAULTS` in `code/tt_diffusion_planner/device.py`; the
|
| 18 |
-
| graph | ONE metal trace per plan: the encoder (6 MLP-Mixer trunks, small encoders, 6 fusion blocks), the cross-attention K / V hoisted once, 11 DiT evaluations with the 10 fp32 DPM-Solver++(2M) updates and the prefix constraint, the turn-indicator head, one packed readback;
|
| 19 |
| hardware | one Blackhole p150 (`hardware: p150`, `mesh_device: P150`, `TT_MESH_SHAPE=1x1`), 12×10 compute grid |
|
| 20 |
|
| 21 |
## 1. Run on the HOST (hardware validation, no Docker)
|
|
@@ -44,10 +44,10 @@ $ROOT/bin/devrun -t 1800 -- bash -c 'env DIFFUSION_PLANNER_DISPATCH=eth DIFFUSIO
|
|
| 44 |
```
|
| 45 |
|
| 46 |
Boot log landmarks (they drive `tt-model serve`'s progress view, `boot_progress.py` TT_DIT_PHASES):
|
| 47 |
-
`Loading weights` -> `Opening device` -> `Warming up: capturing trace ...` -> `Warmup complete` -> uvicorn
|
| 48 |
-
`Application startup complete`. A cold JIT cache compiles every kernel first (
|
| 49 |
-
`from_pretrained`); later boots take about
|
| 50 |
-
fallback). SIGTERM: the lifespan releases the
|
| 51 |
|
| 52 |
Offline override (no Hub): `DIFFUSION_PLANNER_WEIGHTS_DIR=<dir holding diffusion_planner_encoder.onnx, diffusion_planner_decoder.onnx, diffusion_planner_turn_indicator.onnx, diffusion_planner.param.json>`.
|
| 53 |
|
|
@@ -144,11 +144,11 @@ the predicted-agents array to its header):
|
|
| 144 |
"model": "diffusion-planner-p150",
|
| 145 |
"frame_id": "base_link",
|
| 146 |
"meta": {"predicted_agent_columns": ["x", "y", "yaw", "cos", "sin"], "predicted_agent_rows": [0, 1, 2, "... 85 more"], "force_stop": false, "time_from_start_s": [0.1, 0.2, 0.3, "... 77 more"], "valid_counts": {"ego": 1, "neighbor": 88, "static": 0, "lane": 123, "route": 17, "polygon": 0, "line_string": 60, "goal": 1, "ego_shape": 1, "turn": 1}},
|
| 147 |
-
"timing_ms": {"preprocess":
|
| 148 |
"num_poses": 80,
|
| 149 |
"columns": ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"],
|
| 150 |
-
"trajectory": [[0.
|
| 151 |
-
"turn_indicator": {"command": 1, "command_name": "DISABLE", "keep_selected": true, "held": false, "logits": [-15.
|
| 152 |
"predicted_agents": {"format": "npz", "key": "predicted_agents", "dtype": "float32", "shape": [88, 80, 5], "data": "<base64 npz>"}
|
| 153 |
}
|
| 154 |
```
|
|
@@ -168,7 +168,7 @@ the predicted-agents array to its header):
|
|
| 168 |
`~/debug/denoising_steps`).
|
| 169 |
|
| 170 |
`timing_ms`: `decode` (base64 + npz parsing + the schema check), `preprocess` (the node's normalization and the encoder's
|
| 171 |
-
host features), `device` (host tensors + H2D + one trace replay + D2H), `postprocess`, `model_call`, `total` (server
|
| 172 |
side, after the body arrived).
|
| 173 |
|
| 174 |
### 3.3 Errors
|
|
@@ -187,9 +187,10 @@ when the boot failed), **500** `inference failed: <ExceptionType>: <message>` if
|
|
| 187 |
| `TT_MESH_SHAPE` | launcher (`runtime.mesh_shape_env`) | `1x1`; anything else -> RuntimeError at startup |
|
| 188 |
| `TT_DEVICE_ID` | you | chip to open, default 0 |
|
| 189 |
| `DIFFUSION_PLANNER_DISPATCH` | `serve.env` | `eth` (default) \| `worker` (A/B only, 11×10) \| `auto` (ETH if the patch is present) |
|
| 190 |
-
| `DIFFUSION_PLANNER_NUM_CQS` | `serve.env` | `1` (2 CQs
|
| 191 |
| `DIFFUSION_PLANNER_VARIANT` | `serve.env` | `default` (the only v5.0 graph) |
|
| 192 |
| `DIFFUSION_PLANNER_LN_FP32`, `_SPLIT_MATMUL`, `_ATTN_MATMUL`, `_HIDDEN_FP32`, `_ATTN_FP32_ACC` | `serve.env` | the numerics knobs, pinned at their gated defaults (`enc.mixer.*,dec.*`; `enc.island.*,enc.pre.*,dec.*`; `enc.fusion.attn,dec.*`; empty; empty: `tt/config.py` `KNOBS`); changing one changes the numerics and needs the gates re-run |
|
|
|
|
| 193 |
| `DIFFUSION_PLANNER_PRECISION` | `serve.env` | empty: the precision policy `DEFAULT_PRECISION` of `tt/config.py` (HiFi4 + fp32 accumulation everywhere); extra rules such as `dec.*=HiFi2+fp32` are for experiments |
|
| 194 |
| `DIFFUSION_PLANNER_WARMUP` | you | JSON list of warm-up variants, `default` or `none` |
|
| 195 |
| `DIFFUSION_PLANNER_TRACE_REGION`, `DIFFUSION_PLANNER_L1_SMALL`, `DIFFUSION_PLANNER_WORKER_L1_SIZE` | you | device-open overrides (validated values: `DEVICE_DEFAULTS` in `device.py`) |
|
|
@@ -197,8 +198,8 @@ when the boot failed), **500** `inference failed: <ExceptionType>: <message>` if
|
|
| 197 |
| `TT_METAL_VISIBLE_DEVICES`, `MESH_DEVICE`, `HF_HUB_DISABLE_IMPLICIT_TOKEN` | `serve.env` / launcher | `0`, `P150`, `1` (the weights are public: no token is ever sent) |
|
| 198 |
|
| 199 |
Server-side pipeline: JSON -> `tt_diffusion_planner.io` decoders (schema check) -> `DiffusionPlanner.__call__` (the node's
|
| 200 |
-
host pre-processing ported from `planning/autoware_diffusion_planner` -> H2D into
|
| 201 |
-
`execute_trace` of the
|
| 202 |
|
| 203 |
### 3.5 What the client keeps (the API is stateless)
|
| 204 |
|
|
@@ -228,11 +229,12 @@ sequence of plans keeps the same state and puts it into the request or applies i
|
|
| 228 |
|
| 229 |
- Batch 1, one plan per request; concurrent clients queue on the lock.
|
| 230 |
- Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history
|
| 231 |
-
and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the
|
| 232 |
-
the
|
|
|
|
| 233 |
- The node's guidance services (start / stop / centerline guidance) are off (the node's default too); a guided mode
|
| 234 |
would need the host in the solver loop.
|
| 235 |
-
- Warm-up captures the
|
| 236 |
- Weights are pinned by sha in `weights.revision`, exported by the launcher and repeated in `serve.env`;
|
| 237 |
`snapshot_download(..., revision=<sha>)` is a cache hit after `serve`'s pre-download and falls back to
|
| 238 |
`local_files_only=True` if the Hub is unreachable. The sha256 of the four files and `major_version == 5` are checked
|
|
|
|
| 9 |
|
| 10 |
| item | value |
|
| 11 |
|---|---|
|
| 12 |
+
| tt-metal tree | `/home/ubuntu/experiments/tt-models/tt-metal` (main `44d66500520`, `v0.80.0-dev20261006-78-g44d6650052`, + `patches/tt-metal-eth-dispatch.patch` + `patches/tt-metal-reshape-rm-sys1419.patch`; torch 2.11.0+cpu) |
|
| 13 |
| weights | `AutowareFoundation/diffusion_planner` @ `423efde67f5414734da43a7ad856c17ceb8b51aa` (tag `v5.0`; files `diffusion_planner_encoder.onnx`, `diffusion_planner_decoder.onnx`, `diffusion_planner_turn_indicator.onnx`, `diffusion_planner.param.json`; 58.9 MB; Apache-2.0; public, ungated; sha256 of every file checked at load) |
|
| 14 |
| Autoware reference | `planning/autoware_diffusion_planner` @ autoware_universe `9ceaccf026c31ffc5319bc9eeb4bd7bede0af3fd` (package 0.53.0), the node's `multi_step` mode |
|
| 15 |
+
| shared package | `ttaw` 0.23.2 (`common` @ `60b6dd7`) |
|
| 16 |
| app | `tt_diffusion_planner.server.app:app` (uvicorn; the ASGI lifespan does weights -> device -> graph -> trace capture) |
|
| 17 |
+
| device recipe | `ttnn.open_device(device_id, dispatch_core_config=DispatchCoreConfig(ETH), l1_small_size=32768, trace_region_size=192 MiB, num_command_queues=1)` (`DEVICE_DEFAULTS` in `code/tt_diffusion_planner/device.py`; the 6 traces hold 60.3 MB) |
|
| 18 |
+
| graph | ONE metal trace per plan: the encoder (6 MLP-Mixer trunks, small encoders, 6 fusion blocks), the cross-attention K / V hoisted once, 11 DiT evaluations with the 10 fp32 DPM-Solver++(2M) updates and the prefix constraint, the turn-indicator head, one packed readback. Six traces are captured at startup, the full capacity and one per agent bucket (32 / 64 / 96 / 128 / 192 decoder rows); each request replays the smallest that holds its valid neighbour rows (exact). 1,418 device programs at r96 (6,282 in the first release); custom `generic_op` kernels in `code/tt_diffusion_planner/tt/kernels/` |
|
| 19 |
| hardware | one Blackhole p150 (`hardware: p150`, `mesh_device: P150`, `TT_MESH_SHAPE=1x1`), 12×10 compute grid |
|
| 20 |
|
| 21 |
## 1. Run on the HOST (hardware validation, no Docker)
|
|
|
|
| 44 |
```
|
| 45 |
|
| 46 |
Boot log landmarks (they drive `tt-model serve`'s progress view, `boot_progress.py` TT_DIT_PHASES):
|
| 47 |
+
`Loading weights` -> `Opening device` -> `Warming up: capturing trace ...` (once per trace variant) -> `Warmup complete` -> uvicorn
|
| 48 |
+
`Application startup complete`. A cold JIT cache compiles every kernel first (352 s measured for
|
| 49 |
+
`from_pretrained`); later boots take about 13 s to READY (measured on the host, warm JIT cache). The packaged container measured the same on 2026-10-11: `tt-model serve` READY after 6 min 0 s on its first start (server boot 358 s, empty JIT cache) and after 20.9 s on the next one (server boot 19.0 s). Startup failures raise and uvicorn exits non-zero (no CPU
|
| 50 |
+
fallback). SIGTERM: the lifespan releases the traces and closes the chip (within `tt-model stop`'s 120 s budget).
|
| 51 |
|
| 52 |
Offline override (no Hub): `DIFFUSION_PLANNER_WEIGHTS_DIR=<dir holding diffusion_planner_encoder.onnx, diffusion_planner_decoder.onnx, diffusion_planner_turn_indicator.onnx, diffusion_planner.param.json>`.
|
| 53 |
|
|
|
|
| 144 |
"model": "diffusion-planner-p150",
|
| 145 |
"frame_id": "base_link",
|
| 146 |
"meta": {"predicted_agent_columns": ["x", "y", "yaw", "cos", "sin"], "predicted_agent_rows": [0, 1, 2, "... 85 more"], "force_stop": false, "time_from_start_s": [0.1, 0.2, 0.3, "... 77 more"], "valid_counts": {"ego": 1, "neighbor": 88, "static": 0, "lane": 123, "route": 17, "polygon": 0, "line_string": 60, "goal": 1, "ego_shape": 1, "turn": 1}},
|
| 147 |
+
"timing_ms": {"preprocess": 2.4, "device": 22.8, "postprocess": 1.2, "total": 36.9, "decode": 8.4, "model_call": 26.7},
|
| 148 |
"num_poses": 80,
|
| 149 |
"columns": ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"],
|
| 150 |
+
"trajectory": [[0.3644, 0.01, -0.0169, 0.9964, -0.0169, 3.8033, -0.0392], [0.7626, -0.0021, -0.0412, 0.9968, -0.0412, 3.7994, -0.5496], [1.1517, -0.0281, -0.0674, 0.9942, -0.0672, 3.7444, -0.4683], "... 77 more rows"],
|
| 151 |
+
"turn_indicator": {"command": 1, "command_name": "DISABLE", "keep_selected": true, "held": false, "logits": [-15.960197448730469, -5.119836807250977, -5.23490047454834, -0.7636525630950928, 5.502235412597656], "probabilities": [1.655269699085693e-09, 8.448460721410811e-05, 7.530191942350939e-05, 0.00658634165301919, 0.9931545257568359]},
|
| 152 |
"predicted_agents": {"format": "npz", "key": "predicted_agents", "dtype": "float32", "shape": [88, 80, 5], "data": "<base64 npz>"}
|
| 153 |
}
|
| 154 |
```
|
|
|
|
| 168 |
`~/debug/denoising_steps`).
|
| 169 |
|
| 170 |
`timing_ms`: `decode` (base64 + npz parsing + the schema check), `preprocess` (the node's normalization and the encoder's
|
| 171 |
+
host features), `device` (host tensors + H2D + one trace replay of the request's agent bucket + D2H), `postprocess`, `model_call`, `total` (server
|
| 172 |
side, after the body arrived).
|
| 173 |
|
| 174 |
### 3.3 Errors
|
|
|
|
| 187 |
| `TT_MESH_SHAPE` | launcher (`runtime.mesh_shape_env`) | `1x1`; anything else -> RuntimeError at startup |
|
| 188 |
| `TT_DEVICE_ID` | you | chip to open, default 0 |
|
| 189 |
| `DIFFUSION_PLANNER_DISPATCH` | `serve.env` | `eth` (default) \| `worker` (A/B only, 11×10) \| `auto` (ETH if the patch is present) |
|
| 190 |
+
| `DIFFUSION_PLANNER_NUM_CQS` | `serve.env` | `1` (at the first release 2 CQs replayed ~10 % slower on ETH, OPT_BASELINE.md; not re-measured since) |
|
| 191 |
| `DIFFUSION_PLANNER_VARIANT` | `serve.env` | `default` (the only v5.0 graph) |
|
| 192 |
| `DIFFUSION_PLANNER_LN_FP32`, `_SPLIT_MATMUL`, `_ATTN_MATMUL`, `_HIDDEN_FP32`, `_ATTN_FP32_ACC` | `serve.env` | the numerics knobs, pinned at their gated defaults (`enc.mixer.*,dec.*`; `enc.island.*,enc.pre.*,dec.*`; `enc.fusion.attn,dec.*`; empty; empty: `tt/config.py` `KNOBS`); changing one changes the numerics and needs the gates re-run |
|
| 193 |
+
| `DIFFUSION_PLANNER_ENC_CH2D`, `_ATTN_FAST`, `_DEC_MMCFG`, `_LN_KERNEL`, `_LN_RESID`, `_SPLIT_KCAT`, `_ATTN_SMASK`, `_ATTN_SMSM`, `_ATTN_FUSED`, `_KCAT_EMIT`, `_LN_TR`, `_KCAT_L1`, `_ATTN_L1`, `_DEC_L1`, `_ENC_KCAT`, `_KCAT_ACT`, `_LN_SPLIT`, `_LN_SFPU_BCAST`, `_COMPACT`, `_AGENT_BUCKETS`, `_HOST_FAST`, `_INPUT_TRIM` (and the rejected `_KCAT_ACT_ONCE`, `_ENC_L1`, `_FUS_L1`, `_LIN_ACT` = 0) | `serve.env` | the optimization knobs, pinned at the shipped defaults (each one an A/B of `OPT_REPORT.md`; `tt/config.py` `KNOBS`); the gates were run with exactly these values |
|
| 194 |
| `DIFFUSION_PLANNER_PRECISION` | `serve.env` | empty: the precision policy `DEFAULT_PRECISION` of `tt/config.py` (HiFi4 + fp32 accumulation everywhere); extra rules such as `dec.*=HiFi2+fp32` are for experiments |
|
| 195 |
| `DIFFUSION_PLANNER_WARMUP` | you | JSON list of warm-up variants, `default` or `none` |
|
| 196 |
| `DIFFUSION_PLANNER_TRACE_REGION`, `DIFFUSION_PLANNER_L1_SMALL`, `DIFFUSION_PLANNER_WORKER_L1_SIZE` | you | device-open overrides (validated values: `DEVICE_DEFAULTS` in `device.py`) |
|
|
|
|
| 198 |
| `TT_METAL_VISIBLE_DEVICES`, `MESH_DEVICE`, `HF_HUB_DISABLE_IMPLICIT_TOKEN` | `serve.env` / launcher | `0`, `P150`, `1` (the weights are public: no token is ever sent) |
|
| 199 |
|
| 200 |
Server-side pipeline: JSON -> `tt_diffusion_planner.io` decoders (schema check) -> `DiffusionPlanner.__call__` (the node's
|
| 201 |
+
host pre-processing ported from `planning/autoware_diffusion_planner` -> H2D into the persistent device inputs ->
|
| 202 |
+
`execute_trace` of the plan's agent-bucket trace -> one D2H -> the node's host post-processing) under one lock -> `Output.to_dict()`.
|
| 203 |
|
| 204 |
### 3.5 What the client keeps (the API is stateless)
|
| 205 |
|
|
|
|
| 229 |
|
| 230 |
- Batch 1, one plan per request; concurrent clients queue on the lock.
|
| 231 |
- Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history
|
| 232 |
+
and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the traces. The neighbour trunk
|
| 233 |
+
and the decoder run on the smallest agent bucket that holds the scene (exact; one trace per bucket), so the device
|
| 234 |
+
time depends on the neighbour count (17.4-26.7 ms per plan); the map entities are computed at full capacity.
|
| 235 |
- The node's guidance services (start / stop / centerline guidance) are off (the node's default too); a guided mode
|
| 236 |
would need the host in the solver loop.
|
| 237 |
+
- Warm-up captures the traces inside the lifespan, so READY means warm; the first cold boot pays the ttnn JIT.
|
| 238 |
- Weights are pinned by sha in `weights.revision`, exported by the launcher and repeated in `serve.env`;
|
| 239 |
`snapshot_download(..., revision=<sha>)` is a cache hit after `serve`'s pre-download and falls back to
|
| 240 |
`local_files_only=True` if the Hub is unreachable. The sha256 of the four files and `major_version == 5` are checked
|
VERIFICATION_OPT_2026-10-11.md
ADDED
|
@@ -0,0 +1,157 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# diffusion-planner-p150: independent verification of the optimization rounds 1-5, 2026-10-11
|
| 2 |
+
|
| 3 |
+
## Verdict
|
| 4 |
+
|
| 5 |
+
**PASS** (3 low findings, none blocking). Every claim in `OPT_REPORT.md` that I re-measured reproduced on the device
|
| 6 |
+
from a clean checkout of `d9a0436` (the OPT round-5 HEAD), in one session with alternating runs against the published
|
| 7 |
+
baseline `b8114ec`:
|
| 8 |
+
|
| 9 |
+
- **Device suite (gates read-only):** 44 / 44 passed, run twice (`TTAW_GATES_READONLY=1`). All 99 per-scene numbers
|
| 10 |
+
are identical to the round-5 record (`logs/diffusion-planner/opt_r5/e2e_compact2.json`, 0 of 99 differ): worst ego
|
| 11 |
+
max 0.346 m, worst mean 0.154 m (gates 1.0 / 0.3 m), turn command 99 / 99, neighbours ≤ 0.094 m (gate 1.5 m).
|
| 12 |
+
- **No gate loosened:** `test_e2e_device.gates.json`, `test_pcc_device.gates.json`, `test_e2e_device.py` and
|
| 13 |
+
`test_pcc_device.py` are byte-identical between the published commit `b8114ec` and `d9a0436`.
|
| 14 |
+
- **Speed:** the published baseline gave 102.03 ms back-to-back and 113.1 ms e2e. HEAD gives **20.00 ms** b2b and
|
| 15 |
+
26.8-27.2 ms e2e on the bench scene. Both match the claims (20.00 ms and 26.84 ms) to within 0.01 ms on the device
|
| 16 |
+
and within host noise on e2e. The round-end commits r1 to r4 also reproduce to within 0.01 ms.
|
| 17 |
+
- **Held-out (50 scenes outside the 99-scene gate set):** all 50 pass the four gates on HEAD. The set covers the
|
| 18 |
+
33-instant scene-0061 sequence, 4 adversarial scenes and 13 agent-bucket boundary scenes.
|
| 19 |
+
- **Compaction is exact:** HEAD (`COMPACT=2`) and HEAD with `COMPACT=0` give bit-identical public outputs on all 50
|
| 20 |
+
held-out scenes. This includes the r128, r192 and full buckets, which no earlier check had covered.
|
| 21 |
+
- **Soak:** 1,200 `model()` calls cycled over 6 scenes covering the r32, r64, r96, r128, r192 and full plans. There
|
| 22 |
+
were 0 output mismatches against each scene's first output, no hang and no fault. The bench runs also replayed the
|
| 23 |
+
traces back to back more than 1,350 times on HEAD.
|
| 24 |
+
- **Host suite (fake ttnn):** 108 passed, 45 skipped (the 44 device tests and the ONNX Runtime module).
|
| 25 |
+
|
| 26 |
+
## Setup
|
| 27 |
+
|
| 28 |
+
| item | value |
|
| 29 |
+
|---|---|
|
| 30 |
+
| Verified commit | `d9a0436` (OPT round-5 HEAD), checked out as a clean `git worktree` (`dirty=0` printed by every job); the bundle's main working tree carries another agent's uncommitted round-6 work, which was **not** tested |
|
| 31 |
+
| Before | `b8114ec` (the published release, `PUSH_READY`); round-end commits `dd6f249` (r1), `c42c540` (r2), `0b750e3` (r3), `93428bb` (r4) |
|
| 32 |
+
| Configuration | shipped defaults: every `DIFFUSION_PLANNER_*` variable unset (the code defaults equal the `serve.env` pins, `test_serve_env_pins_the_numerics`). ETH dispatch, 1 CQ, 12×10, AICLK 1350 MHz (min 1343) in every run |
|
| 33 |
+
| Device hygiene | every device step ran through `bin/devrun` with explicit test files. `TT_METAL_DISPATCH_TIMEOUT_COMMAND_TO_EXECUTE`, `TT_METAL_OPERATION_TIMEOUT_SECONDS` and `TT_METAL_WATCHER` were unset; pytest used `-o faulthandler_timeout=900`. Device work began only after the orchestrator cleared the CenterPoint FAULT marker of 2026-10-10 20:31 (cleared 00:53 UTC) |
|
| 34 |
+
| Window | 2026-10-11 00:54-01:45 UTC, one session; host load 3.7-5.3 (shared host) |
|
| 35 |
+
| Logs / scripts | `logs/diffusion-planner/verify_opt_r1/` (workspace): `job.sh`, `driver.sh`, `gen_heldout.py`, `dev_heldout.py`, `cmp_heldout.py`, `sum_bench.py`, every log / JSON |
|
| 36 |
+
|
| 37 |
+
## Speed: same-session re-measurement
|
| 38 |
+
|
| 39 |
+
`code/scripts/bench.py --iters 30 --warmup 5 --b2b-iters 50 --b2b-rounds 3` was run on 3 scenes:
|
| 40 |
+
`kashiwanoha_dense` (88 neighbours, bucket r96), `straight_road` (r32) and nuScenes `scene-0103_kf14` (r64). The runs
|
| 41 |
+
alternated base / HEAD, 3 pairs, then r4 / HEAD-with-`COMPACT=0`, 2 pairs. Each value below is the median b2b; the
|
| 42 |
+
spread is the range over the pairs.
|
| 43 |
+
|
| 44 |
+
| build | kashiwanoha b2b | straight_road b2b | scene-0103 b2b | one plan (trace) kashiwanoha | e2e p50 kashiwanoha | claimed in `OPT_REPORT.md` |
|
| 45 |
+
|---|---:|---:|---:|---:|---:|---|
|
| 46 |
+
| published `b8114ec` (3 runs) | 102.03 (±0.00) | 102.03 | 102.03 | 102.09 | 113.11-113.24 | 102.03 / 102.09 / 113.1 (same-session baseline) |
|
| 47 |
+
| r1 `dd6f249` | 70.12 | 70.12 | 70.12 | 70.17 | 80.93 | 70.12 / 70.18 / 81.2 |
|
| 48 |
+
| r2 `c42c540` | 40.89 | 40.89 | 40.89 | 40.95 | 51.77 | 40.89 / 40.95 / 51.98 |
|
| 49 |
+
| r3 `0b750e3` | 30.16 | 30.16 | 30.16 | 30.22 | 41.01 | 30.16 / 30.22 / 40.84 |
|
| 50 |
+
| r4 `93428bb` (2 runs) | 26.62 | 26.62 | 26.62 | 26.68-26.70 | 37.25-37.39 | 26.62 / 26.68 / 37.36 |
|
| 51 |
+
| HEAD `COMPACT=0` (2 runs) | 26.64 | 26.64 | 26.64 | 26.70 | 33.58-33.63 | +0.014 ms over r4 (INPUT_TRIM concat); host items −3.5 ms |
|
| 52 |
+
| **HEAD default (3 runs)** | **20.00** (r96) | **17.44** (r32) | **19.25** (r64) | 20.05-20.06 | **26.82-27.17** | 20.00 / 17.44 / 19.25; 20.06; 26.84 |
|
| 53 |
+
|
| 54 |
+
- Every device number reproduces the report to within ±0.02 ms.
|
| 55 |
+
- Speed-up on the bench scene: **5.10×** b2b (102.03 -> 20.00 ms) and **4.2×** e2e (113.1 -> 26.9 ms).
|
| 56 |
+
- `staged_equals_model` was true in every run: the staged path is bit-identical to `model()`.
|
| 57 |
+
- Load cost: HEAD's warm `from_pretrained` takes 11.3-11.5 s, against 8.3-8.8 s for the base. HEAD captures 6
|
| 58 |
+
traces with 60.3 MB of trace buffers and 480 program-cache entries; the base has 74.6 MB and 309 entries.
|
| 59 |
+
- `timing_ms.device` over the 99 gate scenes (from my suite run): p50 20.41 ms, range 20.19-26.77 ms. The report gives
|
| 60 |
+
p50 20.83 ms. This value includes the upload, so it moves with host load.
|
| 61 |
+
|
| 62 |
+
## Accuracy
|
| 63 |
+
|
| 64 |
+
### Frozen gates (99 scenes + per-module PCC)
|
| 65 |
+
|
| 66 |
+
- 44 passed in 28.3 s (second run; the first run in the same window also passed).
|
| 67 |
+
- The per-scene report (`suite_head.e2e.json`) matches `opt_r5/e2e_compact2.json` on all 99 scenes and all
|
| 68 |
+
fields except device time.
|
| 69 |
+
- The round-3 precision change (`ENC_KCAT`, 0.327 -> 0.346 m worst ego max) is therefore the current state, as
|
| 70 |
+
reported.
|
| 71 |
+
|
| 72 |
+
### Held-out set: 50 scenes outside the 99-scene gate set
|
| 73 |
+
|
| 74 |
+
The CPU references come from the bundle's fp32 `ReferencePlanner` (CPU only, generated by `gen_heldout.py`). The
|
| 75 |
+
device outputs come from the public API, `model(inputs=raw)`, so the bucket is chosen per request. The gates are the
|
| 76 |
+
same as the e2e gates: ego max ≤ 1.0 m, ego mean ≤ 0.3 m, turn command equal, and neighbour median-of-max ≤ 1.5 m.
|
| 77 |
+
|
| 78 |
+
| set | n | HEAD: pass / worst ego max / worst mean / worst nb median | published `b8114ec`: worst ego max / mean / nb | HEAD vs HEAD `COMPACT=0` |
|
| 79 |
+
|---|---:|---|---|---|
|
| 80 |
+
| scene-0061 sequence (2 Hz), not in the gate set | 33 | 33 / 33; 0.112 m; 0.055 m; 0.027 m; turn 33 / 33 | 0.118; 0.049; 0.027 | 33 / 33 bit-identical |
|
| 81 |
+
| adversarial: full capacity (320 neighbours, 140 lanes, 25 route, 10 polygons), ego only, x_T = 0.5 N(0,1), x_T = 1.0 N(0,1) | 4 | 4 / 4; 0.058 m; 0.030 m; 0.101 m | 0.062; 0.031; 0.099 | 4 / 4 bit-identical |
|
| 82 |
+
| bucket boundaries: 31, 32, 63, 64, 95, 96, 127, 128, 191, 192 neighbours (variants r32, r64, r64, r96, r96, r128, r128, r192, r192, full), one far slot (100 -> r128, 319 -> full), 31 neighbours + x_T noise | 13 | 13 / 13; **0.507 m** (n192); 0.174 m; 0.076 m | 0.476 (n191); 0.164; 0.076 | 13 / 13 bit-identical |
|
| 83 |
+
|
| 84 |
+
- The precision changes of r2.3 (`SPLIT_KCAT`) and r3.2 (`ENC_KCAT`) hold on the held-out scenes.
|
| 85 |
+
- Error moves both ways scene by scene, as the report describes. The worst case got worse on dense synthetic scenes:
|
| 86 |
+
n192 went from 0.432 to 0.507 m and n128 from 0.340 to 0.361 m, still at about half the 1.0 m gate.
|
| 87 |
+
- The worst neighbour max, which is not gated, is 6.3 m on HEAD against 5.8 m on the base, both on the dense
|
| 88 |
+
synthetic scenes.
|
| 89 |
+
|
| 90 |
+
### Soak (default path)
|
| 91 |
+
|
| 92 |
+
- 1,200 `model()` calls in 33.9 s, cycled over 6 scenes: ego only (r32), scene-0061_kf14 (r64), temp05_kashi (r96),
|
| 93 |
+
n127 (r128), n191 (r192) and full capacity (plan).
|
| 94 |
+
- Each call's poses, predicted agents and turn command were compared byte for byte with that scene's first call:
|
| 95 |
+
0 mismatches.
|
| 96 |
+
- `timing_ms.device` p50 24.0 ms, max 30.4 ms (full plan).
|
| 97 |
+
- No hang, no fault marker, every devrun job rc = 0.
|
| 98 |
+
|
| 99 |
+
## Audit
|
| 100 |
+
|
| 101 |
+
| check | result |
|
| 102 |
+
|---|---|
|
| 103 |
+
| Gates loosened | none (the gate JSONs and the device test files are unchanged since the publish) |
|
| 104 |
+
| Defaults without an A/B row | none. All 26 knobs added by rounds 1-5 have an `OPT_REPORT.md` row. The four rejected ones (`LIN_ACT`, `KCAT_ACT_ONCE`, `ENC_L1`, `FUS_L1`) default to off, and the `serve.env` pins equal the code defaults (host test) |
|
| 105 |
+
| Custom kernels (P3) | every `generic_op` (fattn, ln32 / ln32s, kcat, smask, smsm) passes the per-core runtime-arg count (`N_RT`) as a compile-time arg; the per-core values are runtime args, the buffer addresses common runtime args |
|
| 106 |
+
| Bit-identity claims | r5.1 / r5.2 compaction re-checked on 50 new scenes (above); r1-r4 b2b numbers reproduce; the 99-scene numbers are identical to round 3, as claimed |
|
| 107 |
+
| Hangs / flakiness | none observed (2 suites, 3 held-out runs, 13 benches, 1,200-call soak) |
|
| 108 |
+
| Precision changes vs held-out | r2.3 and r3.2 pass the held-out set (above); not recorded by the optimizer (finding 2) |
|
| 109 |
+
|
| 110 |
+
## Findings
|
| 111 |
+
|
| 112 |
+
1. **(low, coverage) The device gate set does not exercise the r128, r192 and full-capacity traces.**
|
| 113 |
+
- The 99 gate scenes fall in r32, r64 and r96 only (53 / 44 / 2), and `compact_check.py` covered the same three
|
| 114 |
+
buckets.
|
| 115 |
+
- Three of the six shipped traces therefore had no gate or bit-identity check before this verification.
|
| 116 |
+
- They are correct today: they are bit-identical to `COMPACT=0` and within the gates on my boundary scenes.
|
| 117 |
+
- A future change to a bucket config could still break them without any gate noticing.
|
| 118 |
+
- Recommendation: add boundary scenes for every bucket to the e2e test as an additional case set. This adds
|
| 119 |
+
coverage and does not loosen any threshold. Alternatively, run a `COMPACT=0` vs default bit-identity test over
|
| 120 |
+
scenes that reach every bucket.
|
| 121 |
+
2. **(low, process) No held-out check is recorded for the two kept precision changes (r2.3 `SPLIT_KCAT`, r3.2
|
| 122 |
+
`ENC_KCAT`).**
|
| 123 |
+
- `OPT_PLAN.md` asks for one ("plus a 2-scene held-out check"), and the task rules ask for held-out frames on any
|
| 124 |
+
precision change.
|
| 125 |
+
- My 50-scene held-out run passes, so the changes stand.
|
| 126 |
+
- The worst-case ego error on dense synthetic scenes grew from 0.48 to 0.51 m.
|
| 127 |
+
3. **(info) The bundle's main working tree is dirty.** It holds another agent's uncommitted round-6 work (new
|
| 128 |
+
`fattn_ks_*` kernels and edits to `config.py`, `decoder.py`, `encoder.py` and the ln32s kernels). Anyone who runs
|
| 129 |
+
from that directory runs code that this verification did not cover. Everything here ran on clean worktrees of the
|
| 130 |
+
commits named above.
|
| 131 |
+
|
| 132 |
+
## Reproduce
|
| 133 |
+
|
| 134 |
+
```bash
|
| 135 |
+
S=<scratch>; R=/home/ubuntu/experiments/tt-models
|
| 136 |
+
# clean worktrees under $S/wt/<name>/bundles/diffusion-planner-p150 with research/, assets/, common/ symlinked
|
| 137 |
+
# next to bundles/ (the tests resolve the goldens as PKG.parents[3]/research)
|
| 138 |
+
git -C $R/bundles/diffusion-planner-p150 worktree add --detach $S/wt/head/bundles/diffusion-planner-p150 d9a0436
|
| 139 |
+
tools/research-venv/bin/python gen_heldout.py $S/wt/head/bundles/diffusion-planner-p150/code $S/heldout # CPU
|
| 140 |
+
bash driver.sh # every device step via bin/devrun -t <s> -- bash job.sh <wt> suite|heldout|bench|soak <log> [ENV=..]
|
| 141 |
+
python cmp_heldout.py; python sum_bench.py $S
|
| 142 |
+
```
|
| 143 |
+
|
| 144 |
+
## Addendum: republish prep (2026-10-11, 01:50-02:25 UTC)
|
| 145 |
+
|
| 146 |
+
Re-checked on the code that is packaged for the optimized release, after this verification:
|
| 147 |
+
|
| 148 |
+
| step | commit | result | evidence (`logs/diffusion-planner/republish/`) |
|
| 149 |
+
|---|---|---|---|
|
| 150 |
+
| ttaw re-vendor 0.23.0 -> 0.23.2 (common `60b6dd7`) + the reshape-guard patch shipped | `c91ce91` | host suite 108 passed / 45 skipped; `vendor.py --check` 0 differences | `host_suite_v0232.log` |
|
| 151 |
+
| device suite, gates read-only, alloc tracking, quickstart, stage bench (3 scenes, 100 iterations) | `bc966b4` | 44 passed (twice); 99 per-scene numbers identical to round 5; b2b 20.00 / 17.44 / 19.25 ms, e2e p50 26.73 / 23.13 / 25.65 ms | `job1.out`, `device_suite.log`, `alloc_tracking.log`, `bench_recheck.json` |
|
| 152 |
+
| load times, 132 demo / agreement scenes, served bench | `bc966b4` | 352 s cold / 12.9 s warm; ORT oracle 92 / 92 and 33 / 33 within the gates; served 37.2 ms median, `smoke_test.py` PASS | `job2.out`, `load_*.json`, `gates_*_tt_vs_ort.json`, `served_bench.json` |
|
| 153 |
+
| **new finding (fixed)**: `dispatch="worker"` failed at capture (`TT_FATAL: Num output blocks along x (12) must be smaller than or equal to the number of columns in compute grid (11)`): swept matmul configs assumed the 12×10 ETH grid | `7e003c4` | ETH: 44 passed, 99 numbers identical (configs unchanged on 12×10). WORKER (11×10): `test_e2e_device.py` 2 passed, 99 numbers identical to ETH, b2b 20.61 ms (r96) / 18.01 ms (r32). Host suite 108 passed (+2 new host tests) | `job3.out`, `device_suite_fix.log`, `device_e2e_worker.log`, `bench_worker.log`, `host_suite_fix.log` |
|
| 154 |
+
|
| 155 |
+
The WORKER failure was outside the scope of the audit above (it covered the shipped ETH configuration only) and
|
| 156 |
+
affected only the A/B opt-in and the fallback for a tt-metal without the ETH patch; the published ETH numbers are
|
| 157 |
+
unchanged by the fix.
|
build_info.json
CHANGED
|
@@ -1,13 +1,13 @@
|
|
| 1 |
{
|
| 2 |
"schema": "ttaw-build-info/1",
|
| 3 |
"bundle": "diffusion-planner-p150",
|
| 4 |
-
"recorded_at": "2026-10-
|
| 5 |
"image": {
|
| 6 |
-
"tag": "tt-model/diffusion-planner-p150:
|
| 7 |
-
"digest": "sha256:
|
| 8 |
-
"built_at": "2026-10-
|
| 9 |
-
"code_sha256": "
|
| 10 |
-
"size_bytes":
|
| 11 |
"layers": 23
|
| 12 |
},
|
| 13 |
"base_images": {
|
|
|
|
| 1 |
{
|
| 2 |
"schema": "ttaw-build-info/1",
|
| 3 |
"bundle": "diffusion-planner-p150",
|
| 4 |
+
"recorded_at": "2026-10-11T04:02:55Z",
|
| 5 |
"image": {
|
| 6 |
+
"tag": "tt-model/diffusion-planner-p150:dde78ac2f0be",
|
| 7 |
+
"digest": "sha256:dde78ac2f0be5b7e637ddceba1a7c30fd832c2a50dd3e728acbf187f86354cf4",
|
| 8 |
+
"built_at": "2026-10-11T03:53:34+00:00",
|
| 9 |
+
"code_sha256": "5da22b97bbf89b03f83133089a6a3d01a8862a7a1601437774063d8118d9cae3",
|
| 10 |
+
"size_bytes": 4069615514,
|
| 11 |
"layers": 23
|
| 12 |
},
|
| 13 |
"base_images": {
|
code/PYTHON.md
CHANGED
|
@@ -68,7 +68,7 @@ the three v5.0 ONNX files and `diffusion_planner.param.json`, with an offline fa
|
|
| 68 |
file and the weights' `major_version == 5` are checked), opens the chip (ETH dispatch, 12×10, 1 CQ; the other open
|
| 69 |
parameters are `DEVICE_DEFAULTS` in `tt_diffusion_planner/device.py`, overridable with `DIFFUSION_PLANNER_*`), reads the
|
| 70 |
ONNX initializers as data, uploads the weights and constants (48.1 MB), builds the graph, then compiles and captures
|
| 71 |
-
the metal
|
| 72 |
(`model.info["device"]["fallback"]` names it). Any other keyword argument is a `TypeError`.
|
| 73 |
|
| 74 |
The numerics are not arguments: the published configuration is the default of the `DIFFUSION_PLANNER_LN_FP32`,
|
|
@@ -78,9 +78,9 @@ accuracy figures of the card until the gates are re-run.
|
|
| 78 |
|
| 79 |
## Warm-up
|
| 80 |
|
| 81 |
-
`from_pretrained` returns a warm model: it builds the graph, runs the plan once eagerly (this first run compiles every kernel into the JIT cache), then captures the whole plan as
|
| 82 |
|
| 83 |
-
Measured on the shipped sample (`model.info["warmup_ms"]`, 2026-10-
|
| 84 |
|
| 85 |
## Call: `model(...)`
|
| 86 |
|
|
@@ -143,7 +143,7 @@ plans, the caller keeps the same state (SERVING.md 3.5 has the details):
|
|
| 143 |
|
| 144 |
## Lifetime and information
|
| 145 |
|
| 146 |
-
- `model.close()` releases the
|
| 147 |
idempotent. `with` calls it for you; an unclosed model is closed when Python exits.
|
| 148 |
- `model.info`: weights (repo, tag, revision, path), device (dispatch, grid, CQs, fallback), variant, warm variants,
|
| 149 |
warm-up times, runtime parameter defaults, the input schema, the numerics options and precision policy in effect, the
|
|
@@ -152,28 +152,29 @@ plans, the caller keeps the same state (SERVING.md 3.5 has the details):
|
|
| 152 |
|
| 153 |
## Speed
|
| 154 |
|
| 155 |
-
Warm calls, batch 1, ETH dispatch, 1 CQ, 12×10, the pinned numerics (`code/scripts/bench.py`, 100 iterations
|
| 156 |
|
| 157 |
-
| stage | shipped sample `kashiwanoha_dense` |
|
| 158 |
|---|---:|
|
| 159 |
-
| `.npz` decode + schema check (path inputs only) |
|
| 160 |
-
| host pre-processing (the node's normalization, the encoder's host features, masks, `x_T`) |
|
| 161 |
-
| packing the
|
| 162 |
-
| **device trace, one blocking plan** | **
|
| 163 |
-
| D2H (one packed read
|
| 164 |
-
| host post-processing (trajectory, predicted paths, turn decision) |
|
| 165 |
-
| **`model(inputs=arrays)` end to end** | **
|
| 166 |
-
| `model(inputs=<.npz path>)` |
|
| 167 |
-
| back-to-back replays (device time per plan) |
|
| 168 |
-
|
| 169 |
-
The device time
|
| 170 |
|
| 171 |
## Limits
|
| 172 |
|
| 173 |
- Batch 1 on the chip; one model per process.
|
| 174 |
- Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history
|
| 175 |
-
and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the
|
| 176 |
-
the
|
|
|
|
| 177 |
- The node's guidance services (start / stop / centerline guidance) are not available (the node's default is off).
|
| 178 |
- Accuracy is agreement with the fp32 CPU reference of the same network (README "Demo & Performances"); the planner's
|
| 179 |
driving quality is the weights' (trained by TIER IV on data that is not public).
|
|
|
|
| 68 |
file and the weights' `major_version == 5` are checked), opens the chip (ETH dispatch, 12×10, 1 CQ; the other open
|
| 69 |
parameters are `DEVICE_DEFAULTS` in `tt_diffusion_planner/device.py`, overridable with `DIFFUSION_PLANNER_*`), reads the
|
| 70 |
ONNX initializers as data, uploads the weights and constants (48.1 MB), builds the graph, then compiles and captures
|
| 71 |
+
the metal traces. If ETH dispatch cannot open (tt-metal without the patch), it warns and falls back to WORKER dispatch
|
| 72 |
(`model.info["device"]["fallback"]` names it). Any other keyword argument is a `TypeError`.
|
| 73 |
|
| 74 |
The numerics are not arguments: the published configuration is the default of the `DIFFUSION_PLANNER_LN_FP32`,
|
|
|
|
| 78 |
|
| 79 |
## Warm-up
|
| 80 |
|
| 81 |
+
`from_pretrained` returns a warm model: it builds the graph, runs the plan once eagerly (this first run compiles every kernel into the JIT cache), then captures the whole plan as metal traces (`warmup_variants="default"`: the full-capacity variant `plan` and one variant per agent bucket, `plan_r32`, `plan_r64`, `plan_r96`, `plan_r128`, `plan_r192`; each call replays the smallest that holds the scene, exact) with program-cache misses forbidden, so no later call compiles anything. `model.warmup()` is idempotent; `warmup_variants="none"` defers the capture to `model.warmup()`.
|
| 82 |
|
| 83 |
+
Measured on the shipped sample (`model.info["warmup_ms"]`, 2026-10-11): the load takes 352 s with an empty JIT cache and 12.9 s with a warm one (build 0.76 s: the ONNX initializers read and 118.9 MB of weights and constants uploaded; warm-up and capture of the 6 traces 7.7 s; the rest is the device open). The first call then takes 37 ms and the second 32 ms (an `.npz` path; the stage bench's steady state: 27 ms p50 for decoded arrays, 32 ms for an `.npz` path). The 6 traces hold 60.3 MB of DRAM (`trace_region_size` 192 MiB).
|
| 84 |
|
| 85 |
## Call: `model(...)`
|
| 86 |
|
|
|
|
| 143 |
|
| 144 |
## Lifetime and information
|
| 145 |
|
| 146 |
+
- `model.close()` releases the traces and the persistent device tensors and closes the chip if the model opened it;
|
| 147 |
idempotent. `with` calls it for you; an unclosed model is closed when Python exits.
|
| 148 |
- `model.info`: weights (repo, tag, revision, path), device (dispatch, grid, CQs, fallback), variant, warm variants,
|
| 149 |
warm-up times, runtime parameter defaults, the input schema, the numerics options and precision policy in effect, the
|
|
|
|
| 152 |
|
| 153 |
## Speed
|
| 154 |
|
| 155 |
+
Warm calls, batch 1, ETH dispatch, 1 CQ, 12×10, the pinned numerics (`code/scripts/bench.py`, 100 iterations, 2026-10-11, the optimized release, on a shared host; p50, with p99 in brackets):
|
| 156 |
|
| 157 |
+
| stage | shipped sample `kashiwanoha_dense` (88 neighbours: bucket r96) |
|
| 158 |
|---|---:|
|
| 159 |
+
| `.npz` decode + schema check (path inputs only) | 5.76 (7.33) ms |
|
| 160 |
+
| host pre-processing (the node's normalization, the encoder's host features, masks, `x_T`) | 2.45 (3.16) ms |
|
| 161 |
+
| packing the persistent trace inputs · ttnn host tensors · H2D | 0.52 · 1.57 · 0.76 ms |
|
| 162 |
+
| **device trace, one blocking plan** | **20.05** (20.08) ms |
|
| 163 |
+
| D2H (one packed read) | 0.22 ms |
|
| 164 |
+
| host post-processing (trajectory, predicted paths, turn decision) | 1.39 (1.84) ms |
|
| 165 |
+
| **`model(inputs=arrays)` end to end** | **26.73** (27.48) ms |
|
| 166 |
+
| `model(inputs=<.npz path>)` | 32.27 (33.61) ms |
|
| 167 |
+
| back-to-back replays (device time per plan) | 20.00 ms = 50.0 plans/s |
|
| 168 |
+
|
| 169 |
+
The device time depends on the agent bucket the scene needs (exact compaction: the smallest of 32 / 64 / 96 / 128 / 192 decoder rows that holds the ego and every valid neighbour row, else the full capacity). Back to back: r32 (≤ 31 neighbours, e.g. `straight_road`) 17.44 ms, r64 (a nuScenes instant with 42 neighbours) 19.25 ms, r96 20.00 ms, r128 21.11 ms, r192 22.30 ms, full capacity (> 191) 26.65 ms (the last three: `OPT_REPORT.md` round 5). The map entities are computed at full capacity in every plan. The first release took 102.04 ms for every scene (`OPT_BASELINE.md`). 2 CQs were not re-measured on this release (at the first release they did not help a synchronous request: `OPT_BASELINE.md`). Where the time goes and what comes next: `OPT_REPORT.md`.
|
| 170 |
|
| 171 |
## Limits
|
| 172 |
|
| 173 |
- Batch 1 on the chip; one model per process.
|
| 174 |
- Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history
|
| 175 |
+
and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the traces; the neighbour trunk
|
| 176 |
+
and the decoder run on the smallest agent bucket that holds the scene (exact), so the device time depends on the
|
| 177 |
+
neighbour count (17.4-26.7 ms); the map entities are computed at full capacity.
|
| 178 |
- The node's guidance services (start / stop / centerline guidance) are not available (the node's default is off).
|
| 179 |
- Accuracy is agreement with the fp32 CPU reference of the same network (README "Demo & Performances"); the planner's
|
| 180 |
driving quality is the weights' (trained by TIER IV on data that is not public).
|
code/scripts/bench.py
CHANGED
|
@@ -74,12 +74,12 @@ def counts(arrays: Dict[str, np.ndarray]) -> Dict[str, int]:
|
|
| 74 |
"line_strings": rows(arrays["line_strings"][0], (1, 2))}
|
| 75 |
|
| 76 |
|
| 77 |
-
def unpack(out: Dict[str, np.ndarray]) -> Dict[str, Any]:
|
| 78 |
"""The packed readback -> the raw outputs of ``TtDiffusionPlanner.forward`` (same reshapes)."""
|
| 79 |
from tt_diffusion_planner.reference import config as C
|
| 80 |
from tt_diffusion_planner.tt import config as T
|
| 81 |
|
| 82 |
-
final =
|
| 83 |
steps = out["ego_steps"].reshape(-1, T.STATE_COLS)
|
| 84 |
return {"final_x0": final.reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM).astype(np.float32),
|
| 85 |
"logit": out["logit"].reshape(-1)[:C.TURN_INDICATOR_OUTPUT_DIM].astype(np.float32),
|
|
@@ -101,6 +101,7 @@ def bench_scene(model, scene: Dict[str, Any], a: argparse.Namespace) -> Dict[str
|
|
| 101 |
obs = model.normalization.observation
|
| 102 |
bench = StageBench(f"diffusion-planner {scene['name']}")
|
| 103 |
timing: Dict[str, list] = {}
|
|
|
|
| 104 |
sync = lambda: ttnn.synchronize_device(dev) # noqa: E731
|
| 105 |
for _ in range(a.warmup):
|
| 106 |
ref = model(inputs=arrays)
|
|
@@ -120,7 +121,7 @@ def bench_scene(model, scene: Dict[str, Any], a: argparse.Namespace) -> Dict[str
|
|
| 120 |
with bench.stage("host_pre"):
|
| 121 |
prep = hp.prepare(raw, obs)
|
| 122 |
with bench.stage("pack"):
|
| 123 |
-
packed = I.plan_inputs(prep)
|
| 124 |
with bench.stage("host_in"):
|
| 125 |
host = {k: to_host_tensor(v, slots[k].dtype, slots[k].layout, shape=slots[k].shape)
|
| 126 |
for k, v in packed.items()}
|
|
@@ -128,25 +129,37 @@ def bench_scene(model, scene: Dict[str, Any], a: argparse.Namespace) -> Dict[str
|
|
| 128 |
runner.upload(host)
|
| 129 |
sync()
|
| 130 |
with bench.stage("trace"):
|
| 131 |
-
runner.replay(
|
| 132 |
sync()
|
| 133 |
with bench.stage("d2h"):
|
| 134 |
-
out = runner.read(
|
| 135 |
with bench.stage("host_post"):
|
| 136 |
-
res = model._postprocess(unpack(out), prep, params)
|
| 137 |
if scene["schema_npz"]:
|
| 138 |
for _ in range(a.iters):
|
| 139 |
with bench.stage("e2e_path"):
|
| 140 |
model(inputs=scene["path"])
|
| 141 |
-
rounds = [time_b2b(lambda: runner.replay(
|
| 142 |
for _ in range(a.b2b_rounds)]
|
| 143 |
for r in rounds:
|
| 144 |
bench.add("b2b", r)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 145 |
same = bool(np.array_equal(res.poses, ref.poses) and np.array_equal(res.predicted_agents, ref.predicted_agents)
|
| 146 |
and res.turn_indicator["command"] == ref.turn_indicator["command"])
|
| 147 |
summary = bench.summary()
|
| 148 |
b2b = statistics.median(rounds)
|
| 149 |
-
return {"name": scene["name"], "path": scene["path"], "valid": counts(arrays), "stages_ms": summary,
|
| 150 |
"timing_ms": {k: {"p50": float(np.percentile(v, 50)), "p99": float(np.percentile(v, 99)),
|
| 151 |
"min": float(min(v))} for k, v in timing.items()},
|
| 152 |
"b2b_rounds_ms": rounds, "plans_per_s_b2b": 1000.0 / b2b, "aiclk": clk.summary(),
|
|
@@ -166,6 +179,8 @@ def main() -> None:
|
|
| 166 |
ap.add_argument("--chip", type=int, default=0, help="sysfs chip index for the AICLK sampler")
|
| 167 |
ap.add_argument("--tag", default="")
|
| 168 |
ap.add_argument("--json")
|
|
|
|
|
|
|
| 169 |
a = ap.parse_args()
|
| 170 |
from tt_diffusion_planner.reference.weights import find_weights_dir
|
| 171 |
|
|
@@ -191,7 +206,7 @@ def main() -> None:
|
|
| 191 |
res["scenes"][r["name"]] = r
|
| 192 |
cfg = model.device_info
|
| 193 |
print(f"\n## {r['name']} [{cfg.get('dispatch')} {res['num_command_queues']}CQ {cfg.get('grid')}] "
|
| 194 |
-
f"valid {r['valid']}")
|
| 195 |
print(r["table"])
|
| 196 |
print("timing_ms p50:", {k: round(v["p50"], 3) for k, v in r["timing_ms"].items()},
|
| 197 |
f"| b2b rounds {[round(x, 3) for x in r['b2b_rounds_ms']]} ms -> {r['plans_per_s_b2b']:.2f} plans/s",
|
|
|
|
| 74 |
"line_strings": rows(arrays["line_strings"][0], (1, 2))}
|
| 75 |
|
| 76 |
|
| 77 |
+
def unpack(out: Dict[str, np.ndarray], tt: Any) -> Dict[str, Any]:
|
| 78 |
"""The packed readback -> the raw outputs of ``TtDiffusionPlanner.forward`` (same reshapes)."""
|
| 79 |
from tt_diffusion_planner.reference import config as C
|
| 80 |
from tt_diffusion_planner.tt import config as T
|
| 81 |
|
| 82 |
+
final = tt.unpack(out)
|
| 83 |
steps = out["ego_steps"].reshape(-1, T.STATE_COLS)
|
| 84 |
return {"final_x0": final.reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM).astype(np.float32),
|
| 85 |
"logit": out["logit"].reshape(-1)[:C.TURN_INDICATOR_OUTPUT_DIM].astype(np.float32),
|
|
|
|
| 101 |
obs = model.normalization.observation
|
| 102 |
bench = StageBench(f"diffusion-planner {scene['name']}")
|
| 103 |
timing: Dict[str, list] = {}
|
| 104 |
+
variant = model.tt.variant_for(hp.prepare(load_named_arrays(arrays, C.INPUT_SCHEMA), obs)) # COMPACT bucket
|
| 105 |
sync = lambda: ttnn.synchronize_device(dev) # noqa: E731
|
| 106 |
for _ in range(a.warmup):
|
| 107 |
ref = model(inputs=arrays)
|
|
|
|
| 121 |
with bench.stage("host_pre"):
|
| 122 |
prep = hp.prepare(raw, obs)
|
| 123 |
with bench.stage("pack"):
|
| 124 |
+
packed = model.tt.filter_inputs(I.plan_inputs(prep)) # INPUT_TRIM (as model() does)
|
| 125 |
with bench.stage("host_in"):
|
| 126 |
host = {k: to_host_tensor(v, slots[k].dtype, slots[k].layout, shape=slots[k].shape)
|
| 127 |
for k, v in packed.items()}
|
|
|
|
| 129 |
runner.upload(host)
|
| 130 |
sync()
|
| 131 |
with bench.stage("trace"):
|
| 132 |
+
runner.replay(variant)
|
| 133 |
sync()
|
| 134 |
with bench.stage("d2h"):
|
| 135 |
+
out = runner.read(variant)
|
| 136 |
with bench.stage("host_post"):
|
| 137 |
+
res = model._postprocess(unpack(out, model.tt), prep, params)
|
| 138 |
if scene["schema_npz"]:
|
| 139 |
for _ in range(a.iters):
|
| 140 |
with bench.stage("e2e_path"):
|
| 141 |
model(inputs=scene["path"])
|
| 142 |
+
rounds = [time_b2b(lambda: runner.replay(variant), sync, n=a.b2b_iters, warmup=3)
|
| 143 |
for _ in range(a.b2b_rounds)]
|
| 144 |
for r in rounds:
|
| 145 |
bench.add("b2b", r)
|
| 146 |
+
if getattr(a, "dump", None):
|
| 147 |
+
Path(a.dump).parent.mkdir(parents=True, exist_ok=True)
|
| 148 |
+
flat: Dict[str, np.ndarray] = {}
|
| 149 |
+
|
| 150 |
+
def walk(prefix, v):
|
| 151 |
+
if isinstance(v, dict):
|
| 152 |
+
for k, w in v.items():
|
| 153 |
+
walk(f"{prefix}{k}/", w)
|
| 154 |
+
else:
|
| 155 |
+
flat[prefix.rstrip("/") or "out"] = np.asarray(v)
|
| 156 |
+
walk("", out)
|
| 157 |
+
np.savez(f"{a.dump}.{scene['name'].replace('/', '_')}.npz", **flat)
|
| 158 |
same = bool(np.array_equal(res.poses, ref.poses) and np.array_equal(res.predicted_agents, ref.predicted_agents)
|
| 159 |
and res.turn_indicator["command"] == ref.turn_indicator["command"])
|
| 160 |
summary = bench.summary()
|
| 161 |
b2b = statistics.median(rounds)
|
| 162 |
+
return {"name": scene["name"], "path": scene["path"], "valid": counts(arrays), "variant": variant, "stages_ms": summary,
|
| 163 |
"timing_ms": {k: {"p50": float(np.percentile(v, 50)), "p99": float(np.percentile(v, 99)),
|
| 164 |
"min": float(min(v))} for k, v in timing.items()},
|
| 165 |
"b2b_rounds_ms": rounds, "plans_per_s_b2b": 1000.0 / b2b, "aiclk": clk.summary(),
|
|
|
|
| 179 |
ap.add_argument("--chip", type=int, default=0, help="sysfs chip index for the AICLK sampler")
|
| 180 |
ap.add_argument("--tag", default="")
|
| 181 |
ap.add_argument("--json")
|
| 182 |
+
ap.add_argument("--dump", help="prefix: save each scene's raw packed readback as <prefix>.<scene>.npz "
|
| 183 |
+
"(bit-identity checks of structural rewrites)")
|
| 184 |
a = ap.parse_args()
|
| 185 |
from tt_diffusion_planner.reference.weights import find_weights_dir
|
| 186 |
|
|
|
|
| 206 |
res["scenes"][r["name"]] = r
|
| 207 |
cfg = model.device_info
|
| 208 |
print(f"\n## {r['name']} [{cfg.get('dispatch')} {res['num_command_queues']}CQ {cfg.get('grid')}] "
|
| 209 |
+
f"valid {r['valid']} variant {r['variant']}")
|
| 210 |
print(r["table"])
|
| 211 |
print("timing_ms p50:", {k: round(v["p50"], 3) for k, v in r["timing_ms"].items()},
|
| 212 |
f"| b2b rounds {[round(x, 3) for x in r['b2b_rounds_ms']]} ms -> {r['plans_per_s_b2b']:.2f} plans/s",
|
code/scripts/profile_ops.py
CHANGED
|
@@ -209,7 +209,10 @@ def main() -> None:
|
|
| 209 |
tt, runner, dev = model.tt, model.runner, model.device
|
| 210 |
print("loaded in %.1f s:" % (time.perf_counter() - t0), json.dumps(model.device_info), flush=True)
|
| 211 |
raw = load_named_arrays(a.input, C.INPUT_SCHEMA)
|
| 212 |
-
|
|
|
|
|
|
|
|
|
|
| 213 |
served = model(inputs=raw) # one served plan: upload + replay + readback
|
| 214 |
ttnn.synchronize_device(dev)
|
| 215 |
read_device_profiler(dev) # warm-up / capture / first plan out of the buffer
|
|
@@ -219,7 +222,7 @@ def main() -> None:
|
|
| 219 |
try:
|
| 220 |
t1 = time.perf_counter()
|
| 221 |
with signposted("eager"):
|
| 222 |
-
eager = runner.run_eager(
|
| 223 |
ttnn.synchronize_device(dev)
|
| 224 |
timings["eager_ms"] = (time.perf_counter() - t1) * 1e3
|
| 225 |
finally:
|
|
@@ -232,11 +235,11 @@ def main() -> None:
|
|
| 232 |
read_device_profiler(dev)
|
| 233 |
t1 = time.perf_counter()
|
| 234 |
with signposted("trace"):
|
| 235 |
-
runner.replay(
|
| 236 |
ttnn.synchronize_device(dev)
|
| 237 |
timings["trace_ms"] = (time.perf_counter() - t1) * 1e3 / a.replays
|
| 238 |
read_device_profiler(dev)
|
| 239 |
-
out = runner.read(
|
| 240 |
same = bool(np.array_equal(out["final_x0"], eager["final_x0"])) if a.eager else None
|
| 241 |
desc = {"device": model.device_info, "timings_ms": timings, "replays": a.replays,
|
| 242 |
"replay_equals_eager": same, "turn_command": int(served.turn_indicator["command"]),
|
|
|
|
| 209 |
tt, runner, dev = model.tt, model.runner, model.device
|
| 210 |
print("loaded in %.1f s:" % (time.perf_counter() - t0), json.dumps(model.device_info), flush=True)
|
| 211 |
raw = load_named_arrays(a.input, C.INPUT_SCHEMA)
|
| 212 |
+
prep = hp.prepare(raw, model.normalization.observation)
|
| 213 |
+
inputs = I.plan_inputs(prep)
|
| 214 |
+
variant = tt.variant_for(prep) # COMPACT: the served bucket trace
|
| 215 |
+
print("variant", variant, flush=True)
|
| 216 |
served = model(inputs=raw) # one served plan: upload + replay + readback
|
| 217 |
ttnn.synchronize_device(dev)
|
| 218 |
read_device_profiler(dev) # warm-up / capture / first plan out of the buffer
|
|
|
|
| 222 |
try:
|
| 223 |
t1 = time.perf_counter()
|
| 224 |
with signposted("eager"):
|
| 225 |
+
eager = runner.run_eager(variant, inputs=inputs)
|
| 226 |
ttnn.synchronize_device(dev)
|
| 227 |
timings["eager_ms"] = (time.perf_counter() - t1) * 1e3
|
| 228 |
finally:
|
|
|
|
| 235 |
read_device_profiler(dev)
|
| 236 |
t1 = time.perf_counter()
|
| 237 |
with signposted("trace"):
|
| 238 |
+
runner.replay(variant, n=a.replays)
|
| 239 |
ttnn.synchronize_device(dev)
|
| 240 |
timings["trace_ms"] = (time.perf_counter() - t1) * 1e3 / a.replays
|
| 241 |
read_device_profiler(dev)
|
| 242 |
+
out = runner.read(variant)
|
| 243 |
same = bool(np.array_equal(out["final_x0"], eager["final_x0"])) if a.eager else None
|
| 244 |
desc = {"device": model.device_info, "timings_ms": timings, "replays": a.replays,
|
| 245 |
"replay_equals_eager": same, "turn_command": int(served.turn_indicator["command"]),
|
code/tt_diffusion_planner/__init__.py
CHANGED
|
@@ -16,7 +16,7 @@ never edit it here). ``tt_diffusion_planner.reference`` is the fp32 CPU referenc
|
|
| 16 |
``tt_diffusion_planner.host`` the node's pre- and post-processing (numpy).
|
| 17 |
"""
|
| 18 |
|
| 19 |
-
__version__ = "0.
|
| 20 |
__all__ = ["DiffusionPlanner", "Output", "open_device", "load_inputs", "INPUT_SCHEMA", "__version__"]
|
| 21 |
|
| 22 |
_LAZY = {
|
|
|
|
| 16 |
``tt_diffusion_planner.host`` the node's pre- and post-processing (numpy).
|
| 17 |
"""
|
| 18 |
|
| 19 |
+
__version__ = "0.2.0"
|
| 20 |
__all__ = ["DiffusionPlanner", "Output", "open_device", "load_inputs", "INPUT_SCHEMA", "__version__"]
|
| 21 |
|
| 22 |
_LAZY = {
|
code/tt_diffusion_planner/host/features.py
CHANGED
|
@@ -115,6 +115,18 @@ class DecoderMasks:
|
|
| 115 |
current_states: np.ndarray # [321, 4] normalised (prefix constraint)
|
| 116 |
|
| 117 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 118 |
def _any_nonzero(a: np.ndarray, axes) -> np.ndarray:
|
| 119 |
return np.any(a != 0, axis=axes)
|
| 120 |
|
|
@@ -129,16 +141,31 @@ def encoder_features(norm: Mapping[str, np.ndarray], masks: Mapping[str, np.ndar
|
|
| 129 |
|
| 130 |
# neighbours: keep the 6 newest rows (encoder.py:176-181, 441-451)
|
| 131 |
nb_raw = f32("neighbor_agents_past")
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 142 |
|
| 143 |
static = f32("static_objects")
|
| 144 |
st_valid = _any_nonzero(static[..., :10], -1)
|
|
|
|
| 115 |
current_states: np.ndarray # [321, 4] normalised (prefix constraint)
|
| 116 |
|
| 117 |
|
| 118 |
+
|
| 119 |
+
def _host_fast() -> bool:
|
| 120 |
+
"""``DIFFUSION_PLANNER_HOST_FAST`` (OPT round 5 item 3), read once: the vectorised host functions (bit-exact)
|
| 121 |
+
instead of the first port's ``*_ref`` versions."""
|
| 122 |
+
from ..tt.config import KNOBS
|
| 123 |
+
|
| 124 |
+
return bool(KNOBS.read().HOST_FAST)
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
HOST_FAST = _host_fast()
|
| 128 |
+
|
| 129 |
+
|
| 130 |
def _any_nonzero(a: np.ndarray, axes) -> np.ndarray:
|
| 131 |
return np.any(a != 0, axis=axes)
|
| 132 |
|
|
|
|
| 141 |
|
| 142 |
# neighbours: keep the 6 newest rows (encoder.py:176-181, 441-451)
|
| 143 |
nb_raw = f32("neighbor_agents_past")
|
| 144 |
+
if HOST_FAST:
|
| 145 |
+
K = C.NEIGHBOR_HISTORY_KEEP # the other rows are zero: computed on the kept ones
|
| 146 |
+
x8k = nb_raw[:, K, :8]
|
| 147 |
+
svk = _any_nonzero(x8k, -1) # [320, 6] valid steps
|
| 148 |
+
nb_valid = svk.any(axis=-1) # [320]
|
| 149 |
+
nb_type = nb_raw[:, -1, 8:11].copy() # row 30 is kept
|
| 150 |
+
feat = np.zeros(nb_raw.shape[:2] + (9,), np.float32)
|
| 151 |
+
fk = feat[:, K]
|
| 152 |
+
fk[..., :8] = x8k
|
| 153 |
+
fk[..., 8] = svk
|
| 154 |
+
fk[..., 4:6] = 0.0 # velocities zeroed after the validity test
|
| 155 |
+
fk[~nb_valid] = 0.0
|
| 156 |
+
feat[:, K] = fk
|
| 157 |
+
pos["neighbor"] = _onehot_pos(nb_raw[:, -1, :4], C.POS_CLASS["neighbor"])
|
| 158 |
+
else:
|
| 159 |
+
nb = np.zeros_like(nb_raw)
|
| 160 |
+
nb[:, C.NEIGHBOR_HISTORY_KEEP] = nb_raw[:, C.NEIGHBOR_HISTORY_KEEP]
|
| 161 |
+
nb_type = nb[:, -1, 8:11].copy()
|
| 162 |
+
x8 = nb[..., :8]
|
| 163 |
+
step_valid = _any_nonzero(x8, -1) # [320, 31]
|
| 164 |
+
nb_valid = step_valid.any(axis=-1) # [320]
|
| 165 |
+
feat = np.concatenate([x8, step_valid[..., None].astype(np.float32)], axis=-1)
|
| 166 |
+
feat[..., 4:6] = 0.0 # velocities zeroed after the validity test
|
| 167 |
+
feat[~nb_valid] = 0.0
|
| 168 |
+
pos["neighbor"] = _onehot_pos(x8[:, -1, :4], C.POS_CLASS["neighbor"])
|
| 169 |
|
| 170 |
static = f32("static_objects")
|
| 171 |
st_valid = _any_nonzero(static[..., :10], -1)
|
code/tt_diffusion_planner/host/normalize.py
CHANGED
|
@@ -19,7 +19,18 @@ __all__ = ["FLT_EPSILON", "normalize_inputs", "normalize_array", "speed_masks"]
|
|
| 19 |
FLT_EPSILON = np.float32(np.finfo(np.float32).eps) # std::numeric_limits<float>::epsilon()
|
| 20 |
|
| 21 |
|
| 22 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 23 |
"""``normalize_vector`` of ``preprocessing_utils.cpp:36-63`` on one tensor (a new float32 array).
|
| 24 |
|
| 25 |
Rows have ``std.size`` columns; a single-value mean / std broadcasts over the columns (the C++ ``mean.size() == 1``
|
|
@@ -43,6 +54,57 @@ def normalize_array(value: np.ndarray, mean: np.ndarray, std: np.ndarray) -> np.
|
|
| 43 |
return rows.reshape(data.shape)
|
| 44 |
|
| 45 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
def normalize_inputs(raw: Mapping[str, np.ndarray],
|
| 47 |
observation: Mapping[str, Tuple[np.ndarray, np.ndarray]]) -> Dict[str, np.ndarray]:
|
| 48 |
"""``normalize_input_data(input_data_map, normalization_map)``: every key except the four skipped ones must have
|
|
|
|
| 19 |
FLT_EPSILON = np.float32(np.finfo(np.float32).eps) # std::numeric_limits<float>::epsilon()
|
| 20 |
|
| 21 |
|
| 22 |
+
def _host_fast() -> bool:
|
| 23 |
+
"""``DIFFUSION_PLANNER_HOST_FAST`` (OPT round 5 item 3), read once: the vectorised host functions (bit-exact)
|
| 24 |
+
instead of the first port's ``*_ref`` versions."""
|
| 25 |
+
from ..tt.config import KNOBS
|
| 26 |
+
|
| 27 |
+
return bool(KNOBS.read().HOST_FAST)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
HOST_FAST = _host_fast()
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def normalize_array_ref(value: np.ndarray, mean: np.ndarray, std: np.ndarray) -> np.ndarray:
|
| 34 |
"""``normalize_vector`` of ``preprocessing_utils.cpp:36-63`` on one tensor (a new float32 array).
|
| 35 |
|
| 36 |
Rows have ``std.size`` columns; a single-value mean / std broadcasts over the columns (the C++ ``mean.size() == 1``
|
|
|
|
| 54 |
return rows.reshape(data.shape)
|
| 55 |
|
| 56 |
|
| 57 |
+
_NORMALIZERS: Dict[Tuple[int, int], tuple] = {}
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
def _normalizer(mean: np.ndarray, std: np.ndarray):
|
| 61 |
+
"""``(cols, m, s)`` of a normalizer, validated once per (mean, std) object pair (the weights' constants)."""
|
| 62 |
+
key = (id(mean), id(std))
|
| 63 |
+
hit = _NORMALIZERS.get(key)
|
| 64 |
+
if hit is not None and hit[0] is mean and hit[1] is std:
|
| 65 |
+
return hit[2]
|
| 66 |
+
m32 = np.asarray(mean, np.float32).reshape(-1)
|
| 67 |
+
s32 = np.asarray(std, np.float32).reshape(-1)
|
| 68 |
+
if m32.size != s32.size:
|
| 69 |
+
raise ValueError("Mean and std must be same size")
|
| 70 |
+
cols = s32.size
|
| 71 |
+
if cols and np.any(np.abs(s32) < FLT_EPSILON):
|
| 72 |
+
raise ValueError("Standard deviation is zero, cannot normalize data")
|
| 73 |
+
m = np.ascontiguousarray(np.broadcast_to(m32 if m32.size > 1 else m32[:1], (cols,)))
|
| 74 |
+
s = np.ascontiguousarray(np.broadcast_to(s32 if s32.size > 1 else s32[:1], (cols,)))
|
| 75 |
+
out = (cols, m, s, np.ones(cols, np.uint8) if cols < 256 else None)
|
| 76 |
+
if len(_NORMALIZERS) > 256:
|
| 77 |
+
_NORMALIZERS.clear()
|
| 78 |
+
_NORMALIZERS[key] = (mean, std, out)
|
| 79 |
+
return out
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def normalize_array(value: np.ndarray, mean: np.ndarray, std: np.ndarray) -> np.ndarray:
|
| 83 |
+
"""``normalize_vector`` of ``preprocessing_utils.cpp:36-63`` on one tensor (a new float32 array).
|
| 84 |
+
|
| 85 |
+
Rows have ``std.size`` columns; a single-value mean / std broadcasts over the columns (the C++ ``mean.size() == 1``
|
| 86 |
+
branch); a zero standard deviation raises like the C++. Only the rows that are not all-small are computed
|
| 87 |
+
(``(v - mean) / std`` in float32, element by element); the others are copied unchanged."""
|
| 88 |
+
if not HOST_FAST:
|
| 89 |
+
return normalize_array_ref(value, mean, std)
|
| 90 |
+
data = np.asarray(value, dtype=np.float32) # not modified: the result is a new array
|
| 91 |
+
cols, m, s, ones = _normalizer(mean, std)
|
| 92 |
+
if cols == 0 or data.size % cols:
|
| 93 |
+
raise ValueError(f"data size {data.size} is not divisible by the normalizer size {cols}")
|
| 94 |
+
rows = data.reshape(-1, cols)
|
| 95 |
+
big = ~(np.abs(rows) < FLT_EPSILON) # NaN counts as not small, as in the C++ test
|
| 96 |
+
if ones is not None: # per-row count of the non-small values: a uint8 dot product
|
| 97 |
+
keep = (big.view(np.uint8) @ ones) != 0 # (faster than a short-axis any())
|
| 98 |
+
else:
|
| 99 |
+
keep = big.any(axis=1)
|
| 100 |
+
idx = np.flatnonzero(keep)
|
| 101 |
+
if 2 * idx.size > keep.size: # mostly valid rows: compute all, keep the small rows
|
| 102 |
+
return np.where(keep[:, None], (rows - m) / s, rows).reshape(data.shape)
|
| 103 |
+
out = rows.copy()
|
| 104 |
+
out[idx] = (rows[idx] - m) / s
|
| 105 |
+
return out.reshape(data.shape)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
def normalize_inputs(raw: Mapping[str, np.ndarray],
|
| 109 |
observation: Mapping[str, Tuple[np.ndarray, np.ndarray]]) -> Dict[str, np.ndarray]:
|
| 110 |
"""``normalize_input_data(input_data_map, normalization_map)``: every key except the four skipped ones must have
|
code/tt_diffusion_planner/host/pipeline.py
CHANGED
|
@@ -19,6 +19,7 @@ from ..ttaw.io import InputError, encode_array
|
|
| 19 |
from ..ttaw.outputs import Trajectory
|
| 20 |
from .features import DecoderMasks, EncoderFeatures, decoder_masks, encoder_features
|
| 21 |
from .normalize import normalize_inputs, speed_masks
|
|
|
|
| 22 |
from .postprocess import (TurnIndicatorManager, denoising_steps_ego, denormalize, predicted_paths,
|
| 23 |
trajectory_from_poses)
|
| 24 |
|
|
@@ -83,15 +84,21 @@ def make_output(final_x0: np.ndarray, logit: np.ndarray, prepared: Prepared, nor
|
|
| 83 |
``RUNTIME_PARAMS``; ``denoising_steps``: the 11 iterates ``[321, 81, 4]`` when ``return_denoising_steps``."""
|
| 84 |
x0 = np.asarray(final_x0, np.float32).reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM)
|
| 85 |
mean, std = normalization.state()
|
| 86 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 87 |
traj = trajectory_from_poses(denorm[0], (0.0, 0.0, 0.0),
|
| 88 |
velocity_smoothing_window=int(params["velocity_smoothing_window"]),
|
| 89 |
enable_force_stop=prepared.enable_force_stop,
|
| 90 |
stopping_threshold=float(params["stopping_threshold"]))
|
| 91 |
manager = TurnIndicatorManager(keep_offset=float(params["turn_indicator_keep_offset"]))
|
| 92 |
decision = manager.evaluate(np.asarray(logit, np.float32).reshape(-1), 0.0, prepared.prev_report)
|
| 93 |
-
|
| 94 |
-
|
| 95 |
(0, C.OUTPUT_T, len(PREDICTED_AGENT_COLUMNS)), np.float32)
|
| 96 |
out_meta: Dict[str, Any] = {"predicted_agent_columns": list(PREDICTED_AGENT_COLUMNS),
|
| 97 |
"predicted_agent_rows": [int(r) for r in rows],
|
|
|
|
| 19 |
from ..ttaw.outputs import Trajectory
|
| 20 |
from .features import DecoderMasks, EncoderFeatures, decoder_masks, encoder_features
|
| 21 |
from .normalize import normalize_inputs, speed_masks
|
| 22 |
+
from .postprocess import HOST_FAST
|
| 23 |
from .postprocess import (TurnIndicatorManager, denoising_steps_ego, denormalize, predicted_paths,
|
| 24 |
trajectory_from_poses)
|
| 25 |
|
|
|
|
| 84 |
``RUNTIME_PARAMS``; ``denoising_steps``: the 11 iterates ``[321, 81, 4]`` when ``return_denoising_steps``."""
|
| 85 |
x0 = np.asarray(final_x0, np.float32).reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM)
|
| 86 |
mean, std = normalization.state()
|
| 87 |
+
rows = prepared.neighbor_rows
|
| 88 |
+
if HOST_FAST: # only the rows read below: the ego and the emitted neighbours (element-wise: the same values)
|
| 89 |
+
idx = np.concatenate([[0], np.asarray(rows, np.int64) + 1])
|
| 90 |
+
pick = lambda v: v[idx] if v.shape[0] == C.MAX_NUM_AGENTS else v # noqa: E731
|
| 91 |
+
denorm = (x0[idx, 1:] * pick(std) + pick(mean)).astype(np.float32) # [1 + n, 80, 4]
|
| 92 |
+
else:
|
| 93 |
+
denorm = denormalize(x0, mean, std) # [321, 80, 4]
|
| 94 |
traj = trajectory_from_poses(denorm[0], (0.0, 0.0, 0.0),
|
| 95 |
velocity_smoothing_window=int(params["velocity_smoothing_window"]),
|
| 96 |
enable_force_stop=prepared.enable_force_stop,
|
| 97 |
stopping_threshold=float(params["stopping_threshold"]))
|
| 98 |
manager = TurnIndicatorManager(keep_offset=float(params["turn_indicator_keep_offset"]))
|
| 99 |
decision = manager.evaluate(np.asarray(logit, np.float32).reshape(-1), 0.0, prepared.prev_report)
|
| 100 |
+
paths = (predicted_paths(denorm, rows.size) if HOST_FAST else predicted_paths(denorm, C.MAX_NUM_NEIGHBORS)[rows]) \
|
| 101 |
+
if rows.size else np.zeros(
|
| 102 |
(0, C.OUTPUT_T, len(PREDICTED_AGENT_COLUMNS)), np.float32)
|
| 103 |
out_meta: Dict[str, Any] = {"predicted_agent_columns": list(PREDICTED_AGENT_COLUMNS),
|
| 104 |
"predicted_agent_rows": [int(r) for r in rows],
|
code/tt_diffusion_planner/host/postprocess.py
CHANGED
|
@@ -102,7 +102,19 @@ class EgoTrajectory:
|
|
| 102 |
self.velocity, self.acceleration], axis=-1).astype(np.float32)
|
| 103 |
|
| 104 |
|
| 105 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 106 |
velocity_smoothing_window: int = C.VELOCITY_SMOOTHING_WINDOW,
|
| 107 |
enable_force_stop: bool = True,
|
| 108 |
stopping_threshold: float = C.STOPPING_THRESHOLD) -> EgoTrajectory:
|
|
@@ -151,6 +163,61 @@ def trajectory_from_poses(poses_xycs: np.ndarray, base_position: Tuple[float, fl
|
|
| 151 |
return EgoTrajectory(pos, cs, quat, tf2_yaw(quat), vel, accel, tfs, force_stop)
|
| 152 |
|
| 153 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 154 |
def predicted_paths(denorm: np.ndarray, n_neighbors: int) -> np.ndarray:
|
| 155 |
"""``create_predicted_objects`` poses for the first ``n_neighbors`` neighbours: ``[n, 80, 5]`` float32 =
|
| 156 |
x, y, yaw (``tf2::getYaw``), cos, sin. ``denorm`` is ``[321, 80, 4]`` (agent 0 = ego)."""
|
|
|
|
| 102 |
self.velocity, self.acceleration], axis=-1).astype(np.float32)
|
| 103 |
|
| 104 |
|
| 105 |
+
|
| 106 |
+
def _host_fast() -> bool:
|
| 107 |
+
"""``DIFFUSION_PLANNER_HOST_FAST`` (OPT round 5 item 3), read once: the vectorised host functions (bit-exact)
|
| 108 |
+
instead of the first port's ``*_ref`` versions."""
|
| 109 |
+
from ..tt.config import KNOBS
|
| 110 |
+
|
| 111 |
+
return bool(KNOBS.read().HOST_FAST)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
HOST_FAST = _host_fast()
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def trajectory_from_poses_ref(poses_xycs: np.ndarray, base_position: Tuple[float, float, float] = (0.0, 0.0, 0.0), *,
|
| 118 |
velocity_smoothing_window: int = C.VELOCITY_SMOOTHING_WINDOW,
|
| 119 |
enable_force_stop: bool = True,
|
| 120 |
stopping_threshold: float = C.STOPPING_THRESHOLD) -> EgoTrajectory:
|
|
|
|
| 163 |
return EgoTrajectory(pos, cs, quat, tf2_yaw(quat), vel, accel, tfs, force_stop)
|
| 164 |
|
| 165 |
|
| 166 |
+
def trajectory_from_poses(poses_xycs: np.ndarray, base_position: Tuple[float, float, float] = (0.0, 0.0, 0.0), *,
|
| 167 |
+
velocity_smoothing_window: int = C.VELOCITY_SMOOTHING_WINDOW,
|
| 168 |
+
enable_force_stop: bool = True,
|
| 169 |
+
stopping_threshold: float = C.STOPPING_THRESHOLD) -> EgoTrajectory:
|
| 170 |
+
"""``get_trajectory_from_poses`` (``postprocessing_utils.cpp:362-455``) for poses ``[N, 4]`` = denormalised
|
| 171 |
+
(x, y, cos, sin) float32 of one agent, positions at z = ``base_position[2]`` (identity transform)."""
|
| 172 |
+
if not HOST_FAST:
|
| 173 |
+
return trajectory_from_poses_ref(poses_xycs, base_position, velocity_smoothing_window=velocity_smoothing_window,
|
| 174 |
+
enable_force_stop=enable_force_stop, stopping_threshold=stopping_threshold)
|
| 175 |
+
p = np.asarray(poses_xycs, np.float32)
|
| 176 |
+
n = p.shape[0]
|
| 177 |
+
dt = C.TRAJECTORY_DT
|
| 178 |
+
pos = np.zeros((n, 3), np.float64)
|
| 179 |
+
pos[:, 0] = p[:, 0].astype(np.float64)
|
| 180 |
+
pos[:, 1] = p[:, 1].astype(np.float64)
|
| 181 |
+
pos[:, 2] = float(base_position[2])
|
| 182 |
+
cs = p[:, 2:4].astype(np.float64).copy()
|
| 183 |
+
quat = quaternion_from_cos_sin(cs[:, 0], cs[:, 1])
|
| 184 |
+
# the node's loops, vectorised with the same float64 / float32 operations in the same order (bit-exact)
|
| 185 |
+
prev = np.asarray(base_position, np.float64)
|
| 186 |
+
px, py, pz = float(prev[0]), float(prev[1]), float(prev[2])
|
| 187 |
+
hyp = math.hypot
|
| 188 |
+
dist = []
|
| 189 |
+
for x, y, z in pos.tolist(): # math.hypot (3 args) kept: np.hypot rounds differently
|
| 190 |
+
dist.append(hyp(x - px, y - py, z - pz))
|
| 191 |
+
px, py, pz = x, y, z
|
| 192 |
+
vel = (np.asarray(dist, np.float64) / dt).astype(np.float32)
|
| 193 |
+
w = int(velocity_smoothing_window)
|
| 194 |
+
if n <= w:
|
| 195 |
+
raise ValueError("velocity_smoothing_window must be smaller than number of points")
|
| 196 |
+
thr = np.float32(stopping_threshold)
|
| 197 |
+
m = n - w + 1
|
| 198 |
+
v64 = vel.astype(np.float64)
|
| 199 |
+
acc = np.zeros(m, np.float64)
|
| 200 |
+
for k in range(w): # acc = ((0 + v[i]) + v[i+1]) + ... per window, as the C++ loop
|
| 201 |
+
acc = acc + v64[k:k + m]
|
| 202 |
+
sm = (acc / float(w)).astype(np.float32) # window i reads only original values (it writes vel[i] last)
|
| 203 |
+
vel[:m] = sm
|
| 204 |
+
force_stop = False
|
| 205 |
+
stop = -1
|
| 206 |
+
if enable_force_stop and m > 1:
|
| 207 |
+
hit = np.flatnonzero((np.abs(sm[:-1]) > thr) & (np.abs(sm[1:]) < thr))
|
| 208 |
+
if hit.size:
|
| 209 |
+
stop, force_stop = int(hit[0]) + 1, True
|
| 210 |
+
if force_stop: # from the first stop on: zero velocity, the pose frozen
|
| 211 |
+
vel[stop:] = np.float32(0.0)
|
| 212 |
+
pos[stop:], cs[stop:], quat[stop:] = pos[stop - 1], cs[stop - 1], quat[stop - 1]
|
| 213 |
+
else:
|
| 214 |
+
vel[m:] = vel[m - 1]
|
| 215 |
+
accel = np.zeros(n, np.float32)
|
| 216 |
+
accel[:n - 1] = ((vel[1:].astype(np.float64) - vel[:-1].astype(np.float64)) / dt).astype(np.float32)
|
| 217 |
+
tfs = dt * (np.arange(n, dtype=np.float64) + 1.0)
|
| 218 |
+
return EgoTrajectory(pos, cs, quat, tf2_yaw(quat), vel, accel, tfs, force_stop)
|
| 219 |
+
|
| 220 |
+
|
| 221 |
def predicted_paths(denorm: np.ndarray, n_neighbors: int) -> np.ndarray:
|
| 222 |
"""``create_predicted_objects`` poses for the first ``n_neighbors`` neighbours: ``[n, 80, 5]`` float32 =
|
| 223 |
x, y, yaw (``tf2::getYaw``), cos, sin. ``denorm`` is ``[321, 80, 4]`` (agent 0 = ego)."""
|
code/tt_diffusion_planner/tests/test_bundle_host.py
CHANGED
|
@@ -133,7 +133,11 @@ def test_serve_env_pins_the_numerics(manifest):
|
|
| 133 |
env = manifest["serve"]["env"]
|
| 134 |
pins = KNOBS.serve_env()
|
| 135 |
assert set(pins) == {f"DIFFUSION_PLANNER_{k}" for k in
|
| 136 |
-
("LN_FP32", "HIDDEN_FP32", "SPLIT_MATMUL", "ATTN_FP32_ACC", "ATTN_MATMUL"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 137 |
assert {k: env.get(k) for k in pins} == pins
|
| 138 |
assert env.get("DIFFUSION_PLANNER_PRECISION") == ""
|
| 139 |
for profile in manifest.get("serve_profiles") or []:
|
|
|
|
| 133 |
env = manifest["serve"]["env"]
|
| 134 |
pins = KNOBS.serve_env()
|
| 135 |
assert set(pins) == {f"DIFFUSION_PLANNER_{k}" for k in
|
| 136 |
+
("LN_FP32", "HIDDEN_FP32", "SPLIT_MATMUL", "ATTN_FP32_ACC", "ATTN_MATMUL",
|
| 137 |
+
"ENC_CH2D", "ATTN_FAST",
|
| 138 |
+
"DEC_MMCFG", "LN_KERNEL", "LN_RESID", "SPLIT_KCAT", "LN_SFPU_BCAST", "LN_SPLIT", "KCAT_ACT", "ATTN_SMASK",
|
| 139 |
+
"ATTN_SMSM", "ENC_KCAT", "LIN_ACT", "ATTN_FUSED", "KCAT_EMIT", "LN_TR", "KCAT_L1", "KCAT_ACT_ONCE", "ATTN_L1", "DEC_L1", "ENC_L1", "FUS_L1",
|
| 140 |
+
"COMPACT", "AGENT_BUCKETS", "HOST_FAST", "INPUT_TRIM")}
|
| 141 |
assert {k: env.get(k) for k in pins} == pins
|
| 142 |
assert env.get("DIFFUSION_PLANNER_PRECISION") == ""
|
| 143 |
for profile in manifest.get("serve_profiles") or []:
|
code/tt_diffusion_planner/tests/test_grid_fit_host.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""The swept 2-D multicast configs fit the 12x10 ETH grid and fall back to the auto config on the 11x10 WORKER grid
|
| 3 |
+
(republish 2026-10-11: ``dispatch="worker"`` failed at capture with "Num output blocks along x (12) must be smaller
|
| 4 |
+
than or equal to the number of columns in compute grid (11)")."""
|
| 5 |
+
from types import SimpleNamespace
|
| 6 |
+
|
| 7 |
+
from tt_diffusion_planner.tt.layers import DEC_MM_CONFIGS, DEC_ROWS, KCAT_CONFIGS, fit_pcm, mcast_fits
|
| 8 |
+
|
| 9 |
+
ETH, WORKER = SimpleNamespace(x=12, y=10), SimpleNamespace(x=11, y=10)
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def test_every_swept_config_fits_the_eth_grid():
|
| 13 |
+
for rows in (DEC_ROWS, 32, 64, 96, 128, 192):
|
| 14 |
+
for (k, n, _ps), (tr, pcm, pcn, _kb, _sw) in DEC_MM_CONFIGS.items():
|
| 15 |
+
assert mcast_fits(tr, fit_pcm(tr, pcm, rows, ETH), pcn, rows, n, ETH), (rows, k, n)
|
| 16 |
+
for (rows, _ktp, n), (tr, pcm, pcn, _kb, _sw) in KCAT_CONFIGS.items():
|
| 17 |
+
assert mcast_fits(tr, pcm, pcn, rows, n, ETH), (rows, n)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def test_the_qkv_config_does_not_fit_the_worker_grid():
|
| 21 |
+
tr, pcm, pcn, _kb, _sw = KCAT_CONFIGS[(352, 25, 768)]
|
| 22 |
+
assert not mcast_fits(tr, pcm, pcn, 352, 768, WORKER) # 24 N tiles / 2 = 12 blocks on x > 11
|
| 23 |
+
assert mcast_fits(False, 2, 1, 352, 256, WORKER)
|
code/tt_diffusion_planner/tests/test_tt_params_host.py
CHANGED
|
@@ -156,3 +156,31 @@ def test_plan_inputs_layout(prepared):
|
|
| 156 |
np.testing.assert_array_equal(y0[:C.MAX_NUM_AGENTS, 4:], prepared.x_T.reshape(C.MAX_NUM_AGENTS, -1)[:, 4:])
|
| 157 |
warm = I.warmup_inputs()
|
| 158 |
assert all(warm[k].shape == v for k, v in I.INPUT_SPECS.items())
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 156 |
np.testing.assert_array_equal(y0[:C.MAX_NUM_AGENTS, 4:], prepared.x_T.reshape(C.MAX_NUM_AGENTS, -1)[:, 4:])
|
| 157 |
warm = I.warmup_inputs()
|
| 158 |
assert all(warm[k].shape == v for k, v in I.INPUT_SPECS.items())
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
def test_compact_bucket_covers_every_needed_row():
|
| 162 |
+
"""``COMPACT``: the bucket holds the ego, every valid self-attention key and every emitted neighbour row (the
|
| 163 |
+
rows past it are masked keys, never read back), and falls back to the full 352 rows."""
|
| 164 |
+
from types import SimpleNamespace
|
| 165 |
+
|
| 166 |
+
from tt_diffusion_planner.tt.model import bucket_rows, needed_rows
|
| 167 |
+
|
| 168 |
+
def prep(valid_idx, emitted):
|
| 169 |
+
v = np.zeros(C.MAX_NUM_AGENTS, bool)
|
| 170 |
+
v[0] = True
|
| 171 |
+
v[list(valid_idx)] = True
|
| 172 |
+
nb = np.zeros(C.MAX_NUM_NEIGHBORS, bool)
|
| 173 |
+
nb[[i - 1 for i in valid_idx]] = True
|
| 174 |
+
return SimpleNamespace(decoder=SimpleNamespace(agent_valid=v), neighbor_rows=np.asarray(emitted, int),
|
| 175 |
+
features=SimpleNamespace(valid={"neighbor": nb}))
|
| 176 |
+
|
| 177 |
+
b = (32, 64, 96, 128, 192)
|
| 178 |
+
assert needed_rows(prep([], [])) == 1 and bucket_rows(prep([], []), b) == 32
|
| 179 |
+
assert needed_rows(prep([1, 2, 31], [0, 1, 30])) == 32 and bucket_rows(prep([31], []), b) == 32
|
| 180 |
+
assert bucket_rows(prep([32], []), b) == 64 # the 33rd row is a valid key
|
| 181 |
+
assert bucket_rows(prep([5], [70]), b) == 96 # emitted neighbour 70 = decoder row 71
|
| 182 |
+
assert bucket_rows(prep([200], []), b) == T.AGENTS
|
| 183 |
+
assert bucket_rows(prep([5], []), ()) == T.AGENTS # COMPACT off
|
| 184 |
+
p = prep([3], [])
|
| 185 |
+
p.features.valid["neighbor"][40] = True # a valid neighbour token (row 41)
|
| 186 |
+
assert needed_rows(p) == 42 and bucket_rows(p, b) == 64
|
code/tt_diffusion_planner/tt/attention.py
ADDED
|
@@ -0,0 +1,96 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""The fp32 matmul attention of the fusion transformer and the decoder (C20 ``attention_matmul`` math) with stock-op
|
| 3 |
+
program choices tuned for the planner's shapes (``ATTN_FAST``, OPT round 1 item 4a).
|
| 4 |
+
|
| 5 |
+
``softmax(scale * Q K^T + mask) V`` as in ``ttaw.ops.attention.attention_matmul``, with:
|
| 6 |
+
|
| 7 |
+
- (``mode >= 1``) ``P V`` on a ``MatmulMultiCoreReuseProgramConfig`` with one output tile per core (8 heads x
|
| 8 |
+
Sq / 32 = 88 cores at Sq = 352) instead of the auto config's 11-18 cores; the K-sum in one block
|
| 9 |
+
(``in0_block_w`` = Sk / 32): bit-identical to the auto config in the device sweep (73 -> 21 us self, 71 -> 32 us
|
| 10 |
+
cross, 81 -> 42 us fusion; ``logs/diffusion-planner/opt_r1/attn_sweep.log``);
|
| 11 |
+
- (``mode == 2``) the scale applied to Q (``[1, 8, Sq, 32]``, 11x smaller than the scores) before ``Q K^T``
|
| 12 |
+
instead of to the scores: one small multiply instead of the large one; a precision change (rounding order);
|
| 13 |
+
- (``mode == 3``) the scale as an SFPU pre-activation of the mask add (``input_tensor_a_activations``): one
|
| 14 |
+
program instead of two over the scores (75 -> 42 us), but NOT bit-identical (final_x0 rel 4e-3; e2e ego mean
|
| 15 |
+
0.143 -> 0.173 m over the 99 scenes): measured and rejected (OPT_REPORT round 1).
|
| 16 |
+
- ``Q K^T`` keeps the auto config (the reuse configs of the sweep were slower).
|
| 17 |
+
"""
|
| 18 |
+
from __future__ import annotations
|
| 19 |
+
|
| 20 |
+
from typing import Any, Optional
|
| 21 |
+
|
| 22 |
+
__all__ = ["attention_matmul", "pv_config"]
|
| 23 |
+
|
| 24 |
+
TILE = 32
|
| 25 |
+
# per-core L1 budget for the P block of one core (fp32, double-buffered): 18 tiles x 4 KiB x 2 = 144 KiB
|
| 26 |
+
MAX_KB_TILES = 18
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def pv_config(device: Any, heads: int, sq: int, sk: int):
|
| 30 |
+
"""``MatmulMultiCoreReuseProgramConfig`` for ``P [1, H, Sq, Sk] @ V [1, H, Sk, 32]``: one ``[32, 32]`` output
|
| 31 |
+
tile per core (``per_core_M = 1`` when ``H * Sq / 32`` fits the grid, else 2)."""
|
| 32 |
+
import ttnn
|
| 33 |
+
|
| 34 |
+
grid = device.compute_with_storage_grid_size()
|
| 35 |
+
mt, kt = sq // TILE, sk // TILE
|
| 36 |
+
pcm = 1 if heads * mt <= grid.x * grid.y else 2
|
| 37 |
+
kb = max(d for d in range(1, min(kt, MAX_KB_TILES) + 1) if kt % d == 0)
|
| 38 |
+
return ttnn.MatmulMultiCoreReuseProgramConfig(compute_with_storage_grid_size=grid, in0_block_w=kb,
|
| 39 |
+
out_subblock_h=1, out_subblock_w=1, per_core_M=pcm, per_core_N=1)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def attention_matmul(q: Any, k: Any, v: Any, *, scale: float, attn_mask: Any = None,
|
| 43 |
+
compute_kernel_config: Any = None, mode: int = 1, pv_pc: Optional[Any] = None,
|
| 44 |
+
smask: bool = False, smsm: int = 0):
|
| 45 |
+
"""``q`` ``[1, H, Sq, D]``, ``k`` / ``v`` ``[1, H, Sk, D]`` (tile-aligned ``Sk``) -> ``[1, H, Sq, D]``.
|
| 46 |
+
``mode=0``: exactly ``ttaw.ops.attention.attention_matmul``. ``smask`` (``ATTN_SMASK``, mode 1 with a mask):
|
| 47 |
+
the scale multiply and the mask add as one bit-identical generic_op (``tt/smask_kernel.py``); ``smsm``
|
| 48 |
+
(``ATTN_SMSM`` 1 / 2, with ``smask``): the scale, the mask and the softmax as one generic_op
|
| 49 |
+
(``tt/smsm_kernel.py``; 2 = the scale multiply by an immediate)."""
|
| 50 |
+
import ttnn
|
| 51 |
+
|
| 52 |
+
from ..ttaw.ops import attention as A
|
| 53 |
+
from ..ttaw.precision import compute_kernel_config as ckc
|
| 54 |
+
|
| 55 |
+
cfg = compute_kernel_config if compute_kernel_config is not None else ckc("HiFi4", fp32_acc=True)
|
| 56 |
+
dev = q.device() if callable(getattr(q, "device", None)) else None
|
| 57 |
+
if mode and (dev is None or not hasattr(dev, "compute_with_storage_grid_size")
|
| 58 |
+
or not hasattr(getattr(ttnn, "UnaryOpType", None), "MUL_UNARY_SFPU")):
|
| 59 |
+
mode = 0 # the host fake ttnn: the stock path (same math)
|
| 60 |
+
if not mode:
|
| 61 |
+
return A.attention_matmul(q, k, v, scale=scale, attn_mask=attn_mask, compute_kernel_config=cfg)
|
| 62 |
+
sk = int(k.shape[-2])
|
| 63 |
+
if sk % TILE:
|
| 64 |
+
raise ValueError(f"attention_matmul needs a tile-aligned Sk (got {sk})")
|
| 65 |
+
if mode == 2:
|
| 66 |
+
q = ttnn.multiply(q, float(scale))
|
| 67 |
+
scores = ttnn.matmul(q, k, transpose_b=True, compute_kernel_config=cfg)
|
| 68 |
+
fused = False
|
| 69 |
+
if smsm and smask and mode == 1 and attn_mask is not None:
|
| 70 |
+
from .smsm_kernel import scale_mask_softmax
|
| 71 |
+
from .smsm_kernel import supported as smsm_supported
|
| 72 |
+
|
| 73 |
+
if smsm_supported(scores, attn_mask):
|
| 74 |
+
probs = scale_mask_softmax(scores, scale, attn_mask, scale_mode=int(smsm) - 1)
|
| 75 |
+
if pv_pc is None:
|
| 76 |
+
pv_pc = pv_config(dev, int(q.shape[1]), int(q.shape[-2]), sk)
|
| 77 |
+
return ttnn.matmul(probs, v, compute_kernel_config=cfg, program_config=pv_pc)
|
| 78 |
+
if smask and mode == 1 and attn_mask is not None:
|
| 79 |
+
from .smask_kernel import scale_mask, supported
|
| 80 |
+
|
| 81 |
+
if supported(scores, attn_mask):
|
| 82 |
+
scores, fused = scale_mask(scores, scale, attn_mask), True
|
| 83 |
+
if fused:
|
| 84 |
+
pass
|
| 85 |
+
elif mode == 3 and attn_mask is not None:
|
| 86 |
+
act = [ttnn.UnaryWithParam(ttnn.UnaryOpType.MUL_UNARY_SFPU, float(scale))]
|
| 87 |
+
scores = ttnn.add(scores, attn_mask, input_tensor_a_activations=act)
|
| 88 |
+
else:
|
| 89 |
+
if mode != 2:
|
| 90 |
+
scores = ttnn.multiply(scores, float(scale))
|
| 91 |
+
if attn_mask is not None:
|
| 92 |
+
scores = ttnn.add(scores, attn_mask)
|
| 93 |
+
probs = ttnn.softmax(scores, dim=-1, numeric_stable=True, compute_kernel_config=cfg)
|
| 94 |
+
if pv_pc is None:
|
| 95 |
+
pv_pc = pv_config(dev, int(q.shape[1]), int(q.shape[-2]), sk)
|
| 96 |
+
return ttnn.matmul(probs, v, compute_kernel_config=cfg, program_config=pv_pc)
|
code/tt_diffusion_planner/tt/config.py
CHANGED
|
@@ -97,6 +97,108 @@ def _knobs():
|
|
| 97 |
Knob("ATTN_MATMUL", "enc.fusion.attn,dec.*", "comma-separated module globs (enc.fusion.attn, "
|
| 98 |
"dec.self_attn, dec.cross_attn) whose attention runs as fp32 matmuls + softmax (C20 attention_matmul) "
|
| 99 |
"instead of the bf16 SDPA kernel; 'none' = empty"),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 100 |
])
|
| 101 |
|
| 102 |
|
|
|
|
| 97 |
Knob("ATTN_MATMUL", "enc.fusion.attn,dec.*", "comma-separated module globs (enc.fusion.attn, "
|
| 98 |
"dec.self_attn, dec.cross_attn) whose attention runs as fp32 matmuls + softmax (C20 attention_matmul) "
|
| 99 |
"instead of the bf16 SDPA kernel; 'none' = empty"),
|
| 100 |
+
Knob("ENC_CH2D", True, "the encoder linears on batched [1, E, T, C] activations run as one 2-D matmul "
|
| 101 |
+
"over all rows (free [1, 1, E*T, C] view for T = 64, else a 1-D program config with fuse_batch) "
|
| 102 |
+
"instead of the stock per-element tiling on 4-8 cores (OPT round 1 item 1); 0 = the stock call"),
|
| 103 |
+
Knob("ATTN_FAST", 1, "fp32 matmul attention variant (tt/attention.py, OPT round 1 item 4a): 0 = "
|
| 104 |
+
"ttaw.ops.attention.attention_matmul; 1 = P.V on one output tile per core (bit-identical); 2 = 1 + "
|
| 105 |
+
"the scale on Q; 3 = 1 + the scale fused into the mask add (not bit-identical, rejected)",
|
| 106 |
+
choices=(0, 1, 2, 3)),
|
| 107 |
+
Knob("DEC_MMCFG", True, "explicit 2-D multicast program configs for the decoder's 352-row matmuls "
|
| 108 |
+
"(tt/layers.py DEC_MM_CONFIGS, the fastest bit-identical config per shape and split pass of the device "
|
| 109 |
+
"sweep, OPT round 1 item 2a); 0 = the auto config"),
|
| 110 |
+
Knob("LN_KERNEL", 2, "the fp32 LayerNorm decomposition (LN_FP32 modules) as one fused generic_op "
|
| 111 |
+
"program (tt/ln_kernel.py, the same SFPU LLK sequence, bit-identical: OPT round 2 item 3) instead of "
|
| 112 |
+
"7-9 stock programs: 1 = per-tile unpacker / SFPU inits, 2 = one init per phase; 0 = the stock "
|
| 113 |
+
"decomposition", choices=(0, 1, 2)),
|
| 114 |
+
Knob("LN_RESID", True, "with LN_KERNEL: the residual add in front of a fused fp32 LayerNorm (mixer blocks: "
|
| 115 |
+
"x + y, x + ch2; decoder: h + gate * a, h + m2b, ...) in the same program, which writes the new stream "
|
| 116 |
+
"and the normalised rows (OPT round 2 item 5); 0 = separate add / multiply programs"),
|
| 117 |
+
Knob("SPLIT_KCAT", 2, "decoder split matmuls (SPLIT_MATMUL dec.*, tile-aligned K) as one fp32 matmul over "
|
| 118 |
+
"the concatenated K axis [x_hi|x_hi|x_lo|1] @ [w_hi;w_lo;w_hi;b] (OPT round 2 item 2): 0 = three "
|
| 119 |
+
"passes + adds (the published numerics); 1 = stock-op operand build (typecast / subtract / concat); 2 = "
|
| 120 |
+
"the operand built by one generic_op (tt/kcat_kernel.py); not bit-identical to 0 (one fp32 accumulation, "
|
| 121 |
+
"exact bias), gates green",
|
| 122 |
+
choices=(0, 1, 2)),
|
| 123 |
+
Knob("ATTN_SMASK", True, "fp32 matmul attention (ATTN_FAST=1) with a mask: the score scale multiply and the "
|
| 124 |
+
"mask add as one generic_op (tt/smask_kernel.py, the same SFPU LLK calls: bit-identical; OPT round 2 "
|
| 125 |
+
"item 4b) instead of two binary_ng programs; 0 = the two programs"),
|
| 126 |
+
Knob("ATTN_SMSM", 1, "with ATTN_SMASK: the score scale, the mask add and the softmax as one generic_op "
|
| 127 |
+
"(tt/smsm_kernel.py: the smask SFPU sequence, then the stock softmax kernel_lib calls on the row in L1; "
|
| 128 |
+
"OPT round 3 item 1, bit-identical) instead of the smask program + ttnn.softmax: 1 = the scale as a "
|
| 129 |
+
"multiply by a scale tile (mul_binary_tile, as smask), 2 = the same fp32 SFPU multiply by an immediate "
|
| 130 |
+
"(mul_unary_tile, no scale tile copy); 0 = the two programs", choices=(0, 1, 2)),
|
| 131 |
+
Knob("ATTN_FUSED", True, "with ATTN_SMSM: each fp32 matmul attention (decoder self / cross, fusion) as one "
|
| 132 |
+
"generic_op (tt/fattn_kernel.py: Q K^T into DEST, the smsm phases, P V accumulated in DEST, Q / K / V "
|
| 133 |
+
"read in place and the heads merged on write; OPT round 3 item 4) instead of head split, Q K^T, smsm, "
|
| 134 |
+
"P V and head merge; 0 = those programs"),
|
| 135 |
+
Knob("KCAT_EMIT", True, "with SPLIT_KCAT=2: the producers of the decoder's K-concatenated split linears' "
|
| 136 |
+
"inputs (the split-row LayerNorms -> qkv / mlp fc1 / cross q, the fused attention -> attn out / cross "
|
| 137 |
+
"out) write the split operand [x_hi | x_hi | x_lo | 1] themselves (the kcat LLK calls on their output "
|
| 138 |
+
"tile; OPT round 3 item 5) instead of x + a kcat operand program; 0 = the operand programs"),
|
| 139 |
+
Knob("LN_TR", True, "with LN_KERNEL and LN_RESID: the two transposes around each mixer block's token-mixing "
|
| 140 |
+
"MLP done inside the fused LayerNorm programs (n1 writes its output per-entity transposed, n2 reads the "
|
| 141 |
+
"token-mixing output transposed; the stock transpose LLK, exact; OPT round 3 item 6) instead of two "
|
| 142 |
+
"ttnn.transpose programs (mixer trunks on the one-core-per-row LN kernel); 0 = the transpose programs"),
|
| 143 |
+
Knob("ENC_KCAT", True, "with SPLIT_KCAT=2: the encoder's split linears (SPLIT_MATMUL enc.pre.*, "
|
| 144 |
+
"enc.island.*; any K, blocks padded to whole tiles) as one K-concatenated fp32 matmul too (OPT round 3 "
|
| 145 |
+
"item 2; one fp32 accumulation, exact bias: a precision change of the SPLIT_KCAT kind); 0 = three "
|
| 146 |
+
"passes + adds"),
|
| 147 |
+
Knob("LIN_ACT", False, "the GELU of the encoder mixer linears (token / channel MLP fc1, bf16 output) in the "
|
| 148 |
+
"matmul epilogue (an explicit copy of the stock auto config + fused_activation, OPT round 3 item 3) "
|
| 149 |
+
"instead of a unary program on the rounded bf16 output; a precision change (GELU before the bf16 "
|
| 150 |
+
"rounding); 0 = the unary program"),
|
| 151 |
+
Knob("KCAT_ACT", True, "with SPLIT_KCAT=2: the GELU of a decoder linear that feeds another K-concatenated "
|
| 152 |
+
"split linear (preproj.fc1 -> fc2, mlp fc1 -> fc2) applied inside the next one's operand build (the "
|
| 153 |
+
"same SFPU LLK: bit-identical) instead of its own unary program; 0 = the unary program"),
|
| 154 |
+
Knob("LN_SPLIT", True, "with LN_KERNEL: a fused fp32 LayerNorm over few tile rows (the decoder's 11) spread "
|
| 155 |
+
"over Wt cores per row (tt/ln_kernel.py layer_norm_fp32_split: the root core folds the gathered tiles "
|
| 156 |
+
"in order and broadcasts the statistics, bit-identical) instead of one core per row; 0 = one core per "
|
| 157 |
+
"row"),
|
| 158 |
+
Knob("KCAT_L1", True, "with SPLIT_KCAT=2: the decoder's K = 1024 split operands (mlp fc2, final p4) written "
|
| 159 |
+
"L1 block-sharded by the operand build and read in place by a 2-D multicast matmul with in0 sharded "
|
| 160 |
+
"(tt/layers.py KCAT_L1_CONFIGS; OPT round 4 item 1, bit-identical) instead of a DRAM round trip of "
|
| 161 |
+
"the 4.7 MB operand; 0 = DRAM interleaved"),
|
| 162 |
+
Knob("KCAT_ACT_ONCE", False, "with KCAT_ACT: the operand build applies the deferred GELU to one DEST copy of "
|
| 163 |
+
"each tile and copies it to the second (copy_dest_values; OPT round 4 item 2, bit-identical) instead "
|
| 164 |
+
"of running the GELU on both copies (measured, not kept: +0.06 ms); 0 = twice"),
|
| 165 |
+
Knob("ATTN_L1", 2, "with ATTN_FUSED: the fused attention's inputs in L1 (interleaved) instead of DRAM: the "
|
| 166 |
+
"decoder qkv / cross q projections write L1, the hoisted cross K / V heads and the masks are L1-resident "
|
| 167 |
+
"(OPT round 4 item 3, bit-identical; the 11 query-row cores of a head re-read the same K / V tiles); "
|
| 168 |
+
"1 = also the fusion q / kv projections (the stock linear then runs on 6 cores with a separate bias "
|
| 169 |
+
"add: +0.3 ms, item 8); 2 = not those; 0 = DRAM", choices=(0, 1, 2)),
|
| 170 |
+
Knob("DEC_L1", 2, "the decoder blocks' intermediates (the split-row LN outputs and stream, the fused "
|
| 171 |
+
"attention outputs, the out / mlp / cross-out linear outputs) in L1 (interleaved) instead of DRAM (OPT "
|
| 172 |
+
"round 4 item 4, bit-identical); 2 = also the pre-projection and the final layer (item 6); 0 = DRAM",
|
| 173 |
+
choices=(0, 1, 2)),
|
| 174 |
+
Knob("ENC_L1", False, "the encoder mixer blocks' intermediates (the fused LN outputs and stream, the token / "
|
| 175 |
+
"channel MLP outputs) in L1 (interleaved) instead of DRAM (OPT round 4 item 5, bit-identical; measured, "
|
| 176 |
+
"not kept: +1.67 ms); 0 = DRAM"),
|
| 177 |
+
Knob("FUS_L1", False, "the encoder fusion blocks' intermediates (LN outputs, stream adds, attention / out / "
|
| 178 |
+
"MLP outputs) in L1 (interleaved) instead of DRAM (OPT round 4 item 7; measured, not kept: +0.39 ms and "
|
| 179 |
+
"not bit-identical, the stock ops pick other programs for L1 outputs); 0 = DRAM"),
|
| 180 |
+
Knob("COMPACT", 2, "exact compaction (OPT round 5): besides the full-capacity plan, one trace per agent "
|
| 181 |
+
"bucket (AGENT_BUCKETS), picked per request from the host masks. 1 = the decoder runs on the first R "
|
| 182 |
+
"rows only (the needed rows: ego, every valid self-attention key, every emitted neighbour, every valid "
|
| 183 |
+
"neighbour token; the rows past R are masked keys and never read; item 1); 2 = also the encoder's "
|
| 184 |
+
"neighbour trunk and head on the first R neighbours, the other neighbour tokens zero (invalid: "
|
| 185 |
+
"token_valid zeroes them anyway; item 2). The matmul configs keep their K blocking, so the needed rows "
|
| 186 |
+
"are the same values; 0 = the full 352 rows / 320 neighbours always", choices=(0, 1, 2)),
|
| 187 |
+
Knob("AGENT_BUCKETS", "32,64,96,128,192", "with COMPACT: the decoder row buckets (multiples of 32 below "
|
| 188 |
+
"352), one captured trace each"),
|
| 189 |
+
Knob("HOST_FAST", True, "the host pre- and post-processing vectorised (OPT round 5 item 3: the normalisation "
|
| 190 |
+
"computes only the rows that are not all-small, the neighbour features only the 6 kept history rows, "
|
| 191 |
+
"the trajectory's velocity smoothing / force stop / acceleration as array operations in the same "
|
| 192 |
+
"float64 / float32 order, the denormalisation and predicted paths on the emitted rows only); bit-exact "
|
| 193 |
+
"(the host suite compares the arrays); 0 = the first port's *_ref functions"),
|
| 194 |
+
Knob("INPUT_TRIM", True, "per-request uploads trimmed (OPT round 5 item 4, H5): the solver's current states "
|
| 195 |
+
"``cs`` uploaded as one tile column [352, 32] (the 4 t = 0 columns + zeros) and widened to 324 columns in "
|
| 196 |
+
"the trace by a tile-aligned concat with a zero block (the same values), and an input whose array is all "
|
| 197 |
+
"+0.0 (``y0`` at temperature 0, ``static_x``) not uploaded again while its device buffer already holds "
|
| 198 |
+
"zeros from this model's last upload; 0 = every input every plan, ``cs`` at full width"),
|
| 199 |
+
Knob("LN_SFPU_BCAST", True, "with LN_KERNEL: the row statistics written to every column by the SFPU row "
|
| 200 |
+
"reduce itself (kernels/ln32_sfpu.h, the stock reduce's arithmetic with extra stores: bit-identical) "
|
| 201 |
+
"instead of a RISC-V column fill of the packed column-0 tile; 0 = the fill"),
|
| 202 |
])
|
| 203 |
|
| 204 |
|
code/tt_diffusion_planner/tt/decoder.py
CHANGED
|
@@ -32,7 +32,8 @@ from ..reference import config as C
|
|
| 32 |
from ..ttaw.ops import attention as A
|
| 33 |
from . import config as T
|
| 34 |
from . import params as P
|
| 35 |
-
from .
|
|
|
|
| 36 |
|
| 37 |
__all__ = ["TtDecoder", "TtTurnHead", "STEP_KEYS"]
|
| 38 |
|
|
@@ -42,7 +43,10 @@ STEP_KEYS = ("n1_g", "n1_b", "gate_msa", "n2_g", "n2_b", "gate_mlp")
|
|
| 42 |
class TtDecoder:
|
| 43 |
"""DiT weights, the per-step tables as device rows, and the evaluation / solver graph."""
|
| 44 |
|
| 45 |
-
def __init__(self, build: Build, p: Mapping[str, np.ndarray], tables: P.StepTables
|
|
|
|
|
|
|
|
|
|
| 46 |
self.build = build
|
| 47 |
self.tables = tables
|
| 48 |
D = "dec.block"
|
|
@@ -53,7 +57,8 @@ class TtDecoder:
|
|
| 53 |
out=build.hidden("dec.preproj.fc1"), activation="gelu")
|
| 54 |
self.pre2 = make_linear(build, P.linear(p, "decoder.dit.preproj.fc2", bias=False), "dec.preproj.fc2",
|
| 55 |
out=self.stream)
|
| 56 |
-
self.
|
|
|
|
| 57 |
self.blocks = []
|
| 58 |
for i in range(C.DIT_DEPTH):
|
| 59 |
B = f"decoder.dit.blocks.{i}"
|
|
@@ -92,7 +97,21 @@ class TtDecoder:
|
|
| 92 |
self.attn_mm = {"self": build.attn_matmul("dec.self_attn"), "cross": build.attn_matmul("dec.cross_attn")}
|
| 93 |
if self.attn_mm["cross"]: # fp32 cross-attention over the 576 encoder rows, the 12 pad tokens masked
|
| 94 |
row = np.where(np.arange(T.TOKENS) < T.TOKENS_REAL, 0.0, -np.inf).astype(np.float32)
|
| 95 |
-
self.cross_mask = Const(build, np.broadcast_to(row, (
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
|
| 97 |
def upload_step(self, tables: P.StepTables, k: int) -> Dict[str, Any]:
|
| 98 |
b = self.build
|
|
@@ -110,20 +129,76 @@ class TtDecoder:
|
|
| 110 |
|
| 111 |
if not self.attn_mm["cross"] and int(enc.shape[-2]) != T.TOKENS_REAL:
|
| 112 |
enc = ttnn.slice(enc, [0, 0, 0, 0], [1, 1, T.TOKENS_REAL, C.HIDDEN_DIM])
|
| 113 |
-
|
| 114 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
|
| 116 |
def _attend(self, kind: str, q, k, v, mask):
|
| 117 |
"""Attention of ``kind`` ("self" / "cross") -> concatenated heads ``[1, 1, 352, 256]``."""
|
| 118 |
if self.attn_mm[kind]:
|
| 119 |
-
return A.merge_heads(
|
|
|
|
|
|
|
| 120 |
return A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32[kind])
|
| 121 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 122 |
def evaluate(self, x, rows: Mapping[str, Any], kv: Sequence[Tuple[Any, Any]], self_mask):
|
| 123 |
-
"""One decoder evaluation: ``x`` ``[1, 1,
|
|
|
|
| 124 |
import ttnn
|
| 125 |
|
| 126 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 127 |
for b, r, (kc, vc) in zip(self.blocks, rows["blocks"], kv):
|
| 128 |
q, k, v = A.split_qkv(b["qkv"](b["ln"](h, r["n1_g"], r["n1_b"])), C.NUM_HEADS)
|
| 129 |
a = b["out"](self._attend("self", q, k, v, self_mask))
|
|
@@ -131,7 +206,7 @@ class TtDecoder:
|
|
| 131 |
m = b["m1b"](b["m1a"](b["ln"](h, r["n2_g"], r["n2_b"])))
|
| 132 |
h = ttnn.add(h, ttnn.multiply(m, r["gate_mlp"]))
|
| 133 |
qc = A.split_heads(b["cq"](b["n3"](h)), C.NUM_HEADS)
|
| 134 |
-
cmask = self.cross_mask() if self.attn_mm["cross"] else None
|
| 135 |
h = ttnn.add(h, b["cout"](self._attend("cross", qc, kc, vc, cmask)))
|
| 136 |
h = ttnn.add(h, b["m2b"](b["m2a"](b["n4"](h))))
|
| 137 |
f = self.fin_ln(h, rows["g"], rows["b"])
|
|
|
|
| 32 |
from ..ttaw.ops import attention as A
|
| 33 |
from . import config as T
|
| 34 |
from . import params as P
|
| 35 |
+
from .attention import attention_matmul
|
| 36 |
+
from .layers import KcatOperand, operand_ktp, ATTN, Build, Const, LayerNorm, chain2, make_linear
|
| 37 |
|
| 38 |
__all__ = ["TtDecoder", "TtTurnHead", "STEP_KEYS"]
|
| 39 |
|
|
|
|
| 43 |
class TtDecoder:
|
| 44 |
"""DiT weights, the per-step tables as device rows, and the evaluation / solver graph."""
|
| 45 |
|
| 46 |
+
def __init__(self, build: Build, p: Mapping[str, np.ndarray], tables: P.StepTables,
|
| 47 |
+
rows: Sequence[int] = (T.AGENTS,)):
|
| 48 |
+
"""``rows``: the decoder row counts the graph runs on (``T.AGENTS`` and, with ``COMPACT``, the agent
|
| 49 |
+
buckets): the agent embedding rows and the cross-attention mask are uploaded once per count."""
|
| 50 |
self.build = build
|
| 51 |
self.tables = tables
|
| 52 |
D = "dec.block"
|
|
|
|
| 57 |
out=build.hidden("dec.preproj.fc1"), activation="gelu")
|
| 58 |
self.pre2 = make_linear(build, P.linear(p, "decoder.dit.preproj.fc2", bias=False), "dec.preproj.fc2",
|
| 59 |
out=self.stream)
|
| 60 |
+
self.rows = tuple(sorted({int(r) for r in rows} | {T.AGENTS}))
|
| 61 |
+
self.agent_rows = {r: Const(build, P.agent_rows(p, r), self.stream) for r in self.rows}
|
| 62 |
self.blocks = []
|
| 63 |
for i in range(C.DIT_DEPTH):
|
| 64 |
B = f"decoder.dit.blocks.{i}"
|
|
|
|
| 97 |
self.attn_mm = {"self": build.attn_matmul("dec.self_attn"), "cross": build.attn_matmul("dec.cross_attn")}
|
| 98 |
if self.attn_mm["cross"]: # fp32 cross-attention over the 576 encoder rows, the 12 pad tokens masked
|
| 99 |
row = np.where(np.arange(T.TOKENS) < T.TOKENS_REAL, 0.0, -np.inf).astype(np.float32)
|
| 100 |
+
self.cross_mask = {r: Const(build, np.broadcast_to(row, (r, T.TOKENS)), "bfloat16",
|
| 101 |
+
memory_config=build.attn_mem()) for r in self.rows}
|
| 102 |
+
if build.dec_mem() is not None: # DEC_L1: the blocks' intermediates in L1 (interleaved)
|
| 103 |
+
for blk in self.blocks:
|
| 104 |
+
for key in ("out", "m1a", "m1b", "cout", "m2a", "m2b"):
|
| 105 |
+
blk[key].out_mem = build.dec_mem()
|
| 106 |
+
for key in ("ln", "n3", "n4"):
|
| 107 |
+
blk[key].mem = build.dec_mem()
|
| 108 |
+
self.edge_mem = build.dec_mem() if build.dec_l1 >= 2 else None
|
| 109 |
+
if self.edge_mem is not None: # DEC_L1=2: the pre-projection and the final layer too
|
| 110 |
+
self.pre1.out_mem = self.pre2.out_mem = self.p1.out_mem = self.edge_mem
|
| 111 |
+
self.fin_ln.mem = self.p0.mem = self.p3.mem = self.edge_mem
|
| 112 |
+
if build.attn_mem() is not None: # ATTN_L1: the fused attention reads Q / K / V from L1
|
| 113 |
+
for blk in self.blocks:
|
| 114 |
+
blk["qkv"].out_mem = blk["cq"].out_mem = build.attn_mem()
|
| 115 |
|
| 116 |
def upload_step(self, tables: P.StepTables, k: int) -> Dict[str, Any]:
|
| 117 |
b = self.build
|
|
|
|
| 129 |
|
| 130 |
if not self.attn_mm["cross"] and int(enc.shape[-2]) != T.TOKENS_REAL:
|
| 131 |
enc = ttnn.slice(enc, [0, 0, 0, 0], [1, 1, T.TOKENS_REAL, C.HIDDEN_DIM])
|
| 132 |
+
kv = [(A.split_heads(k(enc), C.NUM_HEADS), A.split_heads(v(enc), C.NUM_HEADS))
|
| 133 |
+
for k, v in zip(self.k_lin, self.v_lin)]
|
| 134 |
+
mem = self.build.attn_mem()
|
| 135 |
+
if mem is not None and self.attn_mm["cross"]: # ATTN_L1: the hoisted heads L1-resident for 11 evaluations
|
| 136 |
+
kv = [(ttnn.to_memory_config(k, mem), ttnn.to_memory_config(v, mem)) for k, v in kv]
|
| 137 |
+
return kv
|
| 138 |
|
| 139 |
def _attend(self, kind: str, q, k, v, mask):
|
| 140 |
"""Attention of ``kind`` ("self" / "cross") -> concatenated heads ``[1, 1, 352, 256]``."""
|
| 141 |
if self.attn_mm[kind]:
|
| 142 |
+
return A.merge_heads(attention_matmul(q, k, v, scale=self.scale, attn_mask=mask,
|
| 143 |
+
mode=self.build.attn_fast, smask=self.build.attn_smask,
|
| 144 |
+
smsm=self.build.attn_smsm))
|
| 145 |
return A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32[kind])
|
| 146 |
|
| 147 |
+
def _fused(self, kind: str, q, k, v, mask, flat_heads, kcat: int = 0):
|
| 148 |
+
"""``ATTN_FUSED``: the attention of ``kind`` as one program (``tt/fattn_kernel.py``) reading Q / K / V in
|
| 149 |
+
place, returning the merged heads ``[1, 1, 352, 256]``; ``flat_heads`` = the head-block column offsets (in
|
| 150 |
+
heads) of Q, K, V inside one projection output (self), None for the cross-attention (Q flat, K / V the
|
| 151 |
+
hoisted heads). None when the fused program does not apply (then the stock chain runs)."""
|
| 152 |
+
if not (self.build.attn_fused and self.attn_mm[kind] and mask is not None and self.build.attn_smask
|
| 153 |
+
and self.build.attn_smsm and self.build.attn_fast == 1):
|
| 154 |
+
return None
|
| 155 |
+
from .fattn_kernel import flat_at, fused_attention, heads_at, supported
|
| 156 |
+
|
| 157 |
+
if not supported((q, k, v), mask):
|
| 158 |
+
return None
|
| 159 |
+
H = C.NUM_HEADS
|
| 160 |
+
if flat_heads is not None:
|
| 161 |
+
ats = [flat_at(t, o * H) for t, o in zip((q, k, v), flat_heads)]
|
| 162 |
+
else:
|
| 163 |
+
ats = [flat_at(q, 0), heads_at(k), heads_at(v)]
|
| 164 |
+
out = fused_attention(q, k, v, mask, self.scale, H, *ats, kcat_ktp=kcat, memory_config=self.build.dec_mem())
|
| 165 |
+
return KcatOperand(out, kcat) if kcat else out
|
| 166 |
+
|
| 167 |
def evaluate(self, x, rows: Mapping[str, Any], kv: Sequence[Tuple[Any, Any]], self_mask):
|
| 168 |
+
"""One decoder evaluation: ``x`` ``[1, 1, R, 324]`` fp32 (prefix-constrained; R = 352 or an agent bucket)
|
| 169 |
+
-> ``m * mask0`` fp32."""
|
| 170 |
import ttnn
|
| 171 |
|
| 172 |
+
R = int(tuple(x.shape)[-2])
|
| 173 |
+
|
| 174 |
+
h = ttnn.add(chain2(self.pre1, self.pre2, x, self.build.kcat_act), self.agent_rows[R](), # [1, 1, R, 256]
|
| 175 |
+
**({} if self.edge_mem is None else {"memory_config": self.edge_mem}))
|
| 176 |
+
if self.build.ln_resid: # LN_RESID: the stream adds fused into the LNs
|
| 177 |
+
pending = None
|
| 178 |
+
for b, r, (kc, vc) in zip(self.blocks, rows["blocks"], kv):
|
| 179 |
+
ek = {key: operand_ktp(self.build, b[key]) for key in ("qkv", "out", "m1a", "cq", "cout", "m2a")}
|
| 180 |
+
if pending is None:
|
| 181 |
+
n1 = b["ln"](h, r["n1_g"], r["n1_b"], kcat=ek["qkv"])
|
| 182 |
+
else:
|
| 183 |
+
h, n1 = b["ln"].residual(h, pending, None, r["n1_g"], r["n1_b"], kcat=ek["qkv"])
|
| 184 |
+
qkv = b["qkv"](n1)
|
| 185 |
+
a = self._fused("self", qkv, qkv, qkv, self_mask, (0, 1, 2), ek["out"])
|
| 186 |
+
if a is None:
|
| 187 |
+
q, k, v = A.split_qkv(qkv, C.NUM_HEADS)
|
| 188 |
+
a = self._attend("self", q, k, v, self_mask)
|
| 189 |
+
a = b["out"](a)
|
| 190 |
+
h, n2 = b["ln"].residual(h, a, r["gate_msa"], r["n2_g"], r["n2_b"], kcat=ek["m1a"])
|
| 191 |
+
m = chain2(b["m1a"], b["m1b"], n2, self.build.kcat_act)
|
| 192 |
+
h, n3 = b["n3"].residual(h, m, r["gate_mlp"], kcat=ek["cq"])
|
| 193 |
+
cq = b["cq"](n3)
|
| 194 |
+
cmask = self.cross_mask[R]() if self.attn_mm["cross"] else None
|
| 195 |
+
c = self._fused("cross", cq, kc, vc, cmask, None, ek["cout"])
|
| 196 |
+
if c is None:
|
| 197 |
+
c = self._attend("cross", A.split_heads(cq, C.NUM_HEADS), kc, vc, cmask)
|
| 198 |
+
h, n4 = b["n4"].residual(h, b["cout"](c), kcat=ek["m2a"])
|
| 199 |
+
pending = chain2(b["m2a"], b["m2b"], n4, self.build.kcat_act)
|
| 200 |
+
_, f = self.fin_ln.residual(h, pending, None, rows["g"], rows["b"], write_h=False)
|
| 201 |
+
return self.p4(self.p3(self.p1(self.p0(f))))
|
| 202 |
for b, r, (kc, vc) in zip(self.blocks, rows["blocks"], kv):
|
| 203 |
q, k, v = A.split_qkv(b["qkv"](b["ln"](h, r["n1_g"], r["n1_b"])), C.NUM_HEADS)
|
| 204 |
a = b["out"](self._attend("self", q, k, v, self_mask))
|
|
|
|
| 206 |
m = b["m1b"](b["m1a"](b["ln"](h, r["n2_g"], r["n2_b"])))
|
| 207 |
h = ttnn.add(h, ttnn.multiply(m, r["gate_mlp"]))
|
| 208 |
qc = A.split_heads(b["cq"](b["n3"](h)), C.NUM_HEADS)
|
| 209 |
+
cmask = self.cross_mask[R]() if self.attn_mm["cross"] else None
|
| 210 |
h = ttnn.add(h, b["cout"](self._attend("cross", qc, kc, vc, cmask)))
|
| 211 |
h = ttnn.add(h, b["m2b"](b["m2a"](b["n4"](h))))
|
| 212 |
f = self.fin_ln(h, rows["g"], rows["b"])
|
code/tt_diffusion_planner/tt/encoder.py
CHANGED
|
@@ -29,6 +29,7 @@ from ..reference import config as C
|
|
| 29 |
from ..ttaw.ops import attention as A
|
| 30 |
from . import config as T
|
| 31 |
from . import params as P
|
|
|
|
| 32 |
from .layers import ATTN, Build, Const, LayerNorm, make_linear
|
| 33 |
|
| 34 |
__all__ = ["MixerTrunk", "TtEncoder", "MIXER_CATS"]
|
|
@@ -48,6 +49,7 @@ class MixerTrunk:
|
|
| 48 |
|
| 49 |
def __init__(self, build: Build, p: Mapping[str, np.ndarray], cat: str):
|
| 50 |
self.cat, self.E, self.T = cat, ENTITIES[cat], T.MIXER_T[cat]
|
|
|
|
| 51 |
self.island = cat in T.ISLAND_ROWS
|
| 52 |
mods = P.mixer_module(p, cat)
|
| 53 |
mix = f"enc.mixer.{cat}"
|
|
@@ -83,18 +85,23 @@ class MixerTrunk:
|
|
| 83 |
"n2": LayerNorm(build, blk["n2"], mix),
|
| 84 |
"ch1": make_linear(build, blk["ch1"], mix, out=build.hidden(mix), activation="gelu"),
|
| 85 |
"ch2": make_linear(build, blk["ch2"], mix, out=self.stream)})
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 86 |
|
| 87 |
# ------------------------------------------------------------------------------------------------------------
|
| 88 |
def _tokens_view(self, x):
|
| 89 |
"""``[1, E, 128, 64]`` <-> ``[1, 1, E*128, 64]`` (free: 128 rows are whole tiles)."""
|
| 90 |
import ttnn
|
| 91 |
|
| 92 |
-
return ttnn.reshape(x, (1, 1,
|
| 93 |
|
| 94 |
def _entities_view(self, x):
|
| 95 |
import ttnn
|
| 96 |
|
| 97 |
-
return ttnn.reshape(x, (1,
|
| 98 |
|
| 99 |
def pre(self, x):
|
| 100 |
"""``[1, E, T, C]`` -> ``x0`` ``[1, E, 64, 128]`` (the ``enc.<cat>.pre`` tap)."""
|
|
@@ -129,6 +136,31 @@ class MixerTrunk:
|
|
| 129 |
"""The 6 MixerBlocks on ``[1, E, 64, 128]`` (the ``enc.<cat>.mixer`` tap)."""
|
| 130 |
import ttnn
|
| 131 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 132 |
for b in self.blocks:
|
| 133 |
y = self._tokens_view(ttnn.transpose(b["n1"](x), -2, -1)) # LN over the 128 channels, then T
|
| 134 |
y = ttnn.transpose(self._entities_view(b["tk2"](b["tk1"](y))), -2, -1)
|
|
@@ -141,7 +173,7 @@ class MixerTrunk:
|
|
| 141 |
import ttnn
|
| 142 |
|
| 143 |
m = ttnn.mean(x, dim=2, keepdim=True) # [1, E, 1, 128]
|
| 144 |
-
return ttnn.reshape(m, (1, 1,
|
| 145 |
|
| 146 |
|
| 147 |
class _Head:
|
|
@@ -194,9 +226,14 @@ class TtEncoder:
|
|
| 194 |
"polygon_x", "line_string_x", "goal_x", "ego_shape_x", "turn_x", "token_valid", "pos_aug",
|
| 195 |
"fusion_key_row")
|
| 196 |
|
| 197 |
-
def __init__(self, build: Build, p: Mapping[str, np.ndarray]):
|
|
|
|
|
|
|
| 198 |
self.build = build
|
| 199 |
self.fstream = build.stream("enc.fusion")
|
|
|
|
|
|
|
|
|
|
| 200 |
self.trunks = {cat: MixerTrunk(build, p, cat) for cat in MIXER_CATS}
|
| 201 |
self.heads = {cat: _Head(build, p, cat, self.fstream) for cat in MIXER_CATS}
|
| 202 |
st = "enc.head.static"
|
|
@@ -222,11 +259,39 @@ class TtEncoder:
|
|
| 222 |
"fc1": make_linear(build, P.linear(p, f"{B}.mlp.fc1"), fu, out=build.hidden(fu),
|
| 223 |
activation="gelu"),
|
| 224 |
"fc2": make_linear(build, P.linear(p, f"{B}.mlp.fc2"), fu, out=self.fstream)})
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 225 |
self.final_norm = LayerNorm(build, P.norm(p, "encoder.fusion.norm"), fu)
|
| 226 |
self.scale = float(C.ATTN_SCALE)
|
| 227 |
self.attn_fp32 = build.attn_fp32_acc("enc.fusion.attn")
|
| 228 |
|
| 229 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 230 |
import ttnn
|
| 231 |
|
| 232 |
def tap(name, t):
|
|
@@ -237,10 +302,17 @@ class TtEncoder:
|
|
| 237 |
out: Dict[str, Any] = {}
|
| 238 |
for cat in MIXER_CATS:
|
| 239 |
tr = self.trunks[cat]
|
| 240 |
-
|
| 241 |
-
x = tap(f"enc.{cat}.mixer", tr.mix(x0))
|
| 242 |
aux = ctx.get(f"{cat}_aux") if cat in ("neighbor", "lane", "route") else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 243 |
out[cat] = self.heads[cat](tr.pool(x), aux)
|
|
|
|
|
|
|
| 244 |
out["static"] = self.static2(self.static1(ctx["static_x"]))
|
| 245 |
for cat, enc in self.small.items():
|
| 246 |
out[cat] = enc(ctx[f"{cat}_x"])
|
|
@@ -249,15 +321,24 @@ class TtEncoder:
|
|
| 249 |
x = ttnn.concat([out[name] for name, _ in C.TOKEN_LAYOUT] + [self.pad_tokens()], dim=2) # [1,1,576,256]
|
| 250 |
x = ttnn.multiply(x, ctx["token_valid"]) # invalid entities -> 0
|
| 251 |
x = tap("enc.tokens", ttnn.add(x, self.pos(ctx["pos_aug"]))) # + valid * (pos W + b)
|
| 252 |
-
mask = A.expand_key_bias(ctx["fusion_key_row"], T.TOKENS
|
|
|
|
| 253 |
for i, b in enumerate(self.blocks):
|
| 254 |
kv = b["kv"](x) # K | V from the un-normalised x
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 258 |
else:
|
|
|
|
| 259 |
a = A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32)
|
| 260 |
-
|
| 261 |
-
x = ttnn.add(x, b["
|
|
|
|
| 262 |
tap(f"enc.fusion.{i}", x)
|
| 263 |
return tap("enc.encoding", self.final_norm(x))
|
|
|
|
| 29 |
from ..ttaw.ops import attention as A
|
| 30 |
from . import config as T
|
| 31 |
from . import params as P
|
| 32 |
+
from .attention import attention_matmul
|
| 33 |
from .layers import ATTN, Build, Const, LayerNorm, make_linear
|
| 34 |
|
| 35 |
__all__ = ["MixerTrunk", "TtEncoder", "MIXER_CATS"]
|
|
|
|
| 49 |
|
| 50 |
def __init__(self, build: Build, p: Mapping[str, np.ndarray], cat: str):
|
| 51 |
self.cat, self.E, self.T = cat, ENTITIES[cat], T.MIXER_T[cat]
|
| 52 |
+
self.build = build
|
| 53 |
self.island = cat in T.ISLAND_ROWS
|
| 54 |
mods = P.mixer_module(p, cat)
|
| 55 |
mix = f"enc.mixer.{cat}"
|
|
|
|
| 85 |
"n2": LayerNorm(build, blk["n2"], mix),
|
| 86 |
"ch1": make_linear(build, blk["ch1"], mix, out=build.hidden(mix), activation="gelu"),
|
| 87 |
"ch2": make_linear(build, blk["ch2"], mix, out=self.stream)})
|
| 88 |
+
if build.enc_mem() is not None: # ENC_L1: the mixer blocks' intermediates in L1 (interleaved)
|
| 89 |
+
for b in self.blocks:
|
| 90 |
+
b["n1"].mem = b["n2"].mem = build.enc_mem()
|
| 91 |
+
for key in ("tk1", "tk2", "ch1", "ch2"):
|
| 92 |
+
b[key].out_mem = build.enc_mem()
|
| 93 |
|
| 94 |
# ------------------------------------------------------------------------------------------------------------
|
| 95 |
def _tokens_view(self, x):
|
| 96 |
"""``[1, E, 128, 64]`` <-> ``[1, 1, E*128, 64]`` (free: 128 rows are whole tiles)."""
|
| 97 |
import ttnn
|
| 98 |
|
| 99 |
+
return ttnn.reshape(x, (1, 1, int(x.shape[1]) * C.MIXER_CHANNELS, x.shape[-1]))
|
| 100 |
|
| 101 |
def _entities_view(self, x):
|
| 102 |
import ttnn
|
| 103 |
|
| 104 |
+
return ttnn.reshape(x, (1, int(x.shape[2]) // C.MIXER_CHANNELS, C.MIXER_CHANNELS, x.shape[-1]))
|
| 105 |
|
| 106 |
def pre(self, x):
|
| 107 |
"""``[1, E, T, C]`` -> ``x0`` ``[1, E, 64, 128]`` (the ``enc.<cat>.pre`` tap)."""
|
|
|
|
| 136 |
"""The 6 MixerBlocks on ``[1, E, 64, 128]`` (the ``enc.<cat>.mixer`` tap)."""
|
| 137 |
import ttnn
|
| 138 |
|
| 139 |
+
if self.build.ln_resid and self.blocks[0]["n1"].can_transpose(x):
|
| 140 |
+
# LN_TR: the transposes around the token-mixing MLP inside the LN programs (n1 writes LN(h)^T, n2 reads
|
| 141 |
+
# the token-mixing output transposed): the same values, two programs less per block
|
| 142 |
+
pending = None
|
| 143 |
+
for b in self.blocks:
|
| 144 |
+
if pending is None:
|
| 145 |
+
n1t = b["n1"].transposed(x)
|
| 146 |
+
else:
|
| 147 |
+
x, n1t = b["n1"].transposed(x, res=pending)
|
| 148 |
+
t = self._entities_view(b["tk2"](b["tk1"](self._tokens_view(n1t)))) # [1, E, 128, 64]
|
| 149 |
+
x, n2 = b["n2"].residual_t(x, t)
|
| 150 |
+
pending = b["ch2"](b["ch1"](n2))
|
| 151 |
+
return ttnn.add(x, pending, memory_config=ttnn.DRAM_MEMORY_CONFIG)
|
| 152 |
+
if self.build.ln_resid: # LN_RESID: adds fused into the LNs
|
| 153 |
+
pending = None
|
| 154 |
+
for b in self.blocks:
|
| 155 |
+
if pending is None:
|
| 156 |
+
n1 = b["n1"](x)
|
| 157 |
+
else:
|
| 158 |
+
x, n1 = b["n1"].residual(x, pending)
|
| 159 |
+
y = self._tokens_view(ttnn.transpose(n1, -2, -1))
|
| 160 |
+
y = ttnn.transpose(self._entities_view(b["tk2"](b["tk1"](y))), -2, -1)
|
| 161 |
+
x, n2 = b["n2"].residual(x, y)
|
| 162 |
+
pending = b["ch2"](b["ch1"](n2))
|
| 163 |
+
return ttnn.add(x, pending)
|
| 164 |
for b in self.blocks:
|
| 165 |
y = self._tokens_view(ttnn.transpose(b["n1"](x), -2, -1)) # LN over the 128 channels, then T
|
| 166 |
y = ttnn.transpose(self._entities_view(b["tk2"](b["tk1"](y))), -2, -1)
|
|
|
|
| 173 |
import ttnn
|
| 174 |
|
| 175 |
m = ttnn.mean(x, dim=2, keepdim=True) # [1, E, 1, 128]
|
| 176 |
+
return ttnn.reshape(m, (1, 1, int(x.shape[1]), C.MIXER_CHANNELS))
|
| 177 |
|
| 178 |
|
| 179 |
class _Head:
|
|
|
|
| 226 |
"polygon_x", "line_string_x", "goal_x", "ego_shape_x", "turn_x", "token_valid", "pos_aug",
|
| 227 |
"fusion_key_row")
|
| 228 |
|
| 229 |
+
def __init__(self, build: Build, p: Mapping[str, np.ndarray], nb_rows=()):
|
| 230 |
+
"""``nb_rows`` (``COMPACT``): the neighbour entity counts the trunk may run on (the agent buckets): a zero
|
| 231 |
+
token block for the rest is uploaded per count."""
|
| 232 |
self.build = build
|
| 233 |
self.fstream = build.stream("enc.fusion")
|
| 234 |
+
E = ENTITIES["neighbor"]
|
| 235 |
+
self.nb_zeros = {int(n): Const(build, np.zeros((E - int(n), C.HIDDEN_DIM), np.float32), self.fstream)
|
| 236 |
+
for n in nb_rows if 0 < int(n) < E}
|
| 237 |
self.trunks = {cat: MixerTrunk(build, p, cat) for cat in MIXER_CATS}
|
| 238 |
self.heads = {cat: _Head(build, p, cat, self.fstream) for cat in MIXER_CATS}
|
| 239 |
st = "enc.head.static"
|
|
|
|
| 259 |
"fc1": make_linear(build, P.linear(p, f"{B}.mlp.fc1"), fu, out=build.hidden(fu),
|
| 260 |
activation="gelu"),
|
| 261 |
"fc2": make_linear(build, P.linear(p, f"{B}.mlp.fc2"), fu, out=self.fstream)})
|
| 262 |
+
if self.attn_mm and build.attn_mem() is not None and build.attn_l1 == 1: # ATTN_L1=1: Q, K | V to L1
|
| 263 |
+
for blk in self.blocks:
|
| 264 |
+
blk["q"].out_mem = blk["kv"].out_mem = build.attn_mem()
|
| 265 |
+
self.fus_mem = None
|
| 266 |
+
if build.fus_l1: # FUS_L1: the fusion blocks' intermediates in L1 (interleaved)
|
| 267 |
+
import ttnn
|
| 268 |
+
|
| 269 |
+
self.fus_mem = ttnn.L1_MEMORY_CONFIG
|
| 270 |
+
for blk in self.blocks:
|
| 271 |
+
blk["n1"].mem = blk["n2"].mem = self.fus_mem
|
| 272 |
+
for key in ("out", "fc1", "fc2"):
|
| 273 |
+
blk[key].out_mem = self.fus_mem
|
| 274 |
self.final_norm = LayerNorm(build, P.norm(p, "encoder.fusion.norm"), fu)
|
| 275 |
self.scale = float(C.ATTN_SCALE)
|
| 276 |
self.attn_fp32 = build.attn_fp32_acc("enc.fusion.attn")
|
| 277 |
|
| 278 |
+
def _fused(self, qs, kv, mask):
|
| 279 |
+
"""``ATTN_FUSED``: the fusion attention as one program (``tt/fattn_kernel.py``; Q from ``[1, 1, 576, 256]``,
|
| 280 |
+
K | V from ``[1, 1, 576, 512]`` in place) -> merged heads, or None (the stock chain runs)."""
|
| 281 |
+
b = self.build
|
| 282 |
+
if not (b.attn_fused and self.attn_mm and b.attn_smask and b.attn_smsm and b.attn_fast == 1):
|
| 283 |
+
return None
|
| 284 |
+
from .fattn_kernel import flat_at, fused_attention, supported
|
| 285 |
+
|
| 286 |
+
if not supported((qs, kv), mask):
|
| 287 |
+
return None
|
| 288 |
+
H = C.NUM_HEADS
|
| 289 |
+
return fused_attention(qs, kv, kv, mask, self.scale, H, flat_at(qs, 0), flat_at(kv, 0), flat_at(kv, H),
|
| 290 |
+
memory_config=self.fus_mem)
|
| 291 |
+
|
| 292 |
+
def forward(self, ctx: Mapping[str, Any], taps: Optional[Dict[str, Any]] = None, nb: Optional[int] = None):
|
| 293 |
+
"""``nb`` (``COMPACT``): run the neighbour trunk and head on the first ``nb`` entities only and fill the other
|
| 294 |
+
neighbour tokens with zeros: they are invalid entities, which ``token_valid`` zeroes anyway."""
|
| 295 |
import ttnn
|
| 296 |
|
| 297 |
def tap(name, t):
|
|
|
|
| 302 |
out: Dict[str, Any] = {}
|
| 303 |
for cat in MIXER_CATS:
|
| 304 |
tr = self.trunks[cat]
|
| 305 |
+
xin = ctx[f"{cat}_x"]
|
|
|
|
| 306 |
aux = ctx.get(f"{cat}_aux") if cat in ("neighbor", "lane", "route") else None
|
| 307 |
+
cut = cat == "neighbor" and nb in self.nb_zeros
|
| 308 |
+
if cut: # COMPACT: the first nb neighbours
|
| 309 |
+
xin = ttnn.slice(xin, [0, 0, 0, 0], [1, nb] + list(xin.shape)[2:])
|
| 310 |
+
aux = ttnn.slice(aux, [0, 0, 0, 0], [1, 1, nb, int(aux.shape[-1])])
|
| 311 |
+
x0 = tap(f"enc.{cat}.pre", tr.pre(xin))
|
| 312 |
+
x = tap(f"enc.{cat}.mixer", tr.mix(x0))
|
| 313 |
out[cat] = self.heads[cat](tr.pool(x), aux)
|
| 314 |
+
if cut:
|
| 315 |
+
out[cat] = ttnn.concat([out[cat], self.nb_zeros[nb]()], dim=2)
|
| 316 |
out["static"] = self.static2(self.static1(ctx["static_x"]))
|
| 317 |
for cat, enc in self.small.items():
|
| 318 |
out[cat] = enc(ctx[f"{cat}_x"])
|
|
|
|
| 321 |
x = ttnn.concat([out[name] for name, _ in C.TOKEN_LAYOUT] + [self.pad_tokens()], dim=2) # [1,1,576,256]
|
| 322 |
x = ttnn.multiply(x, ctx["token_valid"]) # invalid entities -> 0
|
| 323 |
x = tap("enc.tokens", ttnn.add(x, self.pos(ctx["pos_aug"]))) # + valid * (pos W + b)
|
| 324 |
+
mask = A.expand_key_bias(ctx["fusion_key_row"], T.TOKENS, # [1, 1, 576, 576], once per plan
|
| 325 |
+
memory_config=self.build.attn_mem() if self.attn_mm else None)
|
| 326 |
for i, b in enumerate(self.blocks):
|
| 327 |
kv = b["kv"](x) # K | V from the un-normalised x
|
| 328 |
+
qs = b["q"](b["n1"](x))
|
| 329 |
+
a = self._fused(qs, kv, mask)
|
| 330 |
+
if a is not None:
|
| 331 |
+
pass
|
| 332 |
+
elif self.attn_mm:
|
| 333 |
+
q, k, v = A.split_q_kv(qs, kv, C.NUM_HEADS)
|
| 334 |
+
a = A.merge_heads(attention_matmul(q, k, v, scale=self.scale, attn_mask=mask,
|
| 335 |
+
mode=self.build.attn_fast, smask=self.build.attn_smask,
|
| 336 |
+
smsm=self.build.attn_smsm))
|
| 337 |
else:
|
| 338 |
+
q, k, v = A.split_q_kv(qs, kv, C.NUM_HEADS)
|
| 339 |
a = A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32)
|
| 340 |
+
fkw = {} if self.fus_mem is None else {"memory_config": self.fus_mem}
|
| 341 |
+
x = ttnn.add(x, b["out"](a), **fkw)
|
| 342 |
+
x = ttnn.add(x, b["fc2"](b["fc1"](b["n2"](x))), **fkw)
|
| 343 |
tap(f"enc.fusion.{i}", x)
|
| 344 |
return tap("enc.encoding", self.final_norm(x))
|
code/tt_diffusion_planner/tt/fattn_kernel.py
ADDED
|
@@ -0,0 +1,140 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""Fused fp32 matmul attention as one ``ttnn.generic_op`` (``ATTN_FUSED``, OPT round 3 item 4).
|
| 3 |
+
|
| 4 |
+
``fused_attention(q, k, v, mask, scale, heads, q_at, k_at, v_at)`` -> the merged-heads output ``[1, 1, Sq, H * 32]``
|
| 5 |
+
fp32: ``softmax(scale * Q K^T + mask) V`` per head with head dim 32 (one tile), i.e. the programs
|
| 6 |
+
``nlp_create_qkv_heads`` (or ``split_heads``), ``Q K^T`` matmul, the scale + mask + softmax (``ATTN_SMSM``), the
|
| 7 |
+
``P V`` matmul and ``nlp_concat_heads`` of ``tt/attention.py`` in one program (``kernels/fattn_*.cpp``). Each unit
|
| 8 |
+
is one head x one query tile row: ``Q K^T`` into DEST two key tiles at a time, the smask SFPU sequence, the stock
|
| 9 |
+
softmax ``kernel_lib`` calls on the row in L1, then ``P V`` accumulated in DEST in key order, so the scores and the
|
| 10 |
+
probabilities never leave L1 and the math is that of the stock programs (meant to be bit-identical: the device
|
| 11 |
+
check compares with them).
|
| 12 |
+
|
| 13 |
+
Q / K / V are read in place from their producers: ``q_at`` / ``k_at`` / ``v_at`` = ``(base, row_stride,
|
| 14 |
+
head_stride)`` in tiles, e.g. the self-attention ``qkv`` projection ``[1, 1, S, 768]`` gives Q at ``(0, 24, 1)``,
|
| 15 |
+
K at ``(8, 24, 1)``, V at ``(16, 24, 1)``; a ``[1, H, S, 32]`` heads tensor is ``(0, 1, S / 32)``.
|
| 16 |
+
"""
|
| 17 |
+
from __future__ import annotations
|
| 18 |
+
|
| 19 |
+
import os
|
| 20 |
+
import struct
|
| 21 |
+
from typing import Any, Sequence
|
| 22 |
+
|
| 23 |
+
__all__ = ["fused_attention", "supported", "flat_at", "heads_at"]
|
| 24 |
+
|
| 25 |
+
_KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
|
| 26 |
+
TB = 4096
|
| 27 |
+
N_RT = 2
|
| 28 |
+
NDST = 4
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def flat_at(t: Any, col_tile: int):
|
| 32 |
+
"""``(base, row_stride, head_stride)`` of heads stored as consecutive 32-column tiles of ``[1, 1, S, W]``,
|
| 33 |
+
the first head at column tile ``col_tile``."""
|
| 34 |
+
return (int(col_tile), int(t.padded_shape[-1]) // 32, 1)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def heads_at(t: Any):
|
| 38 |
+
"""``(base, row_stride, head_stride)`` of a ``[1, H, S, 32]`` heads tensor."""
|
| 39 |
+
return (0, 1, int(t.padded_shape[-2]) // 32)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def supported(tensors: Sequence[Any], mask: Any) -> bool:
|
| 43 |
+
import ttnn
|
| 44 |
+
|
| 45 |
+
try:
|
| 46 |
+
if not hasattr(ttnn, "generic_op") or mask is None:
|
| 47 |
+
return False
|
| 48 |
+
for t in tensors:
|
| 49 |
+
if t.dtype != ttnn.float32 or t.layout != ttnn.TILE_LAYOUT or t.is_sharded():
|
| 50 |
+
return False
|
| 51 |
+
return (mask.layout == ttnn.TILE_LAYOUT and mask.dtype in (ttnn.bfloat16, ttnn.float32)
|
| 52 |
+
and not mask.is_sharded() and int(mask.padded_shape[1]) == 1)
|
| 53 |
+
except Exception: # noqa: BLE001 - the host fake ttnn
|
| 54 |
+
return False
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def fused_attention(q: Any, k: Any, v: Any, mask: Any, scale: float, heads: int, q_at, k_at, v_at,
|
| 58 |
+
kcat_ktp: int = 0, memory_config: Any = None):
|
| 59 |
+
"""``kcat_ktp`` (``KCAT_EMIT``): write the split operand ``[o_hi | o_hi | o_lo | 1 | 0..]`` of the next
|
| 60 |
+
K-concatenated linear (``[1, 1, Sq, 32 * kcat_ktp]``, the ``kcat_operand`` layout) instead of the output."""
|
| 61 |
+
import ttnn
|
| 62 |
+
|
| 63 |
+
dev = q.device()
|
| 64 |
+
Sq, Sk = int(mask.padded_shape[-2]), int(mask.padded_shape[-1])
|
| 65 |
+
Mt, Wt = Sq // 32, Sk // 32
|
| 66 |
+
units = heads * Mt
|
| 67 |
+
Wp = (Wt + 1) // 2
|
| 68 |
+
rb = -(-Wt // NDST) * NDST
|
| 69 |
+
p_pad = rb - Wt
|
| 70 |
+
ow = 32 * kcat_ktp if kcat_ktp else heads * 32
|
| 71 |
+
out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, Sq, ow]), ttnn.float32, ttnn.TILE_LAYOUT, dev,
|
| 72 |
+
memory_config or ttnn.DRAM_MEMORY_CONFIG)
|
| 73 |
+
g = dev.compute_with_storage_grid_size()
|
| 74 |
+
n = min(units, g.x * g.y)
|
| 75 |
+
cs = [(i % g.x, i // g.x) for i in range(n)]
|
| 76 |
+
full, rem = divmod(n, g.x)
|
| 77 |
+
rs = []
|
| 78 |
+
if full:
|
| 79 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(g.x - 1, full - 1)))
|
| 80 |
+
if rem:
|
| 81 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, full), ttnn.CoreCoord(rem - 1, full)))
|
| 82 |
+
crs = ttnn.CoreRangeSet(set(rs))
|
| 83 |
+
base, extra = divmod(units, n)
|
| 84 |
+
rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
|
| 85 |
+
u0 = 0
|
| 86 |
+
for i, (cx, cy) in enumerate(cs):
|
| 87 |
+
kk = base + (1 if i < extra else 0)
|
| 88 |
+
rd[cx][cy] = [u0, kk]
|
| 89 |
+
wr[cx][cy] = [u0, kk]
|
| 90 |
+
cp[cx][cy] = [kk, 0]
|
| 91 |
+
u0 += kk
|
| 92 |
+
mbf = mask.dtype == ttnn.bfloat16
|
| 93 |
+
tbm = 2048 if mbf else 4096
|
| 94 |
+
|
| 95 |
+
def cb(idx, pages, dt, tb):
|
| 96 |
+
return ttnn.CBDescriptor(total_size=pages * tb, core_ranges=crs, format_descriptors=[
|
| 97 |
+
ttnn.CBFormatDescriptor(buffer_index=idx, data_format=dt, page_size=tb)])
|
| 98 |
+
|
| 99 |
+
f32 = ttnn.float32
|
| 100 |
+
cbs = [cb(0, 2, f32, TB), cb(1, 2 * Wp, mask.dtype, tbm), cb(2, 4, f32, TB), cb(3, 1, f32, TB),
|
| 101 |
+
cb(4, 1, f32, TB), cb(5, Wt, f32, TB), cb(16, 2, f32, TB), cb(24, Wt, f32, TB), cb(25, 1, f32, TB),
|
| 102 |
+
cb(26, Wt, f32, TB), cb(27, 1, f32, TB), cb(28, rb, f32, TB)]
|
| 103 |
+
if kcat_ktp:
|
| 104 |
+
cbs += [cb(17, 2, f32, TB), cb(18, 1, f32, TB), cb(29, 1, f32, TB)]
|
| 105 |
+
if kcat_ktp > 3 * heads + 1:
|
| 106 |
+
cbs.append(cb(19, 1, f32, TB))
|
| 107 |
+
um = [ttnn.UnpackToDestMode.Default] * 64
|
| 108 |
+
if kcat_ktp:
|
| 109 |
+
um[29] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 110 |
+
if not mbf:
|
| 111 |
+
um[1] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 112 |
+
ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True,
|
| 113 |
+
math_approx_mode=False)
|
| 114 |
+
ccfg.unpack_to_dest_mode = um
|
| 115 |
+
bits = int.from_bytes(struct.pack("<f", float(scale)), "little")
|
| 116 |
+
SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH
|
| 117 |
+
|
| 118 |
+
def acc(t):
|
| 119 |
+
return list(ttnn.TensorAccessorArgs(t).get_compile_time_args())
|
| 120 |
+
|
| 121 |
+
def kd(name, ct, rt, common, config):
|
| 122 |
+
return ttnn.KernelDescriptor(kernel_source=os.path.join(_KDIR, name), source_type=SRC, core_ranges=crs,
|
| 123 |
+
compile_time_args=ct, defines=[], runtime_args=rt, common_runtime_args=common,
|
| 124 |
+
config=config)
|
| 125 |
+
|
| 126 |
+
ats = [int(x) for a in (q_at, k_at, v_at) for x in a]
|
| 127 |
+
ks = [kd("fattn_reader.cpp", [Wt, Mt, N_RT, tbm] + ats + acc(q) + acc(k) + acc(v) + acc(mask), rd,
|
| 128 |
+
[q.buffer_address(), k.buffer_address(), v.buffer_address(), mask.buffer_address()],
|
| 129 |
+
ttnn.ReaderConfigDescriptor()),
|
| 130 |
+
kd("fattn_writer.cpp", [Mt, heads, N_RT, int(bool(kcat_ktp)), int(kcat_ktp)] + acc(out), wr,
|
| 131 |
+
[out.buffer_address()],
|
| 132 |
+
ttnn.WriterConfigDescriptor()),
|
| 133 |
+
kd("fattn_compute.cpp", [Wt, N_RT, int(mbf), NDST, p_pad, bits, int(bool(kcat_ktp))], cp, [], ccfg)]
|
| 134 |
+
ins = []
|
| 135 |
+
for t in (q, k, v):
|
| 136 |
+
if all(t is not x for x in ins):
|
| 137 |
+
ins.append(t)
|
| 138 |
+
ins += [mask, out]
|
| 139 |
+
ttnn.generic_op(ins, ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))
|
| 140 |
+
return out
|
code/tt_diffusion_planner/tt/inputs.py
CHANGED
|
@@ -18,10 +18,13 @@ from ..ttaw.ops.attention import key_bias_row
|
|
| 18 |
from . import config as T
|
| 19 |
from .params import lane_aux_features, pad_rows
|
| 20 |
|
| 21 |
-
__all__ = ["plan_inputs", "INPUT_SPECS", "warmup_inputs", "decoder_state", "BF16_INPUTS"]
|
| 22 |
|
| 23 |
f32 = np.float32
|
| 24 |
|
|
|
|
|
|
|
|
|
|
| 25 |
# name -> device shape (all fp32 TILE except the bf16 key-bias rows)
|
| 26 |
INPUT_SPECS: Dict[str, tuple] = {
|
| 27 |
"ego_x": (1, 1, T.MIXER_T["ego"], T.MIXER_CIN["ego"]),
|
|
@@ -41,7 +44,7 @@ INPUT_SPECS: Dict[str, tuple] = {
|
|
| 41 |
"pos_aug": (1, 1, T.TOKENS, T.POS_AUG_DIM),
|
| 42 |
"fusion_key_row": (1, 1, 1, T.TOKENS),
|
| 43 |
"agent_key_row": (1, 1, 1, T.AGENTS),
|
| 44 |
-
"cs": (1, 1, T.AGENTS,
|
| 45 |
"y0": (1, 1, T.AGENTS, T.STATE_COLS),
|
| 46 |
}
|
| 47 |
BF16_INPUTS = ("fusion_key_row", "agent_key_row")
|
|
@@ -84,7 +87,7 @@ def plan_inputs(prep: Prepared) -> Dict[str, np.ndarray]:
|
|
| 84 |
"fusion_key_row": key_bias_row(pad_rows(np.asarray(f.key_valid, bool), T.TOKENS)),
|
| 85 |
"agent_key_row": key_bias_row(pad_rows(np.asarray(prep.decoder.agent_valid, bool), T.AGENTS)),
|
| 86 |
"cs": decoder_state(np.zeros((C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM), f32),
|
| 87 |
-
prep.decoder.current_states),
|
| 88 |
"y0": decoder_state(prep.x_T),
|
| 89 |
}
|
| 90 |
for name, arr in out.items():
|
|
|
|
| 18 |
from . import config as T
|
| 19 |
from .params import lane_aux_features, pad_rows
|
| 20 |
|
| 21 |
+
__all__ = ["plan_inputs", "INPUT_SPECS", "warmup_inputs", "decoder_state", "BF16_INPUTS", "CS_COLS"]
|
| 22 |
|
| 23 |
f32 = np.float32
|
| 24 |
|
| 25 |
+
# INPUT_TRIM: the current states as one tile column (widened to STATE_COLS in the trace)
|
| 26 |
+
CS_COLS = T.TILE if T.KNOBS.read().INPUT_TRIM else T.STATE_COLS
|
| 27 |
+
|
| 28 |
# name -> device shape (all fp32 TILE except the bf16 key-bias rows)
|
| 29 |
INPUT_SPECS: Dict[str, tuple] = {
|
| 30 |
"ego_x": (1, 1, T.MIXER_T["ego"], T.MIXER_CIN["ego"]),
|
|
|
|
| 44 |
"pos_aug": (1, 1, T.TOKENS, T.POS_AUG_DIM),
|
| 45 |
"fusion_key_row": (1, 1, 1, T.TOKENS),
|
| 46 |
"agent_key_row": (1, 1, 1, T.AGENTS),
|
| 47 |
+
"cs": (1, 1, T.AGENTS, CS_COLS),
|
| 48 |
"y0": (1, 1, T.AGENTS, T.STATE_COLS),
|
| 49 |
}
|
| 50 |
BF16_INPUTS = ("fusion_key_row", "agent_key_row")
|
|
|
|
| 87 |
"fusion_key_row": key_bias_row(pad_rows(np.asarray(f.key_valid, bool), T.TOKENS)),
|
| 88 |
"agent_key_row": key_bias_row(pad_rows(np.asarray(prep.decoder.agent_valid, bool), T.AGENTS)),
|
| 89 |
"cs": decoder_state(np.zeros((C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM), f32),
|
| 90 |
+
prep.decoder.current_states)[..., :CS_COLS],
|
| 91 |
"y0": decoder_state(prep.x_T),
|
| 92 |
}
|
| 93 |
for name, arr in out.items():
|
code/tt_diffusion_planner/tt/kcat_kernel.py
ADDED
|
@@ -0,0 +1,109 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""Operand build of the K-concatenated split matmul (``SPLIT_KCAT=2``, OPT round 2 item 2) as one ``generic_op``.
|
| 3 |
+
|
| 4 |
+
``kcat_operand(x)``: fp32 TILE ``[..., M, K]`` (K a multiple of 32) -> fp32 ``[..., M, 3K + 32]`` =
|
| 5 |
+
``[x_hi | x_hi | x_lo | ones]`` with ``x_hi`` = bf16(x) (the ``ttnn.typecast`` LLK) and ``x_lo = x - x_hi``; the
|
| 6 |
+
last tile has columns 0 and 1 = 1 (the exact bias rows of the concatenated weight). It replaces the 5 stock programs
|
| 7 |
+
of ``SPLIT_KCAT=1`` (typecast, typecast, subtract, concat over 4 tensors). x tiles are split in contiguous ranges
|
| 8 |
+
over the grid (``kernels/kcat_*.cpp``).
|
| 9 |
+
"""
|
| 10 |
+
from __future__ import annotations
|
| 11 |
+
|
| 12 |
+
import os
|
| 13 |
+
from typing import Any
|
| 14 |
+
|
| 15 |
+
__all__ = ["kcat_operand", "kcat_tiles", "supported"]
|
| 16 |
+
|
| 17 |
+
_KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
|
| 18 |
+
TB = 4096
|
| 19 |
+
N_RT = 2
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def supported(x: Any) -> bool:
|
| 23 |
+
"""True for a device fp32 TILE interleaved tensor with tile-aligned rows and columns (not the host fake)."""
|
| 24 |
+
import ttnn
|
| 25 |
+
|
| 26 |
+
try:
|
| 27 |
+
shp = list(x.padded_shape)
|
| 28 |
+
return (x.dtype == ttnn.float32 and x.layout == ttnn.TILE_LAYOUT and not x.is_sharded()
|
| 29 |
+
and shp[-1] % 32 == 0 and shp[-2] % 32 == 0 and hasattr(ttnn, "generic_op"))
|
| 30 |
+
except Exception: # noqa: BLE001 - the host fake ttnn
|
| 31 |
+
return False
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def kcat_tiles(k: int, pad: int) -> int:
|
| 35 |
+
"""Row tiles of X' for an input of K = ``k`` columns: 3 Kt + 1, rounded up to a multiple of ``pad``."""
|
| 36 |
+
return -(-(3 * (k // 32) + 1) // pad) * pad
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
ACTS = {None: 0, "gelu": 1, "gelu_tanh": 2}
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def kcat_operand(x: Any, pad: int = 1, act: Any = None, memory_config: Any = None, act_once: bool = False):
|
| 43 |
+
"""``pad``: X' row tiles rounded up to a multiple of ``pad`` with zero tiles (the matmul's K block must divide
|
| 44 |
+
them; 3 Kt + 1 is odd, e.g. 97 for K = 1024). ``act`` (``KCAT_ACT``): ``"gelu"`` / ``"gelu_tanh"`` applied to
|
| 45 |
+
x first (the previous linear's activation, the same LLK as the stock ``ttnn.gelu`` program). ``memory_config``:
|
| 46 |
+
of X' (default DRAM interleaved; ``KCAT_L1``: the L1 block-sharded layout the consumer matmul reads in place, the
|
| 47 |
+
writer's TensorAccessor resolves the shard of each tile). ``act_once`` (``KCAT_ACT_ONCE``): the activation on one
|
| 48 |
+
DEST copy of the tile, copied to the second by ``copy_dest_values`` (instead of on both copies)."""
|
| 49 |
+
import ttnn
|
| 50 |
+
|
| 51 |
+
dev = x.device()
|
| 52 |
+
shp = list(x.padded_shape)
|
| 53 |
+
K = shp[-1]
|
| 54 |
+
assert K % 32 == 0 and shp[-2] % 32 == 0 and x.dtype == ttnn.float32 and x.layout == ttnn.TILE_LAYOUT
|
| 55 |
+
Kt = K // 32
|
| 56 |
+
rows = 1
|
| 57 |
+
for d in shp[:-1]:
|
| 58 |
+
rows *= d
|
| 59 |
+
rows //= 32
|
| 60 |
+
total = rows * Kt
|
| 61 |
+
Ktp = kcat_tiles(K, pad)
|
| 62 |
+
oshape = list(x.shape)[:-1] + [32 * Ktp]
|
| 63 |
+
out = ttnn.allocate_tensor_on_device(ttnn.Shape(oshape), ttnn.float32, ttnn.TILE_LAYOUT, dev,
|
| 64 |
+
memory_config or ttnn.DRAM_MEMORY_CONFIG)
|
| 65 |
+
g = dev.compute_with_storage_grid_size()
|
| 66 |
+
n = min(total, g.x * g.y)
|
| 67 |
+
cs = [(i % g.x, i // g.x) for i in range(n)]
|
| 68 |
+
full, rem = divmod(n, g.x)
|
| 69 |
+
rs = []
|
| 70 |
+
if full:
|
| 71 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(g.x - 1, full - 1)))
|
| 72 |
+
if rem:
|
| 73 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, full), ttnn.CoreCoord(rem - 1, full)))
|
| 74 |
+
crs = ttnn.CoreRangeSet(set(rs))
|
| 75 |
+
base, extra = divmod(total, n)
|
| 76 |
+
rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
|
| 77 |
+
t0 = 0
|
| 78 |
+
for i, (cx, cy) in enumerate(cs):
|
| 79 |
+
k = base + (1 if i < extra else 0)
|
| 80 |
+
rd[cx][cy] = [t0, k]
|
| 81 |
+
wr[cx][cy] = [t0, k]
|
| 82 |
+
cp[cx][cy] = [k, 0]
|
| 83 |
+
t0 += k
|
| 84 |
+
|
| 85 |
+
def cb(idx, pages):
|
| 86 |
+
return ttnn.CBDescriptor(total_size=pages * TB, core_ranges=crs, format_descriptors=[
|
| 87 |
+
ttnn.CBFormatDescriptor(buffer_index=idx, data_format=ttnn.float32, page_size=TB)])
|
| 88 |
+
|
| 89 |
+
cbs = [cb(0, 2), cb(16, 2), cb(17, 2), cb(18, 1)] + ([cb(19, 1)] if Ktp > 3 * Kt + 1 else [])
|
| 90 |
+
um = [ttnn.UnpackToDestMode.Default] * 64
|
| 91 |
+
um[0] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 92 |
+
ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True,
|
| 93 |
+
math_approx_mode=False)
|
| 94 |
+
ccfg.unpack_to_dest_mode = um
|
| 95 |
+
SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH
|
| 96 |
+
|
| 97 |
+
def acc(t):
|
| 98 |
+
return list(ttnn.TensorAccessorArgs(t).get_compile_time_args())
|
| 99 |
+
|
| 100 |
+
def kd(name, ct, rt, common, config):
|
| 101 |
+
return ttnn.KernelDescriptor(kernel_source=os.path.join(_KDIR, name), source_type=SRC, core_ranges=crs,
|
| 102 |
+
compile_time_args=ct, defines=[], runtime_args=rt, common_runtime_args=common,
|
| 103 |
+
config=config)
|
| 104 |
+
|
| 105 |
+
ks = [kd("kcat_reader.cpp", [N_RT] + acc(x), rd, [x.buffer_address()], ttnn.ReaderConfigDescriptor()),
|
| 106 |
+
kd("kcat_writer.cpp", [Kt, N_RT, Ktp] + acc(out), wr, [out.buffer_address()], ttnn.WriterConfigDescriptor()),
|
| 107 |
+
kd("kcat_compute.cpp", [N_RT, ACTS[act], int(bool(act_once))], cp, [], ccfg)]
|
| 108 |
+
ttnn.generic_op([x, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))
|
| 109 |
+
return out
|
code/tt_diffusion_planner/tt/kernels/README.md
CHANGED
|
@@ -1,10 +1,32 @@
|
|
| 1 |
# Custom kernels of diffusion-planner-p150
|
| 2 |
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
# Custom kernels of diffusion-planner-p150
|
| 2 |
|
| 3 |
+
The functional port maps every op to stock ttnn (PLAN.md section 2.12). The optimization rounds add `ttnn.generic_op`
|
| 4 |
+
programs whose sources live here (package data in the repo-root `pyproject.toml`, asserted by a `verify:` line of
|
| 5 |
+
`tt-model.yaml`). Each one has a device check against the stock graph it replaces (bit for bit where the math is
|
| 6 |
+
the same) and went through the hang protocol of PLAN.md section 4.4 (smallest shape under `TT_METAL_WATCHER=2`,
|
| 7 |
+
eager twice, a trace): `logs/diffusion-planner/opt_r2/scripts/ln_check.py`, `kcat_check.py`, `smask_check.py`,
|
| 8 |
+
`logs/diffusion-planner/opt_r3/scripts/smsm_check.py`, `fattn_check.py`, `kemit_check.py`, `lntr_check.py`.
|
| 9 |
+
|
| 10 |
+
Per-core runtime-argument counts are compile-time args (probe P3: the generic_op program hash ignores them).
|
| 11 |
+
|
| 12 |
+
| files | op | used by | knob | status |
|
| 13 |
+
|---|---|---|---|---|
|
| 14 |
+
| `ln32_reader.cpp`, `ln32_compute.cpp`, `ln32_writer.cpp` | fp32 LayerNorm over the last dim (the 9-program `layer_norm_fp32` decomposition in one program, the same SFPU LLK calls in the same order), optionally with the residual add `h = x + r (* gate)` in front | `tt/ln_kernel.py`, `tt/layers.py` `LayerNorm` (mixers, decoder) | `LN_KERNEL`, `LN_RESID` | bit-identical to the stock decomposition (OPT round 2) |
|
| 15 |
+
| `kcat_reader.cpp`, `kcat_compute.cpp`, `kcat_writer.cpp` | operand of the K-concatenated split matmul (optionally after the previous linear's GELU, `KCAT_ACT`): fp32 `x` -> `[bf16(x) \| bf16(x) \| x - bf16(x) \| 1 \| 0...]` | `tt/kcat_kernel.py`, `tt/layers.py` `SplitLinear` (decoder) | `SPLIT_KCAT=2` | bit-identical to the stock-op operand (`SPLIT_KCAT=1`) (OPT round 2) |
|
| 16 |
+
| `smask_reader.cpp`, `smask_compute.cpp`, `smask_writer.cpp` | attention score scale + additive mask: `s * scale + mask` in one pass (the two stock binary_ng programs; with a bf16 mask the stock add runs on the FPU, which truncates the scores to TF32, reproduced with an SFPU AND) | `tt/smask_kernel.py`, `tt/attention.py` (decoder self / cross, fusion) | `ATTN_SMASK` | bit-identical (OPT round 2) |
|
| 17 |
+
| `smsm_reader.cpp`, `smsm_compute.cpp`, `smsm_writer.cpp` | attention score scale + mask + softmax over one tile row (head, 32 queries) per pass: the smask SFPU sequence into an L1 row, then the stock `ttnn.softmax(numeric_stable=True)` compute (`kernel_lib` reduce / bcast / exp calls in the stock order: FPU row max, `exp(x - max)`, FPU row sum + precise fp32 reciprocal, bcast multiply); the scale either by a scale tile (`mul_binary_tile`) or an immediate (`mul_unary_tile`, the same fp32 SFPU multiply) | `tt/smsm_kernel.py`, `tt/attention.py` (decoder self / cross, fusion) | `ATTN_SMSM` | bit-identical to smask + `ttnn.softmax` (OPT round 3) |
|
| 18 |
+
| `fattn_reader.cpp`, `fattn_compute.cpp`, `fattn_writer.cpp` | the whole fp32 matmul attention per (head, query tile row): `Q K^T` into DEST (`matmul_block`, in1 transposed), the smsm phases (scale by an immediate, TF32 truncation, mask, stock softmax `kernel_lib` calls) on the row in L1, `P V` accumulated in DEST in key order; Q / K / V read in place from the projection outputs (or the hoisted cross K / V heads), the output written as tile (r, h) of the merged heads | `tt/fattn_kernel.py`, `tt/decoder.py` / `tt/encoder.py` (decoder self / cross, fusion) | `ATTN_FUSED` | bit-identical to head split + `Q K^T` + smsm + `P V` + head merge (OPT round 3) |
|
| 19 |
+
| `ln32s_reader.cpp`, `ln32s_compute.cpp`, `ln32s_writer.cpp` | the fused fp32 LayerNorm (+ residual) with each tile row spread over Wt cores: member j owns column tile j, the root gathers the tiles (semaphore 0), folds them in order and broadcasts mean / rstd (semaphore 1, monotonic counts) | `tt/ln_kernel.py` `layer_norm_fp32_split`, `LayerNorm` (decoder: 11 tile rows on 88 cores) | `LN_SPLIT` | bit-identical (OPT round 2) |
|
| 20 |
+
| `ln32_sfpu.h` | SFPU row sum of one fp32 tile with the sum stored to every column: the stock `sfpu_reduce<SUM, Float32, REDUCE_ROW>` arithmetic (same loads, replayed adds, butterfly) with three extra stores, so the LN statistics come out already broadcast (no RISC-V column fill) | `ln32_compute.cpp`, `ln32s_compute.cpp` | `LN_SFPU_BCAST` | bit-identical (OPT round 2) |
|
| 21 |
+
|
| 22 |
+
`ln32_sfpu.h` is included by the two LN compute kernels through `compiler_include_paths`; the JIT build cache keys a
|
| 23 |
+
kernel on its own source, compile-time args and defines, so after editing the header, also touch the including
|
| 24 |
+
`.cpp` files (or clear the kernel cache) before measuring.
|
| 25 |
+
|
| 26 |
+
`KCAT_EMIT` (OPT round 3): `ln32s_*.cpp` and `fattn_*.cpp` can write their output tile as the split operand
|
| 27 |
+
`[x_hi | x_hi | x_lo | 1 | 0..]` of the next K-concatenated decoder linear (the `kcat_compute.cpp` LLK calls on the
|
| 28 |
+
packed fp32 tile, the `kcat_writer.cpp` layout) instead of the plain output: bit-identical to output + `kcat_operand`.
|
| 29 |
+
|
| 30 |
+
`LN_TR` (OPT round 3): `ln32_*.cpp` can read the residual per-entity transposed (`[.., E, W, T]`, transposed back with
|
| 31 |
+
`transpose_tile`, the stock `ttnn.transpose` LLK with the fp32 unpack-to-dest: exact) and write the output the same
|
| 32 |
+
way (via c_18): the two transposes around the mixer's token-mixing MLP, bit-identical.
|
code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Fused fp32 matmul attention (tt/fattn_kernel.py, ATTN_FUSED), compute. Per unit (head h, query tile row r; Wt key
|
| 3 |
+
// tiles), the five stock programs of tt/attention.py (ATTN_FAST=1 + ATTN_SMSM) in one pass, the scores and the
|
| 4 |
+
// probabilities kept in L1:
|
| 5 |
+
// phase Q: s_j = Q_r K_j^T in DEST (matmul_block, in1 transposed by the unpacker: the stock Q K^T matmul with its
|
| 6 |
+
// K = 32 in one block), two key tiles at a time;
|
| 7 |
+
// phase A: x_j = trunc_tf32(s_j * scale) + mask_j on the SFPU (kernels/smsm_compute.cpp, scale mode 1), packed to
|
| 8 |
+
// cb_x;
|
| 9 |
+
// phase B: the stock numeric-stable softmax of the row (kernels/smsm_compute.cpp phase B) -> cb_p;
|
| 10 |
+
// phase C: o = sum_j P_j V_j accumulated in DEST in key order (the stock P V matmul: the whole K in one block),
|
| 11 |
+
// packed to cb_o; the writer stores it as tile (r, h) of the merged-heads output.
|
| 12 |
+
// (kcat_out, KCAT_EMIT) the output tile as the operand of the next K-concatenated split linear: o_hi = bf16(o)
|
| 13 |
+
// and o_lo = o - o_hi (kernels/kcat_compute.cpp's LLK calls on the packed fp32 tile) -> cb_o / cb_lo.
|
| 14 |
+
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] trunc_tf32, [3] ndst, [4] P pad tiles per row (to a multiple
|
| 15 |
+
// of ndst), [5] scale bits (fp32), [6] kcat_out. Per-core RT args: [nunits, 0].
|
| 16 |
+
#include <cstdint>
|
| 17 |
+
|
| 18 |
+
#include "api/compute/common.h"
|
| 19 |
+
#include "api/compute/compute_kernel_api.h"
|
| 20 |
+
#include "api/compute/eltwise_binary.h"
|
| 21 |
+
#include "api/compute/eltwise_binary_sfpu.h"
|
| 22 |
+
#include "api/compute/eltwise_unary/binop_with_scalar.h"
|
| 23 |
+
#include "api/compute/eltwise_unary/bitwise.h"
|
| 24 |
+
#include "api/compute/eltwise_unary/eltwise_unary.h"
|
| 25 |
+
#include "api/compute/eltwise_unary/typecast.h"
|
| 26 |
+
#include "api/compute/tile_move_copy.h"
|
| 27 |
+
#include "api/compute/matmul.h"
|
| 28 |
+
#include "api/compute/bcast.h"
|
| 29 |
+
#include "api/compute/softmax.h"
|
| 30 |
+
#include "api/compute/reduce.h"
|
| 31 |
+
#include "ttnn/cpp/ttnn/kernel_lib/reduce_helpers_compute.hpp"
|
| 32 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/chain.hpp"
|
| 33 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/convenience.hpp"
|
| 34 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/unary/math.hpp"
|
| 35 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/core/optional.hpp"
|
| 36 |
+
|
| 37 |
+
namespace ckl = compute_kernel_lib;
|
| 38 |
+
using namespace ckernel;
|
| 39 |
+
|
| 40 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 41 |
+
constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2);
|
| 42 |
+
constexpr uint32_t ndst = get_compile_time_arg_val(3);
|
| 43 |
+
constexpr uint32_t p_pad = get_compile_time_arg_val(4);
|
| 44 |
+
constexpr uint32_t scale_bits = get_compile_time_arg_val(5);
|
| 45 |
+
constexpr uint32_t kcat_out = get_compile_time_arg_val(6);
|
| 46 |
+
constexpr uint32_t cb_q = 0, cb_m = 1, cb_k = 2, cb_max_scaler = 3, cb_sum_scaler = 4, cb_v = 5;
|
| 47 |
+
constexpr uint32_t cb_o = 16, cb_lo = 17, cb_y = 29;
|
| 48 |
+
constexpr uint32_t cb_x = 24, cb_max = 25, cb_exps = 26, cb_recip = 27, cb_p = 28;
|
| 49 |
+
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
|
| 50 |
+
constexpr uint32_t Wp = (Wt + 1) / 2;
|
| 51 |
+
constexpr uint32_t Wr = Wt + p_pad;
|
| 52 |
+
|
| 53 |
+
template <std::uint32_t dfb_in, std::uint32_t dfb_max_scaler, std::uint32_t dfb_max, std::uint32_t dfb_out>
|
| 54 |
+
void calc_numeric_stable(std::uint32_t W, std::uint32_t nd) {
|
| 55 |
+
compute_kernel_lib::reduce<
|
| 56 |
+
PoolType::MAX,
|
| 57 |
+
ReduceDim::REDUCE_ROW,
|
| 58 |
+
dfb_in,
|
| 59 |
+
dfb_max_scaler,
|
| 60 |
+
dfb_max,
|
| 61 |
+
compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop,
|
| 62 |
+
compute_kernel_lib::ReduceDataFormatReconfigMode::INPUT>(compute_kernel_lib::ReduceInputBlockShape::row(W));
|
| 63 |
+
ckl::eltwise_chain(
|
| 64 |
+
ckl::IterationShape::tiles(W).block_size(nd),
|
| 65 |
+
ckl::BinaryFpu<
|
| 66 |
+
ckl::BinaryFpuOp::Sub,
|
| 67 |
+
ckl::input(
|
| 68 |
+
dfb_in,
|
| 69 |
+
ckl::WaitPolicy::Upfront,
|
| 70 |
+
ckl::PopPolicy::AtEnd,
|
| 71 |
+
ckl::InputTileMapping::Block,
|
| 72 |
+
ckl::DataFormatReconfig::Disabled),
|
| 73 |
+
ckl::input(dfb_max, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd)>{},
|
| 74 |
+
ckl::Exp<ckl::Approx::Exact, ckl::Dst::D0>{},
|
| 75 |
+
ckl::PackTile<ckl::output(
|
| 76 |
+
dfb_out,
|
| 77 |
+
ckl::ReservePolicy::PerBlockSize,
|
| 78 |
+
ckl::PushPolicy::PerBlockSize,
|
| 79 |
+
ckl::DataFormatReconfig::Disabled)>{});
|
| 80 |
+
cb_wait_front(dfb_out, W);
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
void kernel_main() {
|
| 84 |
+
const uint32_t nunits = get_arg_val<uint32_t>(0);
|
| 85 |
+
if (nunits == 0) {
|
| 86 |
+
return;
|
| 87 |
+
}
|
| 88 |
+
compute_kernel_hw_startup(cb_q, cb_k, cb_x);
|
| 89 |
+
cb_wait_front(cb_max_scaler, 1);
|
| 90 |
+
cb_wait_front(cb_sum_scaler, 1);
|
| 91 |
+
for (uint32_t u = 0; u < nunits; ++u) {
|
| 92 |
+
// ---- phases Q + A: x_j = trunc(Q K_j^T * scale) + mask_j, two key tiles at a time ----
|
| 93 |
+
cb_wait_front(cb_q, 1);
|
| 94 |
+
cb_wait_front(cb_m, 2 * Wp);
|
| 95 |
+
pack_reconfig_data_format(cb_x);
|
| 96 |
+
for (uint32_t p = 0; p < Wp; ++p) {
|
| 97 |
+
const uint32_t j = 2 * p;
|
| 98 |
+
const bool two = (j + 1) < Wt;
|
| 99 |
+
cb_wait_front(cb_k, 2);
|
| 100 |
+
reconfig_data_format(cb_k, cb_q);
|
| 101 |
+
matmul_block_init(cb_q, cb_k, 1, 1, 1, 1);
|
| 102 |
+
tile_regs_acquire();
|
| 103 |
+
matmul_block(cb_q, cb_k, 0, 0, 0, 1, 1, 1, 1);
|
| 104 |
+
matmul_block(cb_q, cb_k, 0, 1, 1, 1, 1, 1, 1);
|
| 105 |
+
binop_with_scalar_tile_init();
|
| 106 |
+
mul_unary_tile(0, scale_bits);
|
| 107 |
+
mul_unary_tile(1, scale_bits);
|
| 108 |
+
if constexpr (trunc_tf32) {
|
| 109 |
+
bitwise_and_tile_init();
|
| 110 |
+
bitwise_and_tile<DataFormat::Int32>(0, 0xFFFFE000u);
|
| 111 |
+
bitwise_and_tile<DataFormat::Int32>(1, 0xFFFFE000u);
|
| 112 |
+
}
|
| 113 |
+
reconfig_data_format_srca(cb_k, cb_m);
|
| 114 |
+
copy_init(cb_m);
|
| 115 |
+
copy_tile(cb_m, j, 2);
|
| 116 |
+
copy_tile(cb_m, j + 1, 3);
|
| 117 |
+
add_binary_tile_init();
|
| 118 |
+
add_binary_tile<RNE>(0, 2, 0);
|
| 119 |
+
add_binary_tile<RNE>(1, 3, 1);
|
| 120 |
+
const uint32_t np = two ? 2 : 1;
|
| 121 |
+
cb_reserve_back(cb_x, np);
|
| 122 |
+
tile_regs_commit();
|
| 123 |
+
tile_regs_wait();
|
| 124 |
+
pack_tile(0, cb_x);
|
| 125 |
+
if (two) {
|
| 126 |
+
pack_tile(1, cb_x);
|
| 127 |
+
}
|
| 128 |
+
tile_regs_release();
|
| 129 |
+
cb_push_back(cb_x, np);
|
| 130 |
+
cb_pop_front(cb_k, 2);
|
| 131 |
+
}
|
| 132 |
+
cb_pop_front(cb_m, 2 * Wp);
|
| 133 |
+
cb_pop_front(cb_q, 1);
|
| 134 |
+
|
| 135 |
+
// ---- phase B: the stock numeric-stable softmax of the row -> cb_p ----
|
| 136 |
+
reconfig_data_format(cb_x, cb_x);
|
| 137 |
+
pack_reconfig_data_format(cb_exps);
|
| 138 |
+
copy_init(cb_x);
|
| 139 |
+
calc_numeric_stable<cb_x, cb_max_scaler, cb_max, cb_exps>(Wt, ndst);
|
| 140 |
+
reconfig_data_format(cb_exps, cb_sum_scaler);
|
| 141 |
+
compute_kernel_lib::reduce<
|
| 142 |
+
PoolType::SUM,
|
| 143 |
+
ReduceDim::REDUCE_ROW,
|
| 144 |
+
cb_exps,
|
| 145 |
+
cb_sum_scaler,
|
| 146 |
+
cb_recip,
|
| 147 |
+
compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop>(
|
| 148 |
+
compute_kernel_lib::ReduceInputBlockShape::row(Wt),
|
| 149 |
+
compute_kernel_lib::ReduceInputMemoryLayout::contiguous(),
|
| 150 |
+
compute_kernel_lib::NoAccumulation{},
|
| 151 |
+
[](std::uint32_t) {
|
| 152 |
+
if constexpr (DST_ACCUM_MODE) {
|
| 153 |
+
recip_tile_init<ReciprocalDestAcc::FP32, ReciprocalApproxMode::Precise>();
|
| 154 |
+
recip_tile<ReciprocalDestAcc::FP32, ReciprocalApproxMode::Precise>(0);
|
| 155 |
+
} else {
|
| 156 |
+
recip_tile_init();
|
| 157 |
+
recip_tile(0);
|
| 158 |
+
}
|
| 159 |
+
});
|
| 160 |
+
ckl::mul<
|
| 161 |
+
ckl::input(cb_exps, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd, ckl::InputTileMapping::Block),
|
| 162 |
+
ckl::input(cb_recip, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd),
|
| 163 |
+
ckl::output(cb_p, ckl::ReservePolicy::PerBlockSize, ckl::PushPolicy::PerBlockSize)>(
|
| 164 |
+
ckl::IterationShape::tiles(Wt).block_size(ndst));
|
| 165 |
+
if constexpr (p_pad > 0) {
|
| 166 |
+
cb_reserve_back(cb_p, p_pad);
|
| 167 |
+
cb_push_back(cb_p, p_pad);
|
| 168 |
+
}
|
| 169 |
+
|
| 170 |
+
// ---- phase C: o = sum_j P_j V_j (key order, one DEST accumulation) ----
|
| 171 |
+
cb_wait_front(cb_p, Wr);
|
| 172 |
+
reconfig_data_format(cb_v, cb_p);
|
| 173 |
+
pack_reconfig_data_format(cb_o);
|
| 174 |
+
matmul_block_init(cb_p, cb_v, 0, 1, 1, 1);
|
| 175 |
+
tile_regs_acquire();
|
| 176 |
+
for (uint32_t j = 0; j < Wt; ++j) {
|
| 177 |
+
cb_wait_front(cb_v, 1);
|
| 178 |
+
matmul_block(cb_p, cb_v, j, 0, 0, 0, 1, 1, 1);
|
| 179 |
+
cb_pop_front(cb_v, 1);
|
| 180 |
+
}
|
| 181 |
+
if constexpr (kcat_out) {
|
| 182 |
+
cb_reserve_back(cb_y, 1);
|
| 183 |
+
tile_regs_commit();
|
| 184 |
+
tile_regs_wait();
|
| 185 |
+
pack_tile(0, cb_y);
|
| 186 |
+
tile_regs_release();
|
| 187 |
+
cb_push_back(cb_y, 1);
|
| 188 |
+
cb_pop_front(cb_p, Wr);
|
| 189 |
+
// ---- the split operand of the next linear (kernels/kcat_compute.cpp) ----
|
| 190 |
+
cb_wait_front(cb_y, 1);
|
| 191 |
+
reconfig_data_format_srca(cb_y);
|
| 192 |
+
tile_regs_acquire();
|
| 193 |
+
copy_init(cb_y);
|
| 194 |
+
copy_tile(cb_y, 0, 0);
|
| 195 |
+
copy_tile(cb_y, 0, 1);
|
| 196 |
+
typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>();
|
| 197 |
+
typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0);
|
| 198 |
+
sub_binary_tile_init();
|
| 199 |
+
sub_binary_tile<ckernel::DstRoundingMode::NearestEven>(1, 0, 1);
|
| 200 |
+
cb_reserve_back(cb_o, 1);
|
| 201 |
+
cb_reserve_back(cb_lo, 1);
|
| 202 |
+
tile_regs_commit();
|
| 203 |
+
tile_regs_wait();
|
| 204 |
+
pack_tile(0, cb_o);
|
| 205 |
+
pack_tile(1, cb_lo);
|
| 206 |
+
tile_regs_release();
|
| 207 |
+
cb_push_back(cb_o, 1);
|
| 208 |
+
cb_push_back(cb_lo, 1);
|
| 209 |
+
cb_pop_front(cb_y, 1);
|
| 210 |
+
} else {
|
| 211 |
+
cb_reserve_back(cb_o, 1);
|
| 212 |
+
tile_regs_commit();
|
| 213 |
+
tile_regs_wait();
|
| 214 |
+
pack_tile(0, cb_o);
|
| 215 |
+
tile_regs_release();
|
| 216 |
+
cb_push_back(cb_o, 1);
|
| 217 |
+
cb_pop_front(cb_p, Wr);
|
| 218 |
+
}
|
| 219 |
+
}
|
| 220 |
+
cb_pop_front(cb_max_scaler, 1);
|
| 221 |
+
cb_pop_front(cb_sum_scaler, 1);
|
| 222 |
+
}
|
code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Fused fp32 matmul attention (tt/fattn_kernel.py), reader (RISCV_0). Once: the two reduce scaler tiles of the stock
|
| 3 |
+
// softmax (fp32 1.0 in row 0 of every face). Per unit u of this core's range (head h = u / Mt, query tile row
|
| 4 |
+
// r = u % Mt): the Wt mask tiles of row r (+ 1 unused tile when Wt is odd), the Q tile (r, h), the K tiles (j, h) in
|
| 5 |
+
// pairs (the last pair of an odd row carries one unused tile), then the V tiles (j, h) one by one.
|
| 6 |
+
// Tile (row, head) of a source = base + row * row_stride + head * head_stride (tiles): a [1, 1, S, W] projection
|
| 7 |
+
// output read in place (row_stride = W / 32, head_stride = 1, base = the column tile of the first head) or a
|
| 8 |
+
// [1, H, S, 32] heads tensor (row_stride = 1, head_stride = S / 32).
|
| 9 |
+
// CT args: [0] Wt, [1] Mt, [2] per-core RT-arg count (P3), [3] mask tile bytes, [4..12] (base, row_stride,
|
| 10 |
+
// head_stride) of Q, K, V, then the TensorAccessorArgs of q, k, v, mask.
|
| 11 |
+
// Common RT args: [q_addr, k_addr, v_addr, m_addr]. Per-core RT args: [u0, n].
|
| 12 |
+
#include <cstdint>
|
| 13 |
+
|
| 14 |
+
#include "api/dataflow/dataflow_api.h"
|
| 15 |
+
|
| 16 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 17 |
+
constexpr uint32_t Mt = get_compile_time_arg_val(1);
|
| 18 |
+
constexpr uint32_t TBM = get_compile_time_arg_val(3);
|
| 19 |
+
constexpr uint32_t qb = get_compile_time_arg_val(4), qr = get_compile_time_arg_val(5), qh = get_compile_time_arg_val(6);
|
| 20 |
+
constexpr uint32_t kb = get_compile_time_arg_val(7), kr = get_compile_time_arg_val(8), kh = get_compile_time_arg_val(9);
|
| 21 |
+
constexpr uint32_t vb = get_compile_time_arg_val(10), vr = get_compile_time_arg_val(11),
|
| 22 |
+
vh = get_compile_time_arg_val(12);
|
| 23 |
+
constexpr auto q_args = TensorAccessorArgs<13>();
|
| 24 |
+
constexpr auto k_args = TensorAccessorArgs<q_args.next_compile_time_args_offset()>();
|
| 25 |
+
constexpr auto v_args = TensorAccessorArgs<k_args.next_compile_time_args_offset()>();
|
| 26 |
+
constexpr auto m_args = TensorAccessorArgs<v_args.next_compile_time_args_offset()>();
|
| 27 |
+
constexpr uint32_t cb_q = 0, cb_m = 1, cb_k = 2, cb_max_scaler = 3, cb_sum_scaler = 4, cb_v = 5;
|
| 28 |
+
constexpr uint32_t TB = 4096;
|
| 29 |
+
constexpr uint32_t Wp = (Wt + 1) / 2;
|
| 30 |
+
|
| 31 |
+
inline void fill_scaler(uint32_t cb) {
|
| 32 |
+
cb_reserve_back(cb, 1);
|
| 33 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb));
|
| 34 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 35 |
+
p[i] = 0;
|
| 36 |
+
}
|
| 37 |
+
for (uint32_t f = 0; f < 4; ++f) {
|
| 38 |
+
for (uint32_t c = 0; c < 16; ++c) {
|
| 39 |
+
p[f * 256 + c] = 0x3F800000u;
|
| 40 |
+
}
|
| 41 |
+
}
|
| 42 |
+
cb_push_back(cb, 1);
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
void kernel_main() {
|
| 46 |
+
const uint32_t q_addr = get_common_arg_val<uint32_t>(0);
|
| 47 |
+
const uint32_t k_addr = get_common_arg_val<uint32_t>(1);
|
| 48 |
+
const uint32_t v_addr = get_common_arg_val<uint32_t>(2);
|
| 49 |
+
const uint32_t m_addr = get_common_arg_val<uint32_t>(3);
|
| 50 |
+
const uint32_t u0 = get_arg_val<uint32_t>(0);
|
| 51 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 52 |
+
if (n == 0) {
|
| 53 |
+
return;
|
| 54 |
+
}
|
| 55 |
+
const auto qa = TensorAccessor(q_args, q_addr, TB);
|
| 56 |
+
const auto ka = TensorAccessor(k_args, k_addr, TB);
|
| 57 |
+
const auto va = TensorAccessor(v_args, v_addr, TB);
|
| 58 |
+
const auto ma = TensorAccessor(m_args, m_addr, TBM);
|
| 59 |
+
fill_scaler(cb_max_scaler);
|
| 60 |
+
fill_scaler(cb_sum_scaler);
|
| 61 |
+
for (uint32_t u = u0; u < u0 + n; ++u) {
|
| 62 |
+
const uint32_t h = u / Mt;
|
| 63 |
+
const uint32_t r = u % Mt;
|
| 64 |
+
cb_reserve_back(cb_m, 2 * Wp);
|
| 65 |
+
{
|
| 66 |
+
uint32_t p = get_write_ptr(cb_m);
|
| 67 |
+
for (uint32_t j = 0; j < Wt; ++j) {
|
| 68 |
+
noc_async_read(ma.get_noc_addr(r * Wt + j), p, TBM);
|
| 69 |
+
p += TBM;
|
| 70 |
+
}
|
| 71 |
+
}
|
| 72 |
+
cb_reserve_back(cb_q, 1);
|
| 73 |
+
noc_async_read(qa.get_noc_addr(qb + r * qr + h * qh), get_write_ptr(cb_q), TB);
|
| 74 |
+
noc_async_read_barrier();
|
| 75 |
+
cb_push_back(cb_m, 2 * Wp);
|
| 76 |
+
cb_push_back(cb_q, 1);
|
| 77 |
+
for (uint32_t pp = 0; pp < Wp; ++pp) {
|
| 78 |
+
const uint32_t j = 2 * pp;
|
| 79 |
+
cb_reserve_back(cb_k, 2);
|
| 80 |
+
const uint32_t p = get_write_ptr(cb_k);
|
| 81 |
+
noc_async_read(ka.get_noc_addr(kb + j * kr + h * kh), p, TB);
|
| 82 |
+
if (j + 1 < Wt) {
|
| 83 |
+
noc_async_read(ka.get_noc_addr(kb + (j + 1) * kr + h * kh), p + TB, TB);
|
| 84 |
+
}
|
| 85 |
+
noc_async_read_barrier();
|
| 86 |
+
cb_push_back(cb_k, 2);
|
| 87 |
+
}
|
| 88 |
+
for (uint32_t j = 0; j < Wt; ++j) {
|
| 89 |
+
cb_reserve_back(cb_v, 1);
|
| 90 |
+
noc_async_read(va.get_noc_addr(vb + j * vr + h * vh), get_write_ptr(cb_v), TB);
|
| 91 |
+
noc_async_read_barrier();
|
| 92 |
+
cb_push_back(cb_v, 1);
|
| 93 |
+
}
|
| 94 |
+
}
|
| 95 |
+
}
|
code/tt_diffusion_planner/tt/kernels/fattn_writer.cpp
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Fused fp32 matmul attention (tt/fattn_kernel.py), writer (RISCV_1): per unit u (head h = u / Mt, query tile row
|
| 3 |
+
// r = u % Mt) the output tile as tile (r, h) of the merged-heads [1, 1, Sq, H * 32] tensor (what nlp_concat_heads
|
| 4 |
+
// makes of the [1, H, Sq, 32] P V output).
|
| 5 |
+
// With kcat_out (KCAT_EMIT) the output is instead the split operand [o_hi | o_hi | o_lo | 1 | 0..] of the next
|
| 6 |
+
// K-concatenated linear (row tiles Ktp, kernels/kcat_writer.cpp's layout): o_hi to (r, h) and (r, H + h), o_lo to
|
| 7 |
+
// (r, 2 H + h); the unit of head 0 also writes the ones tile (r, 3 H) and the zero pad tiles.
|
| 8 |
+
// CT args: [0] Mt, [1] H, [2] per-core RT-arg count (P3), [3] kcat_out, [4] Ktp, then the TensorAccessorArgs of out.
|
| 9 |
+
// Common RT args: [out_addr]. Per-core RT args: [u0, n].
|
| 10 |
+
#include <cstdint>
|
| 11 |
+
|
| 12 |
+
#include "api/dataflow/dataflow_api.h"
|
| 13 |
+
|
| 14 |
+
constexpr uint32_t Mt = get_compile_time_arg_val(0);
|
| 15 |
+
constexpr uint32_t H = get_compile_time_arg_val(1);
|
| 16 |
+
constexpr uint32_t kcat_out = get_compile_time_arg_val(3);
|
| 17 |
+
constexpr uint32_t Ktp = get_compile_time_arg_val(4);
|
| 18 |
+
constexpr auto o_args = TensorAccessorArgs<5>();
|
| 19 |
+
constexpr uint32_t cb_o = 16, cb_lo = 17, cb_one = 18, cb_zero = 19;
|
| 20 |
+
constexpr uint32_t TB = 4096;
|
| 21 |
+
|
| 22 |
+
void kernel_main() {
|
| 23 |
+
const uint32_t o_addr = get_common_arg_val<uint32_t>(0);
|
| 24 |
+
const uint32_t u0 = get_arg_val<uint32_t>(0);
|
| 25 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 26 |
+
const auto o = TensorAccessor(o_args, o_addr, TB);
|
| 27 |
+
if constexpr (!kcat_out) {
|
| 28 |
+
for (uint32_t u = u0; u < u0 + n; ++u) {
|
| 29 |
+
const uint32_t h = u / Mt;
|
| 30 |
+
const uint32_t r = u % Mt;
|
| 31 |
+
cb_wait_front(cb_o, 1);
|
| 32 |
+
noc_async_write(get_read_ptr(cb_o), o.get_noc_addr(r * H + h), TB);
|
| 33 |
+
noc_async_writes_flushed();
|
| 34 |
+
cb_pop_front(cb_o, 1);
|
| 35 |
+
}
|
| 36 |
+
} else {
|
| 37 |
+
// ones tile: element (i, c) at face (i / 16) * 2 + (c / 16), offset (i % 16) * 16 + c % 16
|
| 38 |
+
uint32_t one_ptr = 0, zero_ptr = 0;
|
| 39 |
+
if (u0 < Mt) {
|
| 40 |
+
cb_reserve_back(cb_one, 1);
|
| 41 |
+
one_ptr = get_write_ptr(cb_one);
|
| 42 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(one_ptr);
|
| 43 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 44 |
+
p[i] = 0;
|
| 45 |
+
}
|
| 46 |
+
for (uint32_t i = 0; i < 32; ++i) {
|
| 47 |
+
const uint32_t base = (i / 16) * 512 + (i % 16) * 16;
|
| 48 |
+
p[base] = 0x3F800000u;
|
| 49 |
+
p[base + 1] = 0x3F800000u;
|
| 50 |
+
}
|
| 51 |
+
if constexpr (Ktp > 3 * H + 1) {
|
| 52 |
+
cb_reserve_back(cb_zero, 1);
|
| 53 |
+
zero_ptr = get_write_ptr(cb_zero);
|
| 54 |
+
auto* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(zero_ptr);
|
| 55 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 56 |
+
z[i] = 0;
|
| 57 |
+
}
|
| 58 |
+
}
|
| 59 |
+
}
|
| 60 |
+
for (uint32_t u = u0; u < u0 + n; ++u) {
|
| 61 |
+
const uint32_t h = u / Mt;
|
| 62 |
+
const uint32_t r = u % Mt;
|
| 63 |
+
const uint32_t ob = r * Ktp;
|
| 64 |
+
cb_wait_front(cb_o, 1);
|
| 65 |
+
const uint32_t hp = get_read_ptr(cb_o);
|
| 66 |
+
noc_async_write(hp, o.get_noc_addr(ob + h), TB);
|
| 67 |
+
noc_async_write(hp, o.get_noc_addr(ob + H + h), TB);
|
| 68 |
+
cb_wait_front(cb_lo, 1);
|
| 69 |
+
noc_async_write(get_read_ptr(cb_lo), o.get_noc_addr(ob + 2 * H + h), TB);
|
| 70 |
+
if (h == 0) {
|
| 71 |
+
noc_async_write(one_ptr, o.get_noc_addr(ob + 3 * H), TB);
|
| 72 |
+
for (uint32_t z = 3 * H + 1; z < Ktp; ++z) {
|
| 73 |
+
noc_async_write(zero_ptr, o.get_noc_addr(ob + z), TB);
|
| 74 |
+
}
|
| 75 |
+
}
|
| 76 |
+
noc_async_writes_flushed();
|
| 77 |
+
cb_pop_front(cb_o, 1);
|
| 78 |
+
cb_pop_front(cb_lo, 1);
|
| 79 |
+
}
|
| 80 |
+
}
|
| 81 |
+
noc_async_write_barrier();
|
| 82 |
+
}
|
code/tt_diffusion_planner/tt/kernels/kcat_compute.cpp
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Split-matmul operand build (tt/kcat_kernel.py), compute: per fp32 x tile, x_hi = bf16(x) (typecast_tile
|
| 3 |
+
// Float32 -> Float16_b, the stock ttnn.typecast LLK) kept as fp32, and x_lo = x - x_hi (sub_binary_tile, exact).
|
| 4 |
+
// With act (KCAT_ACT) the previous linear's activation is applied first, as the stock unary program does it
|
| 5 |
+
// (math_approx_mode false): 1 = gelu_tile<0> (ttnn.gelu, accurate), 2 = gelu_tanh_tile (GeluVariant.Tanh).
|
| 6 |
+
// CT args: [0] per-core RT-arg count (P3), [1] act, [2] once (KCAT_ACT_ONCE: the activation on one DEST copy, then
|
| 7 |
+
// copy_dest_values to the second; the same values). Per-core RT args: [n, 0].
|
| 8 |
+
#include <cstdint>
|
| 9 |
+
|
| 10 |
+
#include "api/compute/common.h"
|
| 11 |
+
#include "api/compute/compute_kernel_api.h"
|
| 12 |
+
#include "api/compute/copy_dest_values.h"
|
| 13 |
+
#include "api/compute/eltwise_binary_sfpu.h"
|
| 14 |
+
#include "api/compute/eltwise_unary/eltwise_unary.h"
|
| 15 |
+
#include "api/compute/eltwise_unary/gelu.h"
|
| 16 |
+
#include "api/compute/eltwise_unary/typecast.h"
|
| 17 |
+
#include "api/compute/tile_move_copy.h"
|
| 18 |
+
|
| 19 |
+
using namespace ckernel;
|
| 20 |
+
|
| 21 |
+
constexpr uint32_t act = get_compile_time_arg_val(1);
|
| 22 |
+
constexpr uint32_t once = get_compile_time_arg_val(2);
|
| 23 |
+
constexpr uint32_t cb_x = 0, cb_hi = 16, cb_lo = 17;
|
| 24 |
+
|
| 25 |
+
void kernel_main() {
|
| 26 |
+
const uint32_t n = get_arg_val<uint32_t>(0);
|
| 27 |
+
if (n == 0) {
|
| 28 |
+
return;
|
| 29 |
+
}
|
| 30 |
+
compute_kernel_hw_startup(cb_x, cb_hi);
|
| 31 |
+
for (uint32_t t = 0; t < n; ++t) {
|
| 32 |
+
cb_wait_front(cb_x, 1);
|
| 33 |
+
tile_regs_acquire();
|
| 34 |
+
copy_init(cb_x);
|
| 35 |
+
copy_tile(cb_x, 0, 0);
|
| 36 |
+
if constexpr (once && act != 0) {
|
| 37 |
+
if constexpr (act == 1) {
|
| 38 |
+
gelu_tile_init<0u>();
|
| 39 |
+
gelu_tile<0u>(0);
|
| 40 |
+
} else {
|
| 41 |
+
gelu_tanh_tile_init();
|
| 42 |
+
gelu_tanh_tile(0);
|
| 43 |
+
}
|
| 44 |
+
copy_dest_values_init();
|
| 45 |
+
copy_dest_values<DataFormat::Float32>(0, 1);
|
| 46 |
+
} else {
|
| 47 |
+
copy_tile(cb_x, 0, 1);
|
| 48 |
+
}
|
| 49 |
+
if constexpr (once && act != 0) {
|
| 50 |
+
} else if constexpr (act == 1) {
|
| 51 |
+
gelu_tile_init<0u>();
|
| 52 |
+
gelu_tile<0u>(0);
|
| 53 |
+
gelu_tile<0u>(1);
|
| 54 |
+
} else if constexpr (act == 2) {
|
| 55 |
+
gelu_tanh_tile_init();
|
| 56 |
+
gelu_tanh_tile(0);
|
| 57 |
+
gelu_tanh_tile(1);
|
| 58 |
+
}
|
| 59 |
+
typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>();
|
| 60 |
+
typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0);
|
| 61 |
+
sub_binary_tile_init();
|
| 62 |
+
sub_binary_tile<ckernel::DstRoundingMode::NearestEven>(1, 0, 1);
|
| 63 |
+
cb_reserve_back(cb_hi, 1);
|
| 64 |
+
cb_reserve_back(cb_lo, 1);
|
| 65 |
+
tile_regs_commit();
|
| 66 |
+
tile_regs_wait();
|
| 67 |
+
pack_tile(0, cb_hi);
|
| 68 |
+
pack_tile(1, cb_lo);
|
| 69 |
+
tile_regs_release();
|
| 70 |
+
cb_push_back(cb_hi, 1);
|
| 71 |
+
cb_push_back(cb_lo, 1);
|
| 72 |
+
cb_pop_front(cb_x, 1);
|
| 73 |
+
}
|
| 74 |
+
}
|
code/tt_diffusion_planner/tt/kernels/kcat_reader.cpp
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Split-matmul operand build (tt/kcat_kernel.py), reader (RISCV_0): this core's contiguous range of x tiles.
|
| 3 |
+
// CT args: [0] per-core RT-arg count (P3), then the TensorAccessorArgs of x.
|
| 4 |
+
// Common RT args: [x_addr]. Per-core RT args: [t0, n] (flattened x tile indices).
|
| 5 |
+
#include <cstdint>
|
| 6 |
+
|
| 7 |
+
#include "api/dataflow/dataflow_api.h"
|
| 8 |
+
|
| 9 |
+
constexpr auto x_args = TensorAccessorArgs<1>();
|
| 10 |
+
constexpr uint32_t cb_x = 0;
|
| 11 |
+
constexpr uint32_t TB = 4096;
|
| 12 |
+
|
| 13 |
+
void kernel_main() {
|
| 14 |
+
const uint32_t x_addr = get_common_arg_val<uint32_t>(0);
|
| 15 |
+
const uint32_t t0 = get_arg_val<uint32_t>(0);
|
| 16 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 17 |
+
const auto x = TensorAccessor(x_args, x_addr, TB);
|
| 18 |
+
for (uint32_t t = t0; t < t0 + n; ++t) {
|
| 19 |
+
cb_reserve_back(cb_x, 1);
|
| 20 |
+
noc_async_read(x.get_noc_addr(t), get_write_ptr(cb_x), TB);
|
| 21 |
+
noc_async_read_barrier();
|
| 22 |
+
cb_push_back(cb_x, 1);
|
| 23 |
+
}
|
| 24 |
+
}
|
code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Split-matmul operand build (tt/kcat_kernel.py), writer (RISCV_1): x tile (r, j) of [.., M, K] goes to
|
| 3 |
+
// X' (r, j) and (r, Kt + j) as x_hi and to (r, 2 Kt + j) as x_lo; the core holding (r, 0) also writes the ones tile
|
| 4 |
+
// (r, 3 Kt) (columns 0 and 1 = 1.0: the bias rows of the concatenated weight) and the zero tiles
|
| 5 |
+
// (r, 3 Kt + 1 .. Ktp - 1) that pad K' to a multiple of the matmul's K block.
|
| 6 |
+
// CT args: [0] Kt, [1] per-core RT-arg count (P3), [2] Ktp (row tiles of X'), then the TensorAccessorArgs of X'.
|
| 7 |
+
// Common RT args: [out_addr]. Per-core RT args: [t0, n].
|
| 8 |
+
#include <cstdint>
|
| 9 |
+
|
| 10 |
+
#include "api/dataflow/dataflow_api.h"
|
| 11 |
+
|
| 12 |
+
constexpr uint32_t Kt = get_compile_time_arg_val(0);
|
| 13 |
+
constexpr uint32_t OKt = get_compile_time_arg_val(2);
|
| 14 |
+
constexpr auto o_args = TensorAccessorArgs<3>();
|
| 15 |
+
constexpr uint32_t cb_hi = 16, cb_lo = 17, cb_one = 18, cb_zero = 19;
|
| 16 |
+
constexpr uint32_t TB = 4096;
|
| 17 |
+
|
| 18 |
+
void kernel_main() {
|
| 19 |
+
const uint32_t o_addr = get_common_arg_val<uint32_t>(0);
|
| 20 |
+
const uint32_t t0 = get_arg_val<uint32_t>(0);
|
| 21 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 22 |
+
const auto o = TensorAccessor(o_args, o_addr, TB);
|
| 23 |
+
// ones tile: element (i, c) at face (i / 16) * 2 + (c / 16), offset (i % 16) * 16 + c % 16
|
| 24 |
+
cb_reserve_back(cb_one, 1);
|
| 25 |
+
const uint32_t one_ptr = get_write_ptr(cb_one);
|
| 26 |
+
{
|
| 27 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(one_ptr);
|
| 28 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 29 |
+
p[i] = 0;
|
| 30 |
+
}
|
| 31 |
+
for (uint32_t i = 0; i < 32; ++i) {
|
| 32 |
+
const uint32_t base = (i / 16) * 512 + (i % 16) * 16;
|
| 33 |
+
p[base] = 0x3F800000u;
|
| 34 |
+
p[base + 1] = 0x3F800000u;
|
| 35 |
+
}
|
| 36 |
+
}
|
| 37 |
+
uint32_t zero_ptr = 0;
|
| 38 |
+
if constexpr (OKt > 3 * Kt + 1) {
|
| 39 |
+
cb_reserve_back(cb_zero, 1);
|
| 40 |
+
zero_ptr = get_write_ptr(cb_zero);
|
| 41 |
+
auto* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(zero_ptr);
|
| 42 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 43 |
+
z[i] = 0;
|
| 44 |
+
}
|
| 45 |
+
}
|
| 46 |
+
for (uint32_t t = t0; t < t0 + n; ++t) {
|
| 47 |
+
const uint32_t r = t / Kt, j = t - r * Kt;
|
| 48 |
+
const uint32_t ob = r * OKt;
|
| 49 |
+
cb_wait_front(cb_hi, 1);
|
| 50 |
+
const uint32_t hp = get_read_ptr(cb_hi);
|
| 51 |
+
noc_async_write(hp, o.get_noc_addr(ob + j), TB);
|
| 52 |
+
noc_async_write(hp, o.get_noc_addr(ob + Kt + j), TB);
|
| 53 |
+
cb_wait_front(cb_lo, 1);
|
| 54 |
+
noc_async_write(get_read_ptr(cb_lo), o.get_noc_addr(ob + 2 * Kt + j), TB);
|
| 55 |
+
if (j == 0) {
|
| 56 |
+
noc_async_write(one_ptr, o.get_noc_addr(ob + 3 * Kt), TB);
|
| 57 |
+
for (uint32_t z = 3 * Kt + 1; z < OKt; ++z) {
|
| 58 |
+
noc_async_write(zero_ptr, o.get_noc_addr(ob + z), TB);
|
| 59 |
+
}
|
| 60 |
+
}
|
| 61 |
+
noc_async_writes_flushed();
|
| 62 |
+
cb_pop_front(cb_hi, 1);
|
| 63 |
+
cb_pop_front(cb_lo, 1);
|
| 64 |
+
}
|
| 65 |
+
noc_async_write_barrier();
|
| 66 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp
ADDED
|
@@ -0,0 +1,282 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Fused fp32 LayerNorm (tt/ln_kernel.py), compute: the exact op sequence of tt/layers.py layer_norm_fp32 (stock
|
| 3 |
+
// ttnn programs: mean -> sub -> mul -> mean -> add eps -> rsqrt -> mul -> mul gamma -> add beta), with the same SFPU
|
| 4 |
+
// LLK calls in the same order, so the result is meant to be bit-identical to the 9-program decomposition:
|
| 5 |
+
// - ttnn.mean (accurate fp32 path, reduce_op.cpp use_sfpu_fp32_mean): copy tile 0, add_binary_tile fold of tiles
|
| 6 |
+
// 1..Wt-1 in order, sfpu_reduce<SUM, Float32, REDUCE_ROW>, mul_unary_tile(1/W) (the AVG post-mul);
|
| 7 |
+
// - binary_ng fp32 ops: sub / add as sub_binary_tile / add_binary_tile<NearestEven>, mul as mul_binary_tile,
|
| 8 |
+
// the broadcast operand built by the dataflow (column / row fill), lhs in DST 0;
|
| 9 |
+
// - ttnn.rsqrt(fast_and_approximate_mode=False): rsqrt_tile<RsqrtMode::Default> (math_approx_mode false).
|
| 10 |
+
// Everything stays fp32 (UnpackToDestFp32 copies, fp32 DST, fp32 CBs): every intermediate the stock graph packs to
|
| 11 |
+
// an fp32 DRAM tensor round-trips losslessly.
|
| 12 |
+
//
|
| 13 |
+
// With has_res (LN_RESID, the residual add of the stream fused in front): h = x + res (* rgate), the stock
|
| 14 |
+
// `ttnn.add(x, ttnn.multiply(res, rgate))` / `ttnn.add(x, res)` as mul_binary_tile / add_binary_tile, packed to
|
| 15 |
+
// c_9 (the LayerNorm input) and to c_17 (written out as the new stream when write_h).
|
| 16 |
+
//
|
| 17 |
+
// CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] 1/W bits, [4] per-core RT-arg count (P3), [5] has_res,
|
| 18 |
+
// [6] has_rgate, [7] write_h, [9] sfpu_bcast (the statistics broadcast over the columns by the SFPU
|
| 19 |
+
// reduce itself, ln32_sfpu.h, and packed straight to c_5; else column 0 packed to c_4 and broadcast by the
|
| 20 |
+
// writer), [8] lean (one copy_init / SFPU-binary init per phase instead of one per
|
| 21 |
+
// tile: every CB copied from is an fp32 UnpackToDestFp32 CB, and add / sub / mul share one init).
|
| 22 |
+
// [10] res_t (LN_TR: the residual tiles arrive transposed and are transposed back with transpose_tile, the
|
| 23 |
+
// stock ttnn.transpose LLK, exact with the fp32 unpack-to-dest; no rgate), [11] out_t (the y tiles are
|
| 24 |
+
// transposed the same way, via c_18, before they are packed for the writer).
|
| 25 |
+
// Per-core RT args: [n_rows].
|
| 26 |
+
#include <cstdint>
|
| 27 |
+
|
| 28 |
+
#include "api/compute/common.h"
|
| 29 |
+
#include "api/compute/compute_kernel_api.h"
|
| 30 |
+
#include "api/compute/eltwise_binary_sfpu.h"
|
| 31 |
+
#include "api/compute/eltwise_unary/binop_with_scalar.h"
|
| 32 |
+
#include "api/compute/eltwise_unary/eltwise_unary.h"
|
| 33 |
+
#include "api/compute/eltwise_unary/rsqrt.h"
|
| 34 |
+
#include "api/compute/tile_move_copy.h"
|
| 35 |
+
#include "api/compute/transpose.h"
|
| 36 |
+
#include "ln32_sfpu.h"
|
| 37 |
+
|
| 38 |
+
using namespace ckernel;
|
| 39 |
+
|
| 40 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 41 |
+
constexpr uint32_t has_gamma = get_compile_time_arg_val(1);
|
| 42 |
+
constexpr uint32_t has_beta = get_compile_time_arg_val(2);
|
| 43 |
+
constexpr uint32_t inv_w_bits = get_compile_time_arg_val(3);
|
| 44 |
+
constexpr uint32_t has_res = get_compile_time_arg_val(5);
|
| 45 |
+
constexpr uint32_t has_rgate = get_compile_time_arg_val(6);
|
| 46 |
+
constexpr uint32_t write_h = get_compile_time_arg_val(7);
|
| 47 |
+
constexpr bool lean = get_compile_time_arg_val(8) != 0;
|
| 48 |
+
constexpr bool sfpu_bcast = get_compile_time_arg_val(9) != 0;
|
| 49 |
+
constexpr uint32_t res_t = get_compile_time_arg_val(10);
|
| 50 |
+
constexpr uint32_t out_t = get_compile_time_arg_val(11);
|
| 51 |
+
|
| 52 |
+
// per-tile (re-)inits, skipped in lean mode (done once at the start of the phase instead)
|
| 53 |
+
ALWI void ci(uint32_t cb) {
|
| 54 |
+
if constexpr (!lean) {
|
| 55 |
+
copy_init(cb);
|
| 56 |
+
}
|
| 57 |
+
}
|
| 58 |
+
#define BIN_INIT(op) \
|
| 59 |
+
do { \
|
| 60 |
+
if constexpr (!lean) { \
|
| 61 |
+
op##_binary_tile_init(); \
|
| 62 |
+
} \
|
| 63 |
+
} while (0)
|
| 64 |
+
|
| 65 |
+
constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_stat = 4, cb_bc = 5, cb_xc = 6, cb_r = 7, cb_rg = 8,
|
| 66 |
+
cb_h = 9, cb_out = 16, cb_hout = 17, cb_yt = 18;
|
| 67 |
+
constexpr uint32_t cb_src = has_res ? cb_h : cb_x;
|
| 68 |
+
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
|
| 69 |
+
|
| 70 |
+
// row mean of the Wt tiles of `cb` (already waited) into DST 0; `square` folds xc * xc instead of x
|
| 71 |
+
template <bool square>
|
| 72 |
+
ALWI void row_sum_to_dst0(uint32_t cb) {
|
| 73 |
+
copy_init(cb);
|
| 74 |
+
if constexpr (lean) {
|
| 75 |
+
add_binary_tile_init();
|
| 76 |
+
}
|
| 77 |
+
if constexpr (square) {
|
| 78 |
+
copy_tile(cb, 0, 0);
|
| 79 |
+
copy_tile(cb, 0, 1);
|
| 80 |
+
BIN_INIT(mul);
|
| 81 |
+
mul_binary_tile(0, 1, 0);
|
| 82 |
+
} else {
|
| 83 |
+
copy_tile(cb, 0, 0);
|
| 84 |
+
}
|
| 85 |
+
for (uint32_t w = 1; w < Wt; ++w) {
|
| 86 |
+
ci(cb);
|
| 87 |
+
copy_tile(cb, w, 1);
|
| 88 |
+
if constexpr (square) {
|
| 89 |
+
copy_tile(cb, w, 2);
|
| 90 |
+
BIN_INIT(mul);
|
| 91 |
+
mul_binary_tile(1, 2, 1);
|
| 92 |
+
}
|
| 93 |
+
BIN_INIT(add);
|
| 94 |
+
add_binary_tile(0, 1, 0);
|
| 95 |
+
}
|
| 96 |
+
sfpu_reduce_init<PoolType::SUM, DataFormat::Float32>();
|
| 97 |
+
if constexpr (sfpu_bcast) {
|
| 98 |
+
ln_row_sum_bcast_tile(0);
|
| 99 |
+
} else {
|
| 100 |
+
sfpu_reduce<PoolType::SUM, DataFormat::Float32, ReduceDim::REDUCE_ROW>(0, 1, 1);
|
| 101 |
+
}
|
| 102 |
+
binop_with_scalar_tile_init();
|
| 103 |
+
mul_unary_tile(0, inv_w_bits);
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
ALWI void pack_stat() {
|
| 107 |
+
constexpr uint32_t cb = sfpu_bcast ? cb_bc : cb_stat;
|
| 108 |
+
cb_reserve_back(cb, 1);
|
| 109 |
+
tile_regs_commit();
|
| 110 |
+
tile_regs_wait();
|
| 111 |
+
pack_tile(0, cb);
|
| 112 |
+
tile_regs_release();
|
| 113 |
+
cb_push_back(cb, 1);
|
| 114 |
+
}
|
| 115 |
+
|
| 116 |
+
void kernel_main() {
|
| 117 |
+
const uint32_t n_rows = get_arg_val<uint32_t>(0);
|
| 118 |
+
if (n_rows == 0) {
|
| 119 |
+
return;
|
| 120 |
+
}
|
| 121 |
+
compute_kernel_hw_startup(cb_x, cb_out);
|
| 122 |
+
// the constant tiles are waited for where they are first used (the reader sends x row 0 first)
|
| 123 |
+
for (uint32_t r = 0; r < n_rows; ++r) {
|
| 124 |
+
cb_wait_front(cb_x, Wt);
|
| 125 |
+
if constexpr (has_res) {
|
| 126 |
+
// h = x + res (* rgate)
|
| 127 |
+
cb_wait_front(cb_r, Wt);
|
| 128 |
+
if constexpr (has_rgate) {
|
| 129 |
+
cb_wait_front(cb_rg, Wt);
|
| 130 |
+
}
|
| 131 |
+
if constexpr (lean) {
|
| 132 |
+
copy_init(cb_r);
|
| 133 |
+
add_binary_tile_init();
|
| 134 |
+
}
|
| 135 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 136 |
+
tile_regs_acquire();
|
| 137 |
+
if constexpr (res_t) {
|
| 138 |
+
transpose_init(cb_r);
|
| 139 |
+
transpose_tile(cb_r, w, 0);
|
| 140 |
+
copy_init(cb_x);
|
| 141 |
+
add_binary_tile_init();
|
| 142 |
+
} else {
|
| 143 |
+
ci(cb_r);
|
| 144 |
+
copy_tile(cb_r, w, 0);
|
| 145 |
+
}
|
| 146 |
+
if constexpr (has_rgate) {
|
| 147 |
+
ci(cb_rg);
|
| 148 |
+
copy_tile(cb_rg, w, 1);
|
| 149 |
+
BIN_INIT(mul);
|
| 150 |
+
mul_binary_tile(0, 1, 0);
|
| 151 |
+
}
|
| 152 |
+
ci(cb_x);
|
| 153 |
+
copy_tile(cb_x, w, 1);
|
| 154 |
+
BIN_INIT(add);
|
| 155 |
+
add_binary_tile<RNE>(1, 0, 1);
|
| 156 |
+
cb_reserve_back(cb_h, 1);
|
| 157 |
+
if constexpr (write_h) {
|
| 158 |
+
cb_reserve_back(cb_hout, 1);
|
| 159 |
+
}
|
| 160 |
+
tile_regs_commit();
|
| 161 |
+
tile_regs_wait();
|
| 162 |
+
pack_tile(1, cb_h);
|
| 163 |
+
if constexpr (write_h) {
|
| 164 |
+
pack_tile(1, cb_hout);
|
| 165 |
+
}
|
| 166 |
+
tile_regs_release();
|
| 167 |
+
cb_push_back(cb_h, 1);
|
| 168 |
+
if constexpr (write_h) {
|
| 169 |
+
cb_push_back(cb_hout, 1);
|
| 170 |
+
}
|
| 171 |
+
}
|
| 172 |
+
cb_pop_front(cb_r, Wt);
|
| 173 |
+
cb_pop_front(cb_x, Wt);
|
| 174 |
+
cb_wait_front(cb_h, Wt);
|
| 175 |
+
}
|
| 176 |
+
// mean
|
| 177 |
+
tile_regs_acquire();
|
| 178 |
+
row_sum_to_dst0<false>(cb_src);
|
| 179 |
+
pack_stat();
|
| 180 |
+
|
| 181 |
+
// xc = x - mean
|
| 182 |
+
cb_wait_front(cb_bc, 1);
|
| 183 |
+
if constexpr (lean) {
|
| 184 |
+
copy_init(cb_src);
|
| 185 |
+
sub_binary_tile_init();
|
| 186 |
+
}
|
| 187 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 188 |
+
tile_regs_acquire();
|
| 189 |
+
ci(cb_src);
|
| 190 |
+
copy_tile(cb_src, w, 0);
|
| 191 |
+
ci(cb_bc);
|
| 192 |
+
copy_tile(cb_bc, 0, 1);
|
| 193 |
+
BIN_INIT(sub);
|
| 194 |
+
sub_binary_tile<RNE>(0, 1, 0);
|
| 195 |
+
cb_reserve_back(cb_xc, 1);
|
| 196 |
+
tile_regs_commit();
|
| 197 |
+
tile_regs_wait();
|
| 198 |
+
pack_tile(0, cb_xc);
|
| 199 |
+
tile_regs_release();
|
| 200 |
+
cb_push_back(cb_xc, 1);
|
| 201 |
+
}
|
| 202 |
+
cb_pop_front(cb_bc, 1);
|
| 203 |
+
cb_pop_front(cb_src, Wt);
|
| 204 |
+
|
| 205 |
+
// rstd = rsqrt(mean(xc * xc) + eps)
|
| 206 |
+
cb_wait_front(cb_xc, Wt);
|
| 207 |
+
tile_regs_acquire();
|
| 208 |
+
row_sum_to_dst0<true>(cb_xc);
|
| 209 |
+
cb_wait_front(cb_eps, 1);
|
| 210 |
+
copy_init(cb_eps);
|
| 211 |
+
copy_tile(cb_eps, 0, 1);
|
| 212 |
+
add_binary_tile_init();
|
| 213 |
+
add_binary_tile<RNE>(0, 1, 0);
|
| 214 |
+
rsqrt_tile_init();
|
| 215 |
+
rsqrt_tile<RsqrtMode::Default>(0);
|
| 216 |
+
pack_stat();
|
| 217 |
+
|
| 218 |
+
// y = xc * rstd (* gamma) (+ beta)
|
| 219 |
+
cb_wait_front(cb_bc, 1);
|
| 220 |
+
if constexpr (has_gamma) {
|
| 221 |
+
cb_wait_front(cb_g, Wt);
|
| 222 |
+
}
|
| 223 |
+
if constexpr (has_beta) {
|
| 224 |
+
cb_wait_front(cb_b, Wt);
|
| 225 |
+
}
|
| 226 |
+
if constexpr (lean) {
|
| 227 |
+
copy_init(cb_xc);
|
| 228 |
+
mul_binary_tile_init();
|
| 229 |
+
}
|
| 230 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 231 |
+
tile_regs_acquire();
|
| 232 |
+
ci(cb_xc);
|
| 233 |
+
copy_tile(cb_xc, w, 0);
|
| 234 |
+
ci(cb_bc);
|
| 235 |
+
copy_tile(cb_bc, 0, 1);
|
| 236 |
+
BIN_INIT(mul);
|
| 237 |
+
mul_binary_tile(0, 1, 0);
|
| 238 |
+
if constexpr (has_gamma) {
|
| 239 |
+
ci(cb_g);
|
| 240 |
+
copy_tile(cb_g, w, 1);
|
| 241 |
+
BIN_INIT(mul);
|
| 242 |
+
mul_binary_tile(0, 1, 0);
|
| 243 |
+
}
|
| 244 |
+
if constexpr (has_beta) {
|
| 245 |
+
ci(cb_b);
|
| 246 |
+
copy_tile(cb_b, w, 1);
|
| 247 |
+
BIN_INIT(add);
|
| 248 |
+
add_binary_tile<RNE>(0, 1, 0);
|
| 249 |
+
}
|
| 250 |
+
if constexpr (out_t) {
|
| 251 |
+
cb_reserve_back(cb_yt, 1);
|
| 252 |
+
tile_regs_commit();
|
| 253 |
+
tile_regs_wait();
|
| 254 |
+
pack_tile(0, cb_yt);
|
| 255 |
+
tile_regs_release();
|
| 256 |
+
cb_push_back(cb_yt, 1);
|
| 257 |
+
cb_wait_front(cb_yt, 1);
|
| 258 |
+
tile_regs_acquire();
|
| 259 |
+
transpose_init(cb_yt);
|
| 260 |
+
transpose_tile(cb_yt, 0, 0);
|
| 261 |
+
cb_reserve_back(cb_out, 1);
|
| 262 |
+
tile_regs_commit();
|
| 263 |
+
tile_regs_wait();
|
| 264 |
+
pack_tile(0, cb_out);
|
| 265 |
+
tile_regs_release();
|
| 266 |
+
cb_push_back(cb_out, 1);
|
| 267 |
+
cb_pop_front(cb_yt, 1);
|
| 268 |
+
copy_init(cb_xc);
|
| 269 |
+
mul_binary_tile_init();
|
| 270 |
+
} else {
|
| 271 |
+
cb_reserve_back(cb_out, 1);
|
| 272 |
+
tile_regs_commit();
|
| 273 |
+
tile_regs_wait();
|
| 274 |
+
pack_tile(0, cb_out);
|
| 275 |
+
tile_regs_release();
|
| 276 |
+
cb_push_back(cb_out, 1);
|
| 277 |
+
}
|
| 278 |
+
}
|
| 279 |
+
cb_pop_front(cb_bc, 1);
|
| 280 |
+
cb_pop_front(cb_xc, Wt);
|
| 281 |
+
}
|
| 282 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32_reader.cpp
ADDED
|
@@ -0,0 +1,123 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Fused fp32 LayerNorm (tt/ln_kernel.py), reader (RISCV_0): the affine rows once, then this core's tile rows of x.
|
| 3 |
+
//
|
| 4 |
+
// CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits (fp32), [4] per-core RT-arg count (P3: part of the
|
| 5 |
+
// program hash), [5] has_res, [6] has_rgate, [7] res_t (LN_TR: res is stored transposed per entity, i.e.
|
| 6 |
+
// [.., E, W, T]: tile (r, w) of res is read from tile (e, w, tr) of that layout, e = r / Tt, tr = r % Tt;
|
| 7 |
+
// the compute transposes it back), [8] Tt (tile rows per entity), then the TensorAccessorArgs of x, gamma,
|
| 8 |
+
// beta, res, rgate (x's when absent).
|
| 9 |
+
// Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr]. Per-core RT args: [row0, n_rows].
|
| 10 |
+
// CBs: c_0 x (Wt tiles per row, double-buffered), c_1 gamma rows (Wt tiles, row 0 replicated over the tile),
|
| 11 |
+
// c_2 beta rows (same), c_3 the eps tile (every element eps), c_7 res (as c_0), c_8 the residual gate rows.
|
| 12 |
+
#include <cstdint>
|
| 13 |
+
|
| 14 |
+
#include "api/dataflow/dataflow_api.h"
|
| 15 |
+
|
| 16 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 17 |
+
constexpr uint32_t has_gamma = get_compile_time_arg_val(1);
|
| 18 |
+
constexpr uint32_t has_beta = get_compile_time_arg_val(2);
|
| 19 |
+
constexpr uint32_t eps_bits = get_compile_time_arg_val(3);
|
| 20 |
+
constexpr uint32_t has_res = get_compile_time_arg_val(5);
|
| 21 |
+
constexpr uint32_t has_rgate = get_compile_time_arg_val(6);
|
| 22 |
+
constexpr uint32_t res_t = get_compile_time_arg_val(7);
|
| 23 |
+
constexpr uint32_t Tt = get_compile_time_arg_val(8);
|
| 24 |
+
constexpr auto x_args = TensorAccessorArgs<9>();
|
| 25 |
+
constexpr auto g_args = TensorAccessorArgs<x_args.next_compile_time_args_offset()>();
|
| 26 |
+
constexpr auto b_args = TensorAccessorArgs<g_args.next_compile_time_args_offset()>();
|
| 27 |
+
constexpr auto r_args = TensorAccessorArgs<b_args.next_compile_time_args_offset()>();
|
| 28 |
+
constexpr auto rg_args = TensorAccessorArgs<r_args.next_compile_time_args_offset()>();
|
| 29 |
+
|
| 30 |
+
constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_r = 7, cb_rg = 8;
|
| 31 |
+
constexpr uint32_t TB = 4096; // fp32 tile bytes
|
| 32 |
+
|
| 33 |
+
// Replicate row 0 of each of the n fp32 tiles at p over its 32 rows (binary_ng fill_tile_with_first_row: local
|
| 34 |
+
// NoC doubling), all tiles per doubling step under one barrier.
|
| 35 |
+
FORCE_INLINE void fill_first_rows(uint32_t p, uint32_t n_tiles) {
|
| 36 |
+
constexpr uint32_t kFace = 1024, kRow = 64;
|
| 37 |
+
for (uint32_t n = kRow; n < kFace; n <<= 1) {
|
| 38 |
+
for (uint32_t t = 0; t < n_tiles; ++t) {
|
| 39 |
+
const uint32_t f0 = p + t * TB, f1 = f0 + kFace;
|
| 40 |
+
noc_async_read(get_noc_addr(f0), f0 + n, n);
|
| 41 |
+
noc_async_read(get_noc_addr(f1), f1 + n, n);
|
| 42 |
+
}
|
| 43 |
+
noc_async_read_barrier();
|
| 44 |
+
}
|
| 45 |
+
for (uint32_t t = 0; t < n_tiles; ++t) {
|
| 46 |
+
const uint32_t f0 = p + t * TB;
|
| 47 |
+
noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace);
|
| 48 |
+
}
|
| 49 |
+
noc_async_read_barrier();
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
template <typename A>
|
| 53 |
+
FORCE_INLINE void read_rows(uint32_t cb, const A& acc) {
|
| 54 |
+
cb_reserve_back(cb, Wt);
|
| 55 |
+
uint32_t p = get_write_ptr(cb);
|
| 56 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 57 |
+
noc_async_read(acc.get_noc_addr(w), p + w * TB, TB);
|
| 58 |
+
}
|
| 59 |
+
noc_async_read_barrier();
|
| 60 |
+
fill_first_rows(p, Wt);
|
| 61 |
+
cb_push_back(cb, Wt);
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
template <typename A, typename R>
|
| 65 |
+
FORCE_INLINE void read_x_row(uint32_t r, const A& x, const R& res) {
|
| 66 |
+
cb_reserve_back(cb_x, Wt);
|
| 67 |
+
uint32_t p = get_write_ptr(cb_x);
|
| 68 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 69 |
+
noc_async_read(x.get_noc_addr(r * Wt + w), p + w * TB, TB);
|
| 70 |
+
}
|
| 71 |
+
if constexpr (has_res) {
|
| 72 |
+
cb_reserve_back(cb_r, Wt);
|
| 73 |
+
uint32_t q = get_write_ptr(cb_r);
|
| 74 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 75 |
+
const uint32_t ri = res_t ? (r / Tt) * (Wt * Tt) + w * Tt + (r % Tt) : r * Wt + w;
|
| 76 |
+
noc_async_read(res.get_noc_addr(ri), q + w * TB, TB);
|
| 77 |
+
}
|
| 78 |
+
noc_async_read_barrier();
|
| 79 |
+
cb_push_back(cb_r, Wt);
|
| 80 |
+
} else {
|
| 81 |
+
noc_async_read_barrier();
|
| 82 |
+
}
|
| 83 |
+
cb_push_back(cb_x, Wt);
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
void kernel_main() {
|
| 87 |
+
const uint32_t x_addr = get_common_arg_val<uint32_t>(0);
|
| 88 |
+
const uint32_t g_addr = get_common_arg_val<uint32_t>(1);
|
| 89 |
+
const uint32_t b_addr = get_common_arg_val<uint32_t>(2);
|
| 90 |
+
const uint32_t r_addr = get_common_arg_val<uint32_t>(3);
|
| 91 |
+
const uint32_t rg_addr = get_common_arg_val<uint32_t>(4);
|
| 92 |
+
const uint32_t row0 = get_arg_val<uint32_t>(0);
|
| 93 |
+
const uint32_t n_rows = get_arg_val<uint32_t>(1);
|
| 94 |
+
if (n_rows == 0) {
|
| 95 |
+
return;
|
| 96 |
+
}
|
| 97 |
+
const auto x = TensorAccessor(x_args, x_addr, TB);
|
| 98 |
+
const auto res = TensorAccessor(r_args, r_addr, TB);
|
| 99 |
+
|
| 100 |
+
// the first x row goes first (the compute starts on it); the affine rows are needed only by its last phase
|
| 101 |
+
read_x_row(row0, x, res);
|
| 102 |
+
// eps tile
|
| 103 |
+
cb_reserve_back(cb_eps, 1);
|
| 104 |
+
{
|
| 105 |
+
auto* e = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_eps));
|
| 106 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 107 |
+
e[i] = eps_bits;
|
| 108 |
+
}
|
| 109 |
+
}
|
| 110 |
+
cb_push_back(cb_eps, 1);
|
| 111 |
+
if constexpr (has_rgate) {
|
| 112 |
+
read_rows(cb_rg, TensorAccessor(rg_args, rg_addr, TB));
|
| 113 |
+
}
|
| 114 |
+
if constexpr (has_gamma) {
|
| 115 |
+
read_rows(cb_g, TensorAccessor(g_args, g_addr, TB));
|
| 116 |
+
}
|
| 117 |
+
if constexpr (has_beta) {
|
| 118 |
+
read_rows(cb_b, TensorAccessor(b_args, b_addr, TB));
|
| 119 |
+
}
|
| 120 |
+
for (uint32_t r = row0 + 1; r < row0 + n_rows; ++r) {
|
| 121 |
+
read_x_row(r, x, res);
|
| 122 |
+
}
|
| 123 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32_sfpu.h
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// SFPU helper of the fused fp32 LayerNorm kernels (ln32_compute.cpp, ln32s_compute.cpp): the row sum of one fp32
|
| 3 |
+
// tile with the result in EVERY column.
|
| 4 |
+
//
|
| 5 |
+
// ln_row_sum_bcast is perform_reduce_row_sum_tile of tt-metal's blackhole ckernel_sfpu_reduce.h (the REDUCE_ROW SUM
|
| 6 |
+
// Float32 single-tile path behind sfpu_reduce<SUM, Float32, REDUCE_ROW>, used by the accurate ttnn.mean): the same
|
| 7 |
+
// loads, the same replayed vertical adds (the replay recorded by sfpu_reduce_init<SUM, Float32>) and the same
|
| 8 |
+
// horizontal butterfly, which leaves each row's full sum in all 8 lanes of its sub-vector. The stock kernel stores
|
| 9 |
+
// that register to the even columns of the left face only (column 0 is what the packer keeps); this one stores it to
|
| 10 |
+
// the four column groups (even / odd x left / right face), so the tile is the column-0 result broadcast over the
|
| 11 |
+
// row: bit for bit the values the stock graph broadcasts with binary_ng's column fill, without a RISC-V fill loop.
|
| 12 |
+
#pragma once
|
| 13 |
+
|
| 14 |
+
#ifdef TRISC_MATH
|
| 15 |
+
namespace ckernel::sfpu {
|
| 16 |
+
|
| 17 |
+
template <bool is_fp32_dest_acc_en>
|
| 18 |
+
inline void ln_row_sum_bcast() {
|
| 19 |
+
constexpr InstrModLoadStore MODE = GetSfpLoadStoreInstrMod<DataFormat::Float32, is_fp32_dest_acc_en>();
|
| 20 |
+
for (std::uint32_t face_pair = 0; face_pair < 2; face_pair++) {
|
| 21 |
+
const std::uint32_t face_pair_base = face_pair * 2 * ROWS_PER_FACE;
|
| 22 |
+
for (std::uint32_t row_group = 0; row_group < 2; row_group++) {
|
| 23 |
+
const std::uint32_t a = face_pair_base + row_group * 8;
|
| 24 |
+
const std::uint32_t b = a + 4;
|
| 25 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG0, MODE, ADDR_MOD_7, a);
|
| 26 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG1, MODE, ADDR_MOD_7, a + 2);
|
| 27 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG2, MODE, ADDR_MOD_7, a + ROWS_PER_FACE);
|
| 28 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG3, MODE, ADDR_MOD_7, a + ROWS_PER_FACE + 2);
|
| 29 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG4, MODE, ADDR_MOD_7, b);
|
| 30 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG5, MODE, ADDR_MOD_7, b + 2);
|
| 31 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG6, MODE, ADDR_MOD_7, b + ROWS_PER_FACE);
|
| 32 |
+
load_and_clear_high_bits<false>(p_sfpu::LREG7, MODE, ADDR_MOD_7, b + ROWS_PER_FACE + 2);
|
| 33 |
+
lltt::replay(0, 6);
|
| 34 |
+
horizontal_reduce<false>();
|
| 35 |
+
TT_SFPSTORE(p_sfpu::LREG0, MODE, ADDR_MOD_7, a);
|
| 36 |
+
TT_SFPSTORE(p_sfpu::LREG4, MODE, ADDR_MOD_7, b);
|
| 37 |
+
TT_SFPSTORE(p_sfpu::LREG0, MODE, ADDR_MOD_7, a + 2);
|
| 38 |
+
TT_SFPSTORE(p_sfpu::LREG4, MODE, ADDR_MOD_7, b + 2);
|
| 39 |
+
TT_SFPSTORE(p_sfpu::LREG0, MODE, ADDR_MOD_7, a + ROWS_PER_FACE);
|
| 40 |
+
TT_SFPSTORE(p_sfpu::LREG4, MODE, ADDR_MOD_7, b + ROWS_PER_FACE);
|
| 41 |
+
TT_SFPSTORE(p_sfpu::LREG0, MODE, ADDR_MOD_7, a + ROWS_PER_FACE + 2);
|
| 42 |
+
TT_SFPSTORE(p_sfpu::LREG4, MODE, ADDR_MOD_7, b + ROWS_PER_FACE + 2);
|
| 43 |
+
}
|
| 44 |
+
}
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
} // namespace ckernel::sfpu
|
| 48 |
+
#endif
|
| 49 |
+
|
| 50 |
+
// row sums of the fp32 tile in DST idst, broadcast over all columns (sfpu_reduce_init<SUM, Float32> first)
|
| 51 |
+
ALWI void ln_row_sum_bcast_tile(uint32_t idst) {
|
| 52 |
+
MATH(SFPU_UNARY_CALL(DST_SYNC_MODE, DST_ACCUM_MODE, ln_row_sum_bcast, (DST_ACCUM_MODE), idst,
|
| 53 |
+
VectorMode::RC_custom));
|
| 54 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32_writer.cpp
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Fused fp32 LayerNorm (tt/ln_kernel.py), writer (RISCV_1): per tile row, broadcasts the two statistic tiles the
|
| 3 |
+
// compute packs (column 0 = the row's mean, then its rsqrt(var + eps)) over all 32 columns, then drains the Wt
|
| 4 |
+
// output tiles to DRAM.
|
| 5 |
+
//
|
| 6 |
+
// With write_h (LN_RESID) the Wt tiles of the new stream h = x + res (* rgate) (c_17) go first, to h_addr.
|
| 7 |
+
//
|
| 8 |
+
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] write_h, [3] sfpu_bcast (no fill here), [4] out_t (LN_TR:
|
| 9 |
+
// the output tiles, transposed by the compute, go to the per-entity transposed layout [.., E, W, T]: tile
|
| 10 |
+
// (r, w) -> (e, w, tr), e = r / Tt, tr = r % Tt), [5] Tt, then the TensorAccessorArgs of out and h (out's
|
| 11 |
+
// when absent).
|
| 12 |
+
// Common RT args: [out_addr, h_addr]. Per-core RT args: [row0, n_rows].
|
| 13 |
+
// CBs: c_4 statistic (from compute, column 0 valid), c_5 broadcast statistic (to compute), c_16 out, c_17 h.
|
| 14 |
+
#include <cstdint>
|
| 15 |
+
|
| 16 |
+
#include "api/dataflow/dataflow_api.h"
|
| 17 |
+
|
| 18 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 19 |
+
constexpr uint32_t write_h = get_compile_time_arg_val(2);
|
| 20 |
+
constexpr uint32_t sfpu_bcast = get_compile_time_arg_val(3); // the compute broadcasts the statistics itself
|
| 21 |
+
constexpr uint32_t out_t = get_compile_time_arg_val(4);
|
| 22 |
+
constexpr uint32_t Tt = get_compile_time_arg_val(5);
|
| 23 |
+
constexpr auto out_args = TensorAccessorArgs<6>();
|
| 24 |
+
constexpr auto h_args = TensorAccessorArgs<out_args.next_compile_time_args_offset()>();
|
| 25 |
+
|
| 26 |
+
constexpr uint32_t cb_stat = 4, cb_bc = 5, cb_out = 16, cb_hout = 17;
|
| 27 |
+
constexpr uint32_t TB = 4096;
|
| 28 |
+
|
| 29 |
+
// dst tile = src tile's column 0 replicated over the 32 columns (binary_ng fill_tile_with_first_column, fp32;
|
| 30 |
+
// faces 0/1 take face 0's column 0, faces 2/3 face 2's).
|
| 31 |
+
FORCE_INLINE void bcast_col(uint32_t src, uint32_t dst) {
|
| 32 |
+
auto* s = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(src);
|
| 33 |
+
auto* d = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(dst);
|
| 34 |
+
for (uint32_t fo = 0; fo < 1024; fo += 512) {
|
| 35 |
+
for (uint32_t ro = 0; ro < 256; ro += 16) {
|
| 36 |
+
const uint32_t v = s[fo + ro];
|
| 37 |
+
volatile tt_l1_ptr uint32_t* l = d + fo + ro;
|
| 38 |
+
volatile tt_l1_ptr uint32_t* r = l + 256;
|
| 39 |
+
for (uint32_t c = 0; c < 16; ++c) {
|
| 40 |
+
l[c] = v;
|
| 41 |
+
r[c] = v;
|
| 42 |
+
}
|
| 43 |
+
}
|
| 44 |
+
}
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
FORCE_INLINE void stat() {
|
| 48 |
+
cb_wait_front(cb_stat, 1);
|
| 49 |
+
cb_reserve_back(cb_bc, 1);
|
| 50 |
+
bcast_col(get_read_ptr(cb_stat), get_write_ptr(cb_bc));
|
| 51 |
+
cb_push_back(cb_bc, 1);
|
| 52 |
+
cb_pop_front(cb_stat, 1);
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
void kernel_main() {
|
| 56 |
+
const uint32_t out_addr = get_common_arg_val<uint32_t>(0);
|
| 57 |
+
const uint32_t h_addr = get_common_arg_val<uint32_t>(1);
|
| 58 |
+
const auto hacc = TensorAccessor(h_args, h_addr, TB);
|
| 59 |
+
const uint32_t row0 = get_arg_val<uint32_t>(0);
|
| 60 |
+
const uint32_t n_rows = get_arg_val<uint32_t>(1);
|
| 61 |
+
const auto out = TensorAccessor(out_args, out_addr, TB);
|
| 62 |
+
for (uint32_t r = row0; r < row0 + n_rows; ++r) {
|
| 63 |
+
if constexpr (write_h) {
|
| 64 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 65 |
+
cb_wait_front(cb_hout, 1);
|
| 66 |
+
noc_async_write(get_read_ptr(cb_hout), hacc.get_noc_addr(r * Wt + w), TB);
|
| 67 |
+
noc_async_writes_flushed();
|
| 68 |
+
cb_pop_front(cb_hout, 1);
|
| 69 |
+
}
|
| 70 |
+
}
|
| 71 |
+
if constexpr (!sfpu_bcast) {
|
| 72 |
+
stat(); // mean
|
| 73 |
+
stat(); // rsqrt(var + eps)
|
| 74 |
+
}
|
| 75 |
+
for (uint32_t w = 0; w < Wt; ++w) {
|
| 76 |
+
cb_wait_front(cb_out, 1);
|
| 77 |
+
const uint32_t oi = out_t ? (r / Tt) * (Wt * Tt) + w * Tt + (r % Tt) : r * Wt + w;
|
| 78 |
+
noc_async_write(get_read_ptr(cb_out), out.get_noc_addr(oi), TB);
|
| 79 |
+
noc_async_writes_flushed();
|
| 80 |
+
cb_pop_front(cb_out, 1);
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
noc_async_write_barrier();
|
| 84 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), compute: member j of a tile row owns column tile j.
|
| 3 |
+
// The same LLK sequence as ln32_compute.cpp (and the stock 9-program decomposition): the folds over the Wt tiles run
|
| 4 |
+
// on the root, in tile order, over the gathered tiles (c_10), so every statistic is bit for bit the stock one.
|
| 5 |
+
// h = x (+ res (* rgate)) -> c_11 (to the root's gather) [+ c_17 for write_h]
|
| 6 |
+
// root: mean = fold(c_10) / W -> c_4 (the writer broadcasts it into every member's c_5 page 0)
|
| 7 |
+
// xc = h - mean; sq = xc * xc -> c_6 (kept), sq -> c_11
|
| 8 |
+
// root: rstd = rsqrt(fold(c_10) / W + eps) -> c_4 (broadcast into c_5 page 1)
|
| 9 |
+
// y = xc * rstd (* gamma) (+ beta) -> c_16
|
| 10 |
+
// CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] 1/W bits, [4] per-core RT-arg count (P3), [5] has_res,
|
| 11 |
+
// [6] has_rgate, [7] write_h, [8] sfpu_bcast (the root's statistics broadcast over the columns by the
|
| 12 |
+
// SFPU reduce, ln32_sfpu.h; else column 0 and the writer fills), [9] kcat_out (KCAT_EMIT: y is emitted
|
| 13 |
+
// as the split operand of the next K-concatenated linear, y_hi = bf16(y) -> c_16 and y_lo = y - y_hi ->
|
| 14 |
+
// c_19, the LLK calls of kernels/kcat_compute.cpp on the packed fp32 y tile c_18).
|
| 15 |
+
// Per-core RT args: [n_rows, is_root].
|
| 16 |
+
#include <cstdint>
|
| 17 |
+
|
| 18 |
+
#include "api/compute/common.h"
|
| 19 |
+
#include "api/compute/compute_kernel_api.h"
|
| 20 |
+
#include "api/compute/eltwise_binary_sfpu.h"
|
| 21 |
+
#include "api/compute/eltwise_unary/binop_with_scalar.h"
|
| 22 |
+
#include "api/compute/eltwise_unary/eltwise_unary.h"
|
| 23 |
+
#include "api/compute/eltwise_unary/rsqrt.h"
|
| 24 |
+
#include "api/compute/eltwise_unary/typecast.h"
|
| 25 |
+
#include "api/compute/tile_move_copy.h"
|
| 26 |
+
#include "ln32_sfpu.h"
|
| 27 |
+
|
| 28 |
+
using namespace ckernel;
|
| 29 |
+
|
| 30 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 31 |
+
constexpr uint32_t has_gamma = get_compile_time_arg_val(1);
|
| 32 |
+
constexpr uint32_t has_beta = get_compile_time_arg_val(2);
|
| 33 |
+
constexpr uint32_t inv_w_bits = get_compile_time_arg_val(3);
|
| 34 |
+
constexpr uint32_t has_res = get_compile_time_arg_val(5);
|
| 35 |
+
constexpr uint32_t has_rgate = get_compile_time_arg_val(6);
|
| 36 |
+
constexpr uint32_t write_h = get_compile_time_arg_val(7);
|
| 37 |
+
constexpr bool sfpu_bcast = get_compile_time_arg_val(8) != 0;
|
| 38 |
+
constexpr uint32_t kcat_out = get_compile_time_arg_val(9);
|
| 39 |
+
|
| 40 |
+
constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_stat = 4, cb_bc = 5, cb_xc = 6, cb_r = 7, cb_rg = 8,
|
| 41 |
+
cb_h = 9, cb_gat = 10, cb_snd = 11, cb_out = 16, cb_hout = 17, cb_y = 18, cb_lo = 19;
|
| 42 |
+
constexpr uint32_t cb_src = has_res ? cb_h : cb_x;
|
| 43 |
+
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
|
| 44 |
+
|
| 45 |
+
// fold of the Wt gathered tiles (in order) -> row sum * 1/W in DST 0
|
| 46 |
+
ALWI void fold_gathered() {
|
| 47 |
+
cb_wait_front(cb_gat, Wt);
|
| 48 |
+
copy_init(cb_gat);
|
| 49 |
+
add_binary_tile_init();
|
| 50 |
+
copy_tile(cb_gat, 0, 0);
|
| 51 |
+
for (uint32_t w = 1; w < Wt; ++w) {
|
| 52 |
+
copy_tile(cb_gat, w, 1);
|
| 53 |
+
add_binary_tile(0, 1, 0);
|
| 54 |
+
}
|
| 55 |
+
sfpu_reduce_init<PoolType::SUM, DataFormat::Float32>();
|
| 56 |
+
if constexpr (sfpu_bcast) {
|
| 57 |
+
ln_row_sum_bcast_tile(0);
|
| 58 |
+
} else {
|
| 59 |
+
sfpu_reduce<PoolType::SUM, DataFormat::Float32, ReduceDim::REDUCE_ROW>(0, 1, 1);
|
| 60 |
+
}
|
| 61 |
+
binop_with_scalar_tile_init();
|
| 62 |
+
mul_unary_tile(0, inv_w_bits);
|
| 63 |
+
}
|
| 64 |
+
|
| 65 |
+
ALWI void pack_one(uint32_t dst, uint32_t cb) {
|
| 66 |
+
cb_reserve_back(cb, 1);
|
| 67 |
+
tile_regs_commit();
|
| 68 |
+
tile_regs_wait();
|
| 69 |
+
pack_tile(dst, cb);
|
| 70 |
+
tile_regs_release();
|
| 71 |
+
cb_push_back(cb, 1);
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
void kernel_main() {
|
| 75 |
+
const uint32_t n_rows = get_arg_val<uint32_t>(0);
|
| 76 |
+
const bool root = get_arg_val<uint32_t>(1) != 0;
|
| 77 |
+
if (n_rows == 0) {
|
| 78 |
+
return;
|
| 79 |
+
}
|
| 80 |
+
compute_kernel_hw_startup(cb_x, cb_out);
|
| 81 |
+
for (uint32_t r = 0; r < n_rows; ++r) {
|
| 82 |
+
// h (this member's tile) -> c_11 for the gather (and c_9 / c_17)
|
| 83 |
+
cb_wait_front(cb_x, 1);
|
| 84 |
+
tile_regs_acquire();
|
| 85 |
+
if constexpr (has_res) {
|
| 86 |
+
cb_wait_front(cb_r, 1);
|
| 87 |
+
if constexpr (has_rgate) {
|
| 88 |
+
cb_wait_front(cb_rg, 1);
|
| 89 |
+
}
|
| 90 |
+
copy_init(cb_r);
|
| 91 |
+
copy_tile(cb_r, 0, 0);
|
| 92 |
+
if constexpr (has_rgate) {
|
| 93 |
+
copy_tile(cb_rg, 0, 1);
|
| 94 |
+
mul_binary_tile_init();
|
| 95 |
+
mul_binary_tile(0, 1, 0);
|
| 96 |
+
}
|
| 97 |
+
copy_tile(cb_x, 0, 1);
|
| 98 |
+
add_binary_tile_init();
|
| 99 |
+
add_binary_tile<RNE>(1, 0, 1);
|
| 100 |
+
cb_reserve_back(cb_h, 1);
|
| 101 |
+
cb_reserve_back(cb_snd, 1);
|
| 102 |
+
if constexpr (write_h) {
|
| 103 |
+
cb_reserve_back(cb_hout, 1);
|
| 104 |
+
}
|
| 105 |
+
tile_regs_commit();
|
| 106 |
+
tile_regs_wait();
|
| 107 |
+
pack_tile(1, cb_h);
|
| 108 |
+
pack_tile(1, cb_snd);
|
| 109 |
+
if constexpr (write_h) {
|
| 110 |
+
pack_tile(1, cb_hout);
|
| 111 |
+
}
|
| 112 |
+
tile_regs_release();
|
| 113 |
+
cb_push_back(cb_h, 1);
|
| 114 |
+
cb_push_back(cb_snd, 1);
|
| 115 |
+
if constexpr (write_h) {
|
| 116 |
+
cb_push_back(cb_hout, 1);
|
| 117 |
+
}
|
| 118 |
+
cb_pop_front(cb_r, 1);
|
| 119 |
+
cb_pop_front(cb_x, 1);
|
| 120 |
+
cb_wait_front(cb_h, 1);
|
| 121 |
+
} else {
|
| 122 |
+
copy_init(cb_x);
|
| 123 |
+
copy_tile(cb_x, 0, 0);
|
| 124 |
+
pack_one(0, cb_snd);
|
| 125 |
+
}
|
| 126 |
+
// root: mean
|
| 127 |
+
if (root) {
|
| 128 |
+
tile_regs_acquire();
|
| 129 |
+
fold_gathered();
|
| 130 |
+
pack_one(0, cb_stat);
|
| 131 |
+
cb_pop_front(cb_gat, Wt);
|
| 132 |
+
}
|
| 133 |
+
// xc = h - mean -> c_6; sq = xc * xc -> c_11
|
| 134 |
+
cb_wait_front(cb_bc, 1);
|
| 135 |
+
tile_regs_acquire();
|
| 136 |
+
copy_init(cb_src);
|
| 137 |
+
copy_tile(cb_src, 0, 0);
|
| 138 |
+
copy_tile(cb_bc, 0, 1);
|
| 139 |
+
sub_binary_tile_init();
|
| 140 |
+
sub_binary_tile<RNE>(0, 1, 0);
|
| 141 |
+
pack_one(0, cb_xc);
|
| 142 |
+
cb_pop_front(cb_bc, 1);
|
| 143 |
+
cb_pop_front(cb_src, 1);
|
| 144 |
+
cb_wait_front(cb_xc, 1);
|
| 145 |
+
tile_regs_acquire();
|
| 146 |
+
copy_init(cb_xc);
|
| 147 |
+
copy_tile(cb_xc, 0, 0);
|
| 148 |
+
copy_tile(cb_xc, 0, 1);
|
| 149 |
+
mul_binary_tile_init();
|
| 150 |
+
mul_binary_tile(0, 1, 0);
|
| 151 |
+
pack_one(0, cb_snd);
|
| 152 |
+
// root: rstd
|
| 153 |
+
if (root) {
|
| 154 |
+
tile_regs_acquire();
|
| 155 |
+
fold_gathered();
|
| 156 |
+
cb_wait_front(cb_eps, 1);
|
| 157 |
+
copy_init(cb_eps);
|
| 158 |
+
copy_tile(cb_eps, 0, 1);
|
| 159 |
+
add_binary_tile_init();
|
| 160 |
+
add_binary_tile<RNE>(0, 1, 0);
|
| 161 |
+
rsqrt_tile_init();
|
| 162 |
+
rsqrt_tile<RsqrtMode::Default>(0);
|
| 163 |
+
pack_one(0, cb_stat);
|
| 164 |
+
cb_pop_front(cb_gat, Wt);
|
| 165 |
+
}
|
| 166 |
+
// y = xc * rstd (* gamma) (+ beta)
|
| 167 |
+
cb_wait_front(cb_bc, 1);
|
| 168 |
+
tile_regs_acquire();
|
| 169 |
+
copy_init(cb_xc);
|
| 170 |
+
copy_tile(cb_xc, 0, 0);
|
| 171 |
+
copy_tile(cb_bc, 0, 1);
|
| 172 |
+
mul_binary_tile_init();
|
| 173 |
+
mul_binary_tile(0, 1, 0);
|
| 174 |
+
if constexpr (has_gamma) {
|
| 175 |
+
cb_wait_front(cb_g, 1);
|
| 176 |
+
copy_tile(cb_g, 0, 1);
|
| 177 |
+
mul_binary_tile(0, 1, 0);
|
| 178 |
+
}
|
| 179 |
+
if constexpr (has_beta) {
|
| 180 |
+
cb_wait_front(cb_b, 1);
|
| 181 |
+
copy_tile(cb_b, 0, 1);
|
| 182 |
+
add_binary_tile_init();
|
| 183 |
+
add_binary_tile<RNE>(0, 1, 0);
|
| 184 |
+
}
|
| 185 |
+
if constexpr (kcat_out) {
|
| 186 |
+
pack_one(0, cb_y);
|
| 187 |
+
cb_pop_front(cb_bc, 1);
|
| 188 |
+
cb_pop_front(cb_xc, 1);
|
| 189 |
+
cb_wait_front(cb_y, 1);
|
| 190 |
+
tile_regs_acquire();
|
| 191 |
+
copy_init(cb_y);
|
| 192 |
+
copy_tile(cb_y, 0, 0);
|
| 193 |
+
copy_tile(cb_y, 0, 1);
|
| 194 |
+
typecast_tile_init<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>();
|
| 195 |
+
typecast_tile<(uint32_t)DataFormat::Float32, (uint32_t)DataFormat::Float16_b>(0);
|
| 196 |
+
sub_binary_tile_init();
|
| 197 |
+
sub_binary_tile<RNE>(1, 0, 1);
|
| 198 |
+
cb_reserve_back(cb_out, 1);
|
| 199 |
+
cb_reserve_back(cb_lo, 1);
|
| 200 |
+
tile_regs_commit();
|
| 201 |
+
tile_regs_wait();
|
| 202 |
+
pack_tile(0, cb_out);
|
| 203 |
+
pack_tile(1, cb_lo);
|
| 204 |
+
tile_regs_release();
|
| 205 |
+
cb_push_back(cb_out, 1);
|
| 206 |
+
cb_push_back(cb_lo, 1);
|
| 207 |
+
cb_pop_front(cb_y, 1);
|
| 208 |
+
} else {
|
| 209 |
+
pack_one(0, cb_out);
|
| 210 |
+
cb_pop_front(cb_bc, 1);
|
| 211 |
+
cb_pop_front(cb_xc, 1);
|
| 212 |
+
}
|
| 213 |
+
}
|
| 214 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp
ADDED
|
@@ -0,0 +1,130 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), reader (RISCV_0). One tile row is spread over Wt
|
| 3 |
+
// cores (member j owns column tile j); member 0 (the root) gathers the Wt tiles that the stock reduction folds in
|
| 4 |
+
// order and broadcasts the statistics back, so the arithmetic stays the stock decomposition's (bit-identical).
|
| 5 |
+
// Per row this kernel: reads x (r, j) (+ res (r, j)); on the root, waits for the Wt gathered tiles (semaphore 0,
|
| 6 |
+
// monotonic count) and hands them to the compute (c_10); then waits for the two statistic broadcasts (semaphore 1,
|
| 7 |
+
// monotonic count) and hands each to the compute (c_5, 2 pages: mean, rstd).
|
| 8 |
+
// CT args: [0] Wt, [1] has_gamma, [2] has_beta, [3] eps bits, [4] per-core RT-arg count (P3), [5] has_res,
|
| 9 |
+
// [6] has_rgate, then the TensorAccessorArgs of x, gamma, beta, res, rgate.
|
| 10 |
+
// Common RT args: [x_addr, gamma_addr, beta_addr, res_addr, rgate_addr].
|
| 11 |
+
// Per-core RT args: [row0, n_rows, row_stride, j, root_x, root_y, (member x, y) x Wt].
|
| 12 |
+
#include <cstdint>
|
| 13 |
+
|
| 14 |
+
#include "api/dataflow/dataflow_api.h"
|
| 15 |
+
|
| 16 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 17 |
+
constexpr uint32_t has_gamma = get_compile_time_arg_val(1);
|
| 18 |
+
constexpr uint32_t has_beta = get_compile_time_arg_val(2);
|
| 19 |
+
constexpr uint32_t eps_bits = get_compile_time_arg_val(3);
|
| 20 |
+
constexpr uint32_t has_res = get_compile_time_arg_val(5);
|
| 21 |
+
constexpr uint32_t has_rgate = get_compile_time_arg_val(6);
|
| 22 |
+
constexpr auto x_args = TensorAccessorArgs<7>();
|
| 23 |
+
constexpr auto g_args = TensorAccessorArgs<x_args.next_compile_time_args_offset()>();
|
| 24 |
+
constexpr auto b_args = TensorAccessorArgs<g_args.next_compile_time_args_offset()>();
|
| 25 |
+
constexpr auto r_args = TensorAccessorArgs<b_args.next_compile_time_args_offset()>();
|
| 26 |
+
constexpr auto rg_args = TensorAccessorArgs<r_args.next_compile_time_args_offset()>();
|
| 27 |
+
|
| 28 |
+
constexpr uint32_t cb_x = 0, cb_g = 1, cb_b = 2, cb_eps = 3, cb_bc = 5, cb_r = 7, cb_rg = 8, cb_gat = 10;
|
| 29 |
+
constexpr uint32_t TB = 4096;
|
| 30 |
+
|
| 31 |
+
FORCE_INLINE void fill_first_row(uint32_t p) {
|
| 32 |
+
constexpr uint32_t kFace = 1024, kRow = 64;
|
| 33 |
+
const uint32_t f0 = p, f1 = p + kFace;
|
| 34 |
+
for (uint32_t n = kRow; n < kFace; n <<= 1) {
|
| 35 |
+
noc_async_read(get_noc_addr(f0), f0 + n, n);
|
| 36 |
+
noc_async_read(get_noc_addr(f1), f1 + n, n);
|
| 37 |
+
noc_async_read_barrier();
|
| 38 |
+
}
|
| 39 |
+
noc_async_read(get_noc_addr(f0), f0 + 2 * kFace, 2 * kFace);
|
| 40 |
+
noc_async_read_barrier();
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
template <typename A>
|
| 44 |
+
FORCE_INLINE void read_row_tile(uint32_t cb, const A& acc, uint32_t j) {
|
| 45 |
+
cb_reserve_back(cb, 1);
|
| 46 |
+
const uint32_t p = get_write_ptr(cb);
|
| 47 |
+
noc_async_read(acc.get_noc_addr(j), p, TB);
|
| 48 |
+
noc_async_read_barrier();
|
| 49 |
+
fill_first_row(p);
|
| 50 |
+
cb_push_back(cb, 1);
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
void kernel_main() {
|
| 54 |
+
const uint32_t x_addr = get_common_arg_val<uint32_t>(0);
|
| 55 |
+
const uint32_t g_addr = get_common_arg_val<uint32_t>(1);
|
| 56 |
+
const uint32_t b_addr = get_common_arg_val<uint32_t>(2);
|
| 57 |
+
const uint32_t r_addr = get_common_arg_val<uint32_t>(3);
|
| 58 |
+
const uint32_t rg_addr = get_common_arg_val<uint32_t>(4);
|
| 59 |
+
const uint32_t row0 = get_arg_val<uint32_t>(0);
|
| 60 |
+
const uint32_t n_rows = get_arg_val<uint32_t>(1);
|
| 61 |
+
const uint32_t stride = get_arg_val<uint32_t>(2);
|
| 62 |
+
const uint32_t j = get_arg_val<uint32_t>(3);
|
| 63 |
+
if (n_rows == 0) {
|
| 64 |
+
return;
|
| 65 |
+
}
|
| 66 |
+
const bool root = j == 0;
|
| 67 |
+
const auto x = TensorAccessor(x_args, x_addr, TB);
|
| 68 |
+
const auto res = TensorAccessor(r_args, r_addr, TB);
|
| 69 |
+
volatile tt_l1_ptr uint32_t* sem_g = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(0));
|
| 70 |
+
volatile tt_l1_ptr uint32_t* sem_b = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_semaphore(1));
|
| 71 |
+
uint32_t n_gat = 0, n_bc = 0;
|
| 72 |
+
|
| 73 |
+
auto read_x = [&](uint32_t r) {
|
| 74 |
+
cb_reserve_back(cb_x, 1);
|
| 75 |
+
noc_async_read(x.get_noc_addr(r * Wt + j), get_write_ptr(cb_x), TB);
|
| 76 |
+
if constexpr (has_res) {
|
| 77 |
+
cb_reserve_back(cb_r, 1);
|
| 78 |
+
noc_async_read(res.get_noc_addr(r * Wt + j), get_write_ptr(cb_r), TB);
|
| 79 |
+
noc_async_read_barrier();
|
| 80 |
+
cb_push_back(cb_r, 1);
|
| 81 |
+
} else {
|
| 82 |
+
noc_async_read_barrier();
|
| 83 |
+
}
|
| 84 |
+
cb_push_back(cb_x, 1);
|
| 85 |
+
};
|
| 86 |
+
auto gather = [&]() {
|
| 87 |
+
cb_reserve_back(cb_gat, Wt);
|
| 88 |
+
n_gat += Wt;
|
| 89 |
+
noc_semaphore_wait_min(sem_g, n_gat);
|
| 90 |
+
cb_push_back(cb_gat, Wt);
|
| 91 |
+
};
|
| 92 |
+
auto bcast_in = [&]() {
|
| 93 |
+
cb_reserve_back(cb_bc, 1);
|
| 94 |
+
n_bc += 1;
|
| 95 |
+
noc_semaphore_wait_min(sem_b, n_bc);
|
| 96 |
+
cb_push_back(cb_bc, 1);
|
| 97 |
+
};
|
| 98 |
+
|
| 99 |
+
read_x(row0);
|
| 100 |
+
if constexpr (has_rgate) {
|
| 101 |
+
read_row_tile(cb_rg, TensorAccessor(rg_args, rg_addr, TB), j);
|
| 102 |
+
}
|
| 103 |
+
cb_reserve_back(cb_eps, 1);
|
| 104 |
+
{
|
| 105 |
+
auto* e = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_eps));
|
| 106 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 107 |
+
e[i] = eps_bits;
|
| 108 |
+
}
|
| 109 |
+
}
|
| 110 |
+
cb_push_back(cb_eps, 1);
|
| 111 |
+
if constexpr (has_gamma) {
|
| 112 |
+
read_row_tile(cb_g, TensorAccessor(g_args, g_addr, TB), j);
|
| 113 |
+
}
|
| 114 |
+
if constexpr (has_beta) {
|
| 115 |
+
read_row_tile(cb_b, TensorAccessor(b_args, b_addr, TB), j);
|
| 116 |
+
}
|
| 117 |
+
for (uint32_t i = 0; i < n_rows; ++i) {
|
| 118 |
+
if (i > 0) {
|
| 119 |
+
read_x(row0 + i * stride);
|
| 120 |
+
}
|
| 121 |
+
if (root) {
|
| 122 |
+
gather();
|
| 123 |
+
}
|
| 124 |
+
bcast_in();
|
| 125 |
+
if (root) {
|
| 126 |
+
gather();
|
| 127 |
+
}
|
| 128 |
+
bcast_in();
|
| 129 |
+
}
|
| 130 |
+
}
|
code/tt_diffusion_planner/tt/kernels/ln32s_writer.cpp
ADDED
|
@@ -0,0 +1,161 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Split-row fused fp32 LayerNorm (tt/ln_kernel.py, LN_SPLIT), writer (RISCV_1). Per row: (write_h) the new stream
|
| 3 |
+
// tile; sends this member's tile for the mean fold (c_11, from the compute) to page j of the root's gather CB c_10
|
| 4 |
+
// and raises the root's semaphore 0; on the root, broadcasts the packed mean (c_4, column 0) over the 32 columns and
|
| 5 |
+
// sends it to page 0 of every member's c_5 (semaphore 1 raised); the same for the squared tile and rstd (page 1);
|
| 6 |
+
// then writes the output tile.
|
| 7 |
+
// With kcat_out (KCAT_EMIT) out is the split operand [y_hi | y_hi | y_lo | 1 | 0..] of the next K-concatenated
|
| 8 |
+
// linear (row tiles Ktp, the kcat_writer layout): y_hi to (r, j) and (r, Wt + j), y_lo to (r, 2 Wt + j); the root
|
| 9 |
+
// also writes the ones tile (r, 3 Wt) and the zero pad tiles.
|
| 10 |
+
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] write_h, [3] sfpu_bcast, [4] kcat_out, [5] Ktp, then the
|
| 11 |
+
// TensorAccessorArgs of out and h.
|
| 12 |
+
// Common RT args: [out_addr, h_addr]. Per-core RT args: as the reader's.
|
| 13 |
+
#include <cstdint>
|
| 14 |
+
|
| 15 |
+
#include "api/dataflow/dataflow_api.h"
|
| 16 |
+
|
| 17 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 18 |
+
constexpr uint32_t write_h = get_compile_time_arg_val(2);
|
| 19 |
+
constexpr uint32_t sfpu_bcast = get_compile_time_arg_val(3); // the root's compute packs a broadcast tile
|
| 20 |
+
constexpr uint32_t kcat_out = get_compile_time_arg_val(4);
|
| 21 |
+
constexpr uint32_t Ktp = get_compile_time_arg_val(5);
|
| 22 |
+
constexpr auto out_args = TensorAccessorArgs<6>();
|
| 23 |
+
constexpr auto h_args = TensorAccessorArgs<out_args.next_compile_time_args_offset()>();
|
| 24 |
+
|
| 25 |
+
constexpr uint32_t cb_stat = 4, cb_bc = 5, cb_gat = 10, cb_snd = 11, cb_scr = 12, cb_out = 16, cb_hout = 17,
|
| 26 |
+
cb_lo = 19, cb_one = 20, cb_zero = 21;
|
| 27 |
+
constexpr uint32_t TB = 4096;
|
| 28 |
+
|
| 29 |
+
FORCE_INLINE void bcast_col(uint32_t src, uint32_t dst) {
|
| 30 |
+
auto* s = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(src);
|
| 31 |
+
auto* d = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(dst);
|
| 32 |
+
for (uint32_t fo = 0; fo < 1024; fo += 512) {
|
| 33 |
+
for (uint32_t ro = 0; ro < 256; ro += 16) {
|
| 34 |
+
const uint32_t v = s[fo + ro];
|
| 35 |
+
volatile tt_l1_ptr uint32_t* l = d + fo + ro;
|
| 36 |
+
volatile tt_l1_ptr uint32_t* r = l + 256;
|
| 37 |
+
for (uint32_t c = 0; c < 16; ++c) {
|
| 38 |
+
l[c] = v;
|
| 39 |
+
r[c] = v;
|
| 40 |
+
}
|
| 41 |
+
}
|
| 42 |
+
}
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
void kernel_main() {
|
| 46 |
+
const uint32_t out_addr = get_common_arg_val<uint32_t>(0);
|
| 47 |
+
const uint32_t h_addr = get_common_arg_val<uint32_t>(1);
|
| 48 |
+
const uint32_t row0 = get_arg_val<uint32_t>(0);
|
| 49 |
+
const uint32_t n_rows = get_arg_val<uint32_t>(1);
|
| 50 |
+
const uint32_t stride = get_arg_val<uint32_t>(2);
|
| 51 |
+
const uint32_t j = get_arg_val<uint32_t>(3);
|
| 52 |
+
const uint32_t root_x = get_arg_val<uint32_t>(4);
|
| 53 |
+
const uint32_t root_y = get_arg_val<uint32_t>(5);
|
| 54 |
+
if (n_rows == 0) {
|
| 55 |
+
return;
|
| 56 |
+
}
|
| 57 |
+
const bool root = j == 0;
|
| 58 |
+
const auto out = TensorAccessor(out_args, out_addr, TB);
|
| 59 |
+
const auto hacc = TensorAccessor(h_args, h_addr, TB);
|
| 60 |
+
const uint32_t gat_base = get_write_ptr(cb_gat); // the same L1 address on every core of the program
|
| 61 |
+
const uint32_t bc_base = get_write_ptr(cb_bc);
|
| 62 |
+
const uint64_t root_sem = get_noc_addr(root_x, root_y, get_semaphore(0));
|
| 63 |
+
const uint32_t sem_b_addr = get_semaphore(1);
|
| 64 |
+
cb_reserve_back(cb_scr, 1);
|
| 65 |
+
const uint32_t scr = get_write_ptr(cb_scr);
|
| 66 |
+
|
| 67 |
+
auto send_to_root = [&]() {
|
| 68 |
+
cb_wait_front(cb_snd, 1);
|
| 69 |
+
noc_async_write(get_read_ptr(cb_snd), get_noc_addr(root_x, root_y, gat_base + j * TB), TB);
|
| 70 |
+
noc_async_write_barrier();
|
| 71 |
+
noc_semaphore_inc(root_sem, 1);
|
| 72 |
+
cb_pop_front(cb_snd, 1);
|
| 73 |
+
};
|
| 74 |
+
auto bcast_out = [&](uint32_t page) {
|
| 75 |
+
cb_wait_front(cb_stat, 1);
|
| 76 |
+
uint32_t src = get_read_ptr(cb_stat);
|
| 77 |
+
if constexpr (!sfpu_bcast) {
|
| 78 |
+
bcast_col(src, scr);
|
| 79 |
+
src = scr;
|
| 80 |
+
}
|
| 81 |
+
for (uint32_t m = 0; m < Wt; ++m) {
|
| 82 |
+
const uint32_t mx = get_arg_val<uint32_t>(6 + 2 * m), my = get_arg_val<uint32_t>(7 + 2 * m);
|
| 83 |
+
noc_async_write(src, get_noc_addr(mx, my, bc_base + page * TB), TB);
|
| 84 |
+
}
|
| 85 |
+
noc_async_write_barrier();
|
| 86 |
+
cb_pop_front(cb_stat, 1);
|
| 87 |
+
for (uint32_t m = 0; m < Wt; ++m) {
|
| 88 |
+
const uint32_t mx = get_arg_val<uint32_t>(6 + 2 * m), my = get_arg_val<uint32_t>(7 + 2 * m);
|
| 89 |
+
noc_semaphore_inc(get_noc_addr(mx, my, sem_b_addr), 1);
|
| 90 |
+
}
|
| 91 |
+
};
|
| 92 |
+
|
| 93 |
+
uint32_t one_ptr = 0, zero_ptr = 0;
|
| 94 |
+
if constexpr (kcat_out) {
|
| 95 |
+
if (root) {
|
| 96 |
+
// ones tile: element (i, c) at face (i / 16) * 2 + (c / 16), offset (i % 16) * 16 + c % 16
|
| 97 |
+
cb_reserve_back(cb_one, 1);
|
| 98 |
+
one_ptr = get_write_ptr(cb_one);
|
| 99 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(one_ptr);
|
| 100 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 101 |
+
p[i] = 0;
|
| 102 |
+
}
|
| 103 |
+
for (uint32_t i = 0; i < 32; ++i) {
|
| 104 |
+
const uint32_t base = (i / 16) * 512 + (i % 16) * 16;
|
| 105 |
+
p[base] = 0x3F800000u;
|
| 106 |
+
p[base + 1] = 0x3F800000u;
|
| 107 |
+
}
|
| 108 |
+
if constexpr (Ktp > 3 * Wt + 1) {
|
| 109 |
+
cb_reserve_back(cb_zero, 1);
|
| 110 |
+
zero_ptr = get_write_ptr(cb_zero);
|
| 111 |
+
auto* z = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(zero_ptr);
|
| 112 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 113 |
+
z[i] = 0;
|
| 114 |
+
}
|
| 115 |
+
}
|
| 116 |
+
}
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
for (uint32_t i = 0; i < n_rows; ++i) {
|
| 120 |
+
const uint32_t r = row0 + i * stride;
|
| 121 |
+
if constexpr (write_h) {
|
| 122 |
+
cb_wait_front(cb_hout, 1);
|
| 123 |
+
noc_async_write(get_read_ptr(cb_hout), hacc.get_noc_addr(r * Wt + j), TB);
|
| 124 |
+
noc_async_writes_flushed();
|
| 125 |
+
cb_pop_front(cb_hout, 1);
|
| 126 |
+
}
|
| 127 |
+
send_to_root();
|
| 128 |
+
if (root) {
|
| 129 |
+
bcast_out(0);
|
| 130 |
+
}
|
| 131 |
+
send_to_root();
|
| 132 |
+
if (root) {
|
| 133 |
+
bcast_out(1);
|
| 134 |
+
}
|
| 135 |
+
if constexpr (kcat_out) {
|
| 136 |
+
const uint32_t ob = r * Ktp;
|
| 137 |
+
cb_wait_front(cb_out, 1);
|
| 138 |
+
const uint32_t hp = get_read_ptr(cb_out);
|
| 139 |
+
noc_async_write(hp, out.get_noc_addr(ob + j), TB);
|
| 140 |
+
noc_async_write(hp, out.get_noc_addr(ob + Wt + j), TB);
|
| 141 |
+
cb_wait_front(cb_lo, 1);
|
| 142 |
+
noc_async_write(get_read_ptr(cb_lo), out.get_noc_addr(ob + 2 * Wt + j), TB);
|
| 143 |
+
if (root) {
|
| 144 |
+
noc_async_write(one_ptr, out.get_noc_addr(ob + 3 * Wt), TB);
|
| 145 |
+
for (uint32_t z = 3 * Wt + 1; z < Ktp; ++z) {
|
| 146 |
+
noc_async_write(zero_ptr, out.get_noc_addr(ob + z), TB);
|
| 147 |
+
}
|
| 148 |
+
}
|
| 149 |
+
noc_async_writes_flushed();
|
| 150 |
+
cb_pop_front(cb_out, 1);
|
| 151 |
+
cb_pop_front(cb_lo, 1);
|
| 152 |
+
} else {
|
| 153 |
+
cb_wait_front(cb_out, 1);
|
| 154 |
+
noc_async_write(get_read_ptr(cb_out), out.get_noc_addr(r * Wt + j), TB);
|
| 155 |
+
noc_async_writes_flushed();
|
| 156 |
+
cb_pop_front(cb_out, 1);
|
| 157 |
+
}
|
| 158 |
+
}
|
| 159 |
+
noc_async_write_barrier();
|
| 160 |
+
noc_async_atomic_barrier();
|
| 161 |
+
}
|
code/tt_diffusion_planner/tt/kernels/smask_compute.cpp
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Attention score scale + mask (tt/smask_kernel.py), compute: out = s * scale + mask, the two stock binary_ng fp32
|
| 3 |
+
// programs `ttnn.multiply(s, scale)` (SFPU mul_binary_tile against a tile filled with the fp32 scale) and
|
| 4 |
+
// `ttnn.add(., mask)`, in one pass. With an fp32 mask the stock add is the SFPU add_binary_tile<NearestEven>; with a
|
| 5 |
+
// bf16 mask (trunc_tf32 = 1) binary_ng takes its FPU path, whose SrcA truncates the fp32 scores to TF32 (measured:
|
| 6 |
+
// logs/diffusion-planner/opt_r2/smask_diag.log, the stock result == (exact & 0xFFFFE000) for every finite element),
|
| 7 |
+
// so the scaled scores are truncated the same way (SFPU bitwise AND on the raw bits) before the exact add of the
|
| 8 |
+
// 0 / -inf mask. Bit-identical to the two programs in both cases.
|
| 9 |
+
// CT args: [0] H (even), [1] per-core RT-arg count (P3), [2] trunc_tf32. Per-core RT args: [n, 0].
|
| 10 |
+
#include <cstdint>
|
| 11 |
+
|
| 12 |
+
#include "api/compute/common.h"
|
| 13 |
+
#include "api/compute/compute_kernel_api.h"
|
| 14 |
+
#include "api/compute/eltwise_binary_sfpu.h"
|
| 15 |
+
#include "api/compute/eltwise_unary/bitwise.h"
|
| 16 |
+
#include "api/compute/eltwise_unary/eltwise_unary.h"
|
| 17 |
+
#include "api/compute/tile_move_copy.h"
|
| 18 |
+
|
| 19 |
+
using namespace ckernel;
|
| 20 |
+
|
| 21 |
+
constexpr uint32_t H = get_compile_time_arg_val(0);
|
| 22 |
+
constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2);
|
| 23 |
+
constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_out = 16;
|
| 24 |
+
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
|
| 25 |
+
|
| 26 |
+
void kernel_main() {
|
| 27 |
+
const uint32_t n = get_arg_val<uint32_t>(0);
|
| 28 |
+
if (n == 0) {
|
| 29 |
+
return;
|
| 30 |
+
}
|
| 31 |
+
compute_kernel_hw_startup(cb_s, cb_out);
|
| 32 |
+
cb_wait_front(cb_c, 1);
|
| 33 |
+
for (uint32_t t = 0; t < n; ++t) {
|
| 34 |
+
cb_wait_front(cb_m, 1);
|
| 35 |
+
cb_wait_front(cb_s, H);
|
| 36 |
+
for (uint32_t h = 0; h < H; h += 2) {
|
| 37 |
+
tile_regs_acquire();
|
| 38 |
+
copy_init(cb_s);
|
| 39 |
+
copy_tile(cb_s, h, 0);
|
| 40 |
+
copy_tile(cb_s, h + 1, 1);
|
| 41 |
+
copy_init(cb_c);
|
| 42 |
+
copy_tile(cb_c, 0, 2);
|
| 43 |
+
mul_binary_tile_init();
|
| 44 |
+
mul_binary_tile(0, 2, 0);
|
| 45 |
+
mul_binary_tile(1, 2, 1);
|
| 46 |
+
if constexpr (trunc_tf32) {
|
| 47 |
+
bitwise_and_tile_init();
|
| 48 |
+
bitwise_and_tile<DataFormat::Int32>(0, 0xFFFFE000u);
|
| 49 |
+
bitwise_and_tile<DataFormat::Int32>(1, 0xFFFFE000u);
|
| 50 |
+
}
|
| 51 |
+
reconfig_data_format_srca(cb_c, cb_m);
|
| 52 |
+
copy_init(cb_m);
|
| 53 |
+
copy_tile(cb_m, 0, 3);
|
| 54 |
+
reconfig_data_format_srca(cb_m, cb_s);
|
| 55 |
+
add_binary_tile_init();
|
| 56 |
+
add_binary_tile<RNE>(0, 3, 0);
|
| 57 |
+
add_binary_tile<RNE>(1, 3, 1);
|
| 58 |
+
cb_reserve_back(cb_out, 2);
|
| 59 |
+
tile_regs_commit();
|
| 60 |
+
tile_regs_wait();
|
| 61 |
+
pack_tile(0, cb_out);
|
| 62 |
+
pack_tile(1, cb_out);
|
| 63 |
+
tile_regs_release();
|
| 64 |
+
cb_push_back(cb_out, 2);
|
| 65 |
+
}
|
| 66 |
+
cb_pop_front(cb_s, H);
|
| 67 |
+
cb_pop_front(cb_m, 1);
|
| 68 |
+
}
|
| 69 |
+
}
|
code/tt_diffusion_planner/tt/kernels/smask_reader.cpp
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Attention score scale + mask (tt/smask_kernel.py), reader (RISCV_0): the scale tile once, then per mask tile m of
|
| 3 |
+
// this core's range the bf16 mask tile and the H fp32 score tiles (h, m) (the mask broadcasts over the heads).
|
| 4 |
+
// CT args: [0] H, [1] mask tiles per head plane (Mt * Kt), [2] scale bits (fp32), [3] per-core RT-arg count (P3),
|
| 5 |
+
// [4] mask tile bytes (bf16 2048 / fp32 4096), then the TensorAccessorArgs of scores and mask.
|
| 6 |
+
// Common RT args: [s_addr, m_addr]. Per-core RT args: [m0, n].
|
| 7 |
+
#include <cstdint>
|
| 8 |
+
|
| 9 |
+
#include "api/dataflow/dataflow_api.h"
|
| 10 |
+
|
| 11 |
+
constexpr uint32_t H = get_compile_time_arg_val(0);
|
| 12 |
+
constexpr uint32_t plane = get_compile_time_arg_val(1);
|
| 13 |
+
constexpr uint32_t scale_bits = get_compile_time_arg_val(2);
|
| 14 |
+
constexpr uint32_t TBM = get_compile_time_arg_val(4);
|
| 15 |
+
constexpr auto s_args = TensorAccessorArgs<5>();
|
| 16 |
+
constexpr auto m_args = TensorAccessorArgs<s_args.next_compile_time_args_offset()>();
|
| 17 |
+
constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2;
|
| 18 |
+
constexpr uint32_t TB = 4096;
|
| 19 |
+
|
| 20 |
+
void kernel_main() {
|
| 21 |
+
const uint32_t s_addr = get_common_arg_val<uint32_t>(0);
|
| 22 |
+
const uint32_t m_addr = get_common_arg_val<uint32_t>(1);
|
| 23 |
+
const uint32_t m0 = get_arg_val<uint32_t>(0);
|
| 24 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 25 |
+
if (n == 0) {
|
| 26 |
+
return;
|
| 27 |
+
}
|
| 28 |
+
const auto s = TensorAccessor(s_args, s_addr, TB);
|
| 29 |
+
const auto mk = TensorAccessor(m_args, m_addr, TBM);
|
| 30 |
+
cb_reserve_back(cb_c, 1);
|
| 31 |
+
{
|
| 32 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_c));
|
| 33 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 34 |
+
p[i] = scale_bits;
|
| 35 |
+
}
|
| 36 |
+
}
|
| 37 |
+
cb_push_back(cb_c, 1);
|
| 38 |
+
for (uint32_t m = m0; m < m0 + n; ++m) {
|
| 39 |
+
cb_reserve_back(cb_m, 1);
|
| 40 |
+
noc_async_read(mk.get_noc_addr(m), get_write_ptr(cb_m), TBM);
|
| 41 |
+
cb_reserve_back(cb_s, H);
|
| 42 |
+
const uint32_t p = get_write_ptr(cb_s);
|
| 43 |
+
for (uint32_t h = 0; h < H; ++h) {
|
| 44 |
+
noc_async_read(s.get_noc_addr(h * plane + m), p + h * TB, TB);
|
| 45 |
+
}
|
| 46 |
+
noc_async_read_barrier();
|
| 47 |
+
cb_push_back(cb_m, 1);
|
| 48 |
+
cb_push_back(cb_s, H);
|
| 49 |
+
}
|
| 50 |
+
}
|
code/tt_diffusion_planner/tt/kernels/smask_writer.cpp
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Attention score scale + mask (tt/smask_kernel.py), writer (RISCV_1): the H output tiles (h, m) per mask tile m.
|
| 3 |
+
// CT args: [0] H, [1] plane (Mt * Kt), [2] per-core RT-arg count (P3), then the TensorAccessorArgs of out.
|
| 4 |
+
// Common RT args: [out_addr]. Per-core RT args: [m0, n].
|
| 5 |
+
#include <cstdint>
|
| 6 |
+
|
| 7 |
+
#include "api/dataflow/dataflow_api.h"
|
| 8 |
+
|
| 9 |
+
constexpr uint32_t H = get_compile_time_arg_val(0);
|
| 10 |
+
constexpr uint32_t plane = get_compile_time_arg_val(1);
|
| 11 |
+
constexpr auto o_args = TensorAccessorArgs<3>();
|
| 12 |
+
constexpr uint32_t cb_out = 16;
|
| 13 |
+
constexpr uint32_t TB = 4096;
|
| 14 |
+
|
| 15 |
+
void kernel_main() {
|
| 16 |
+
const uint32_t o_addr = get_common_arg_val<uint32_t>(0);
|
| 17 |
+
const uint32_t m0 = get_arg_val<uint32_t>(0);
|
| 18 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 19 |
+
const auto o = TensorAccessor(o_args, o_addr, TB);
|
| 20 |
+
for (uint32_t m = m0; m < m0 + n; ++m) {
|
| 21 |
+
for (uint32_t h = 0; h < H; h += 2) {
|
| 22 |
+
cb_wait_front(cb_out, 2);
|
| 23 |
+
const uint32_t p = get_read_ptr(cb_out);
|
| 24 |
+
noc_async_write(p, o.get_noc_addr(h * plane + m), TB);
|
| 25 |
+
noc_async_write(p + TB, o.get_noc_addr((h + 1) * plane + m), TB);
|
| 26 |
+
noc_async_writes_flushed();
|
| 27 |
+
cb_pop_front(cb_out, 2);
|
| 28 |
+
}
|
| 29 |
+
}
|
| 30 |
+
noc_async_write_barrier();
|
| 31 |
+
}
|
code/tt_diffusion_planner/tt/kernels/smsm_compute.cpp
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Attention score scale + mask + softmax (tt/smsm_kernel.py, ATTN_SMSM), compute. Per tile row (one head, 32 query
|
| 3 |
+
// rows, Wt key tiles):
|
| 4 |
+
// phase A: x = trunc_tf32(s * scale) + mask, the smask kernel's SFPU sequence (kernels/smask_compute.cpp: the two
|
| 5 |
+
// stock binary_ng programs bit for bit), packed to cb_x (L1) instead of DRAM;
|
| 6 |
+
// phase B: the stock ttnn.softmax(numeric_stable=True) of that row, i.e. the no-mask path of
|
| 7 |
+
// ttnn/.../softmax/device/kernels/attention/compute/softmax.cpp with the same kernel_lib calls: row max
|
| 8 |
+
// (FPU reduce), exp(x - max) (FPU bcast sub + SFPU exp), row sum (FPU reduce) + precise fp32
|
| 9 |
+
// reciprocal, x * 1/sum (FPU bcast mul).
|
| 10 |
+
// The stock softmax unpacks its fp32 input to SrcA (TF32); x holds TF32 values already (phase A truncates them and
|
| 11 |
+
// the mask is 0 / -inf), so the stock program saw exactly these operands.
|
| 12 |
+
// CT args: [0] Wt, [1] per-core RT-arg count (P3), [2] trunc_tf32, [3] ndst (block), [4] out pad tiles per row,
|
| 13 |
+
// [5] scale mode: 0 = multiply by a tile filled with the scale (mul_binary_tile, as the stock binary_ng
|
| 14 |
+
// program), 1 = mul_unary_tile with the scale bits (the same SFPU fp32 multiply, no scale tile copy).
|
| 15 |
+
// Per-core RT args: [nrows, 0].
|
| 16 |
+
#include <cstdint>
|
| 17 |
+
|
| 18 |
+
#include "api/compute/common.h"
|
| 19 |
+
#include "api/compute/compute_kernel_api.h"
|
| 20 |
+
#include "api/compute/eltwise_binary.h"
|
| 21 |
+
#include "api/compute/eltwise_binary_sfpu.h"
|
| 22 |
+
#include "api/compute/eltwise_unary/binop_with_scalar.h"
|
| 23 |
+
#include "api/compute/eltwise_unary/bitwise.h"
|
| 24 |
+
#include "api/compute/eltwise_unary/eltwise_unary.h"
|
| 25 |
+
#include "api/compute/tile_move_copy.h"
|
| 26 |
+
#include "api/compute/bcast.h"
|
| 27 |
+
#include "api/compute/softmax.h"
|
| 28 |
+
#include "api/compute/reduce.h"
|
| 29 |
+
#include "ttnn/cpp/ttnn/kernel_lib/reduce_helpers_compute.hpp"
|
| 30 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/chain.hpp"
|
| 31 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/api/convenience.hpp"
|
| 32 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/unary/math.hpp"
|
| 33 |
+
#include "ttnn/cpp/ttnn/kernel_lib/eltwise/core/optional.hpp"
|
| 34 |
+
|
| 35 |
+
namespace ckl = compute_kernel_lib;
|
| 36 |
+
using namespace ckernel;
|
| 37 |
+
|
| 38 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 39 |
+
constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2);
|
| 40 |
+
constexpr uint32_t ndst = get_compile_time_arg_val(3);
|
| 41 |
+
constexpr uint32_t out_pad = get_compile_time_arg_val(4);
|
| 42 |
+
constexpr uint32_t scale_mode = get_compile_time_arg_val(5);
|
| 43 |
+
constexpr uint32_t scale_bits = get_compile_time_arg_val(6);
|
| 44 |
+
constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_max_scaler = 3, cb_sum_scaler = 4;
|
| 45 |
+
constexpr uint32_t cb_out = 16;
|
| 46 |
+
constexpr uint32_t cb_x = 24, cb_max = 25, cb_exps = 26, cb_recip = 27;
|
| 47 |
+
constexpr auto RNE = ckernel::DstRoundingMode::NearestEven;
|
| 48 |
+
constexpr uint32_t Wp = (Wt + 1) / 2; // tile pairs per row (the last one half-used when Wt is odd)
|
| 49 |
+
|
| 50 |
+
// the stock calc_numeric_stable (softmax.cpp), CB ids instead of DFB handles
|
| 51 |
+
template <std::uint32_t dfb_in, std::uint32_t dfb_max_scaler, std::uint32_t dfb_max, std::uint32_t dfb_out>
|
| 52 |
+
void calc_numeric_stable(std::uint32_t W, std::uint32_t nd) {
|
| 53 |
+
compute_kernel_lib::reduce<
|
| 54 |
+
PoolType::MAX,
|
| 55 |
+
ReduceDim::REDUCE_ROW,
|
| 56 |
+
dfb_in,
|
| 57 |
+
dfb_max_scaler,
|
| 58 |
+
dfb_max,
|
| 59 |
+
compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop,
|
| 60 |
+
compute_kernel_lib::ReduceDataFormatReconfigMode::INPUT>(compute_kernel_lib::ReduceInputBlockShape::row(W));
|
| 61 |
+
ckl::eltwise_chain(
|
| 62 |
+
ckl::IterationShape::tiles(W).block_size(nd),
|
| 63 |
+
ckl::BinaryFpu<
|
| 64 |
+
ckl::BinaryFpuOp::Sub,
|
| 65 |
+
ckl::input(
|
| 66 |
+
dfb_in,
|
| 67 |
+
ckl::WaitPolicy::Upfront,
|
| 68 |
+
ckl::PopPolicy::AtEnd,
|
| 69 |
+
ckl::InputTileMapping::Block,
|
| 70 |
+
ckl::DataFormatReconfig::Disabled),
|
| 71 |
+
ckl::input(dfb_max, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd)>{},
|
| 72 |
+
ckl::Exp<ckl::Approx::Exact, ckl::Dst::D0>{},
|
| 73 |
+
ckl::PackTile<ckl::output(
|
| 74 |
+
dfb_out,
|
| 75 |
+
ckl::ReservePolicy::PerBlockSize,
|
| 76 |
+
ckl::PushPolicy::PerBlockSize,
|
| 77 |
+
ckl::DataFormatReconfig::Disabled)>{});
|
| 78 |
+
cb_wait_front(dfb_out, W);
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
void kernel_main() {
|
| 82 |
+
const uint32_t nrows = get_arg_val<uint32_t>(0);
|
| 83 |
+
if (nrows == 0) {
|
| 84 |
+
return;
|
| 85 |
+
}
|
| 86 |
+
compute_kernel_hw_startup(cb_s, cb_c, cb_x);
|
| 87 |
+
cb_wait_front(cb_c, 1);
|
| 88 |
+
cb_wait_front(cb_max_scaler, 1);
|
| 89 |
+
cb_wait_front(cb_sum_scaler, 1);
|
| 90 |
+
for (uint32_t r = 0; r < nrows; ++r) {
|
| 91 |
+
// ---- phase A: x = trunc(s * scale) + mask (pairs of tiles) ----
|
| 92 |
+
reconfig_data_format(cb_s, cb_c);
|
| 93 |
+
pack_reconfig_data_format(cb_x);
|
| 94 |
+
cb_wait_front(cb_m, 2 * Wp);
|
| 95 |
+
for (uint32_t p = 0; p < Wp; ++p) {
|
| 96 |
+
const uint32_t j = 2 * p;
|
| 97 |
+
const bool two = (j + 1) < Wt;
|
| 98 |
+
cb_wait_front(cb_s, 2);
|
| 99 |
+
tile_regs_acquire();
|
| 100 |
+
copy_init(cb_s);
|
| 101 |
+
copy_tile(cb_s, 0, 0);
|
| 102 |
+
copy_tile(cb_s, 1, 1);
|
| 103 |
+
if constexpr (scale_mode == 0) {
|
| 104 |
+
copy_init(cb_c);
|
| 105 |
+
copy_tile(cb_c, 0, 2);
|
| 106 |
+
mul_binary_tile_init();
|
| 107 |
+
mul_binary_tile(0, 2, 0);
|
| 108 |
+
mul_binary_tile(1, 2, 1);
|
| 109 |
+
} else {
|
| 110 |
+
binop_with_scalar_tile_init();
|
| 111 |
+
mul_unary_tile(0, scale_bits);
|
| 112 |
+
mul_unary_tile(1, scale_bits);
|
| 113 |
+
}
|
| 114 |
+
if constexpr (trunc_tf32) {
|
| 115 |
+
bitwise_and_tile_init();
|
| 116 |
+
bitwise_and_tile<DataFormat::Int32>(0, 0xFFFFE000u);
|
| 117 |
+
bitwise_and_tile<DataFormat::Int32>(1, 0xFFFFE000u);
|
| 118 |
+
}
|
| 119 |
+
reconfig_data_format_srca(scale_mode == 0 ? cb_c : cb_s, cb_m);
|
| 120 |
+
copy_init(cb_m);
|
| 121 |
+
copy_tile(cb_m, j, 2);
|
| 122 |
+
copy_tile(cb_m, j + 1, 3);
|
| 123 |
+
reconfig_data_format_srca(cb_m, cb_s);
|
| 124 |
+
add_binary_tile_init();
|
| 125 |
+
add_binary_tile<RNE>(0, 2, 0);
|
| 126 |
+
add_binary_tile<RNE>(1, 3, 1);
|
| 127 |
+
const uint32_t np = two ? 2 : 1;
|
| 128 |
+
cb_reserve_back(cb_x, np);
|
| 129 |
+
tile_regs_commit();
|
| 130 |
+
tile_regs_wait();
|
| 131 |
+
pack_tile(0, cb_x);
|
| 132 |
+
if (two) {
|
| 133 |
+
pack_tile(1, cb_x);
|
| 134 |
+
}
|
| 135 |
+
tile_regs_release();
|
| 136 |
+
cb_push_back(cb_x, np);
|
| 137 |
+
cb_pop_front(cb_s, 2);
|
| 138 |
+
}
|
| 139 |
+
cb_pop_front(cb_m, 2 * Wp);
|
| 140 |
+
|
| 141 |
+
// ---- phase B: the stock numeric-stable softmax of the row ----
|
| 142 |
+
reconfig_data_format(cb_x, cb_x);
|
| 143 |
+
pack_reconfig_data_format(cb_exps);
|
| 144 |
+
copy_init(cb_x);
|
| 145 |
+
calc_numeric_stable<cb_x, cb_max_scaler, cb_max, cb_exps>(Wt, ndst);
|
| 146 |
+
reconfig_data_format(cb_exps, cb_sum_scaler);
|
| 147 |
+
compute_kernel_lib::reduce<
|
| 148 |
+
PoolType::SUM,
|
| 149 |
+
ReduceDim::REDUCE_ROW,
|
| 150 |
+
cb_exps,
|
| 151 |
+
cb_sum_scaler,
|
| 152 |
+
cb_recip,
|
| 153 |
+
compute_kernel_lib::ReduceInputPolicy::WaitUpfrontNoPop>(
|
| 154 |
+
compute_kernel_lib::ReduceInputBlockShape::row(Wt),
|
| 155 |
+
compute_kernel_lib::ReduceInputMemoryLayout::contiguous(),
|
| 156 |
+
compute_kernel_lib::NoAccumulation{},
|
| 157 |
+
[](std::uint32_t) {
|
| 158 |
+
if constexpr (DST_ACCUM_MODE) {
|
| 159 |
+
recip_tile_init<ReciprocalDestAcc::FP32, ReciprocalApproxMode::Precise>();
|
| 160 |
+
recip_tile<ReciprocalDestAcc::FP32, ReciprocalApproxMode::Precise>(0);
|
| 161 |
+
} else {
|
| 162 |
+
recip_tile_init();
|
| 163 |
+
recip_tile(0);
|
| 164 |
+
}
|
| 165 |
+
});
|
| 166 |
+
ckl::mul<
|
| 167 |
+
ckl::input(cb_exps, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd, ckl::InputTileMapping::Block),
|
| 168 |
+
ckl::input(cb_recip, ckl::BroadcastDim::Col, ckl::WaitPolicy::Upfront, ckl::PopPolicy::AtEnd),
|
| 169 |
+
ckl::output(cb_out, ckl::ReservePolicy::PerBlockSize, ckl::PushPolicy::PerBlockSize)>(
|
| 170 |
+
ckl::IterationShape::tiles(Wt).block_size(ndst));
|
| 171 |
+
if constexpr (out_pad > 0) {
|
| 172 |
+
cb_reserve_back(cb_out, out_pad);
|
| 173 |
+
cb_push_back(cb_out, out_pad);
|
| 174 |
+
}
|
| 175 |
+
}
|
| 176 |
+
cb_pop_front(cb_c, 1);
|
| 177 |
+
cb_pop_front(cb_max_scaler, 1);
|
| 178 |
+
cb_pop_front(cb_sum_scaler, 1);
|
| 179 |
+
}
|
code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Attention score scale + mask + softmax (tt/smsm_kernel.py), reader (RISCV_0). Once: the scale tile (fp32 filled
|
| 3 |
+
// with the scale) and the two reduce scaler tiles of the stock softmax reader (fp32 1.0 in row 0 of every face, the
|
| 4 |
+
// rest 0: dataflow_kernel_lib::calculate_and_prepare_reduce_scaler for MAX / SUM, REDUCE_ROW). Per tile row r of
|
| 5 |
+
// this core's range (head h = r / Mt, query tile row mr = r % Mt): the Wt mask tiles of row mr (+ 1 unused tile when
|
| 6 |
+
// Wt is odd, so the compute can take the mask in pairs), then the Wt fp32 score tiles in pairs (the last pair of an
|
| 7 |
+
// odd row carries one unused tile).
|
| 8 |
+
// CT args: [0] Wt, [1] Mt, [2] scale bits (fp32), [3] per-core RT-arg count (P3), [4] mask tile bytes, then the
|
| 9 |
+
// TensorAccessorArgs of scores and mask. Common RT args: [s_addr, m_addr]. Per-core RT args: [r0, n].
|
| 10 |
+
#include <cstdint>
|
| 11 |
+
|
| 12 |
+
#include "api/dataflow/dataflow_api.h"
|
| 13 |
+
|
| 14 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 15 |
+
constexpr uint32_t Mt = get_compile_time_arg_val(1);
|
| 16 |
+
constexpr uint32_t scale_bits = get_compile_time_arg_val(2);
|
| 17 |
+
constexpr uint32_t TBM = get_compile_time_arg_val(4);
|
| 18 |
+
constexpr auto s_args = TensorAccessorArgs<5>();
|
| 19 |
+
constexpr auto m_args = TensorAccessorArgs<s_args.next_compile_time_args_offset()>();
|
| 20 |
+
constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_max_scaler = 3, cb_sum_scaler = 4;
|
| 21 |
+
constexpr uint32_t TB = 4096;
|
| 22 |
+
constexpr uint32_t Wp = (Wt + 1) / 2;
|
| 23 |
+
|
| 24 |
+
inline void fill_scaler(uint32_t cb) {
|
| 25 |
+
cb_reserve_back(cb, 1);
|
| 26 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb));
|
| 27 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 28 |
+
p[i] = 0;
|
| 29 |
+
}
|
| 30 |
+
for (uint32_t f = 0; f < 4; ++f) {
|
| 31 |
+
for (uint32_t c = 0; c < 16; ++c) {
|
| 32 |
+
p[f * 256 + c] = 0x3F800000u;
|
| 33 |
+
}
|
| 34 |
+
}
|
| 35 |
+
cb_push_back(cb, 1);
|
| 36 |
+
}
|
| 37 |
+
|
| 38 |
+
void kernel_main() {
|
| 39 |
+
const uint32_t s_addr = get_common_arg_val<uint32_t>(0);
|
| 40 |
+
const uint32_t m_addr = get_common_arg_val<uint32_t>(1);
|
| 41 |
+
const uint32_t r0 = get_arg_val<uint32_t>(0);
|
| 42 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 43 |
+
if (n == 0) {
|
| 44 |
+
return;
|
| 45 |
+
}
|
| 46 |
+
const auto s = TensorAccessor(s_args, s_addr, TB);
|
| 47 |
+
const auto mk = TensorAccessor(m_args, m_addr, TBM);
|
| 48 |
+
cb_reserve_back(cb_c, 1);
|
| 49 |
+
{
|
| 50 |
+
auto* p = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_c));
|
| 51 |
+
for (uint32_t i = 0; i < 1024; ++i) {
|
| 52 |
+
p[i] = scale_bits;
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
cb_push_back(cb_c, 1);
|
| 56 |
+
fill_scaler(cb_max_scaler);
|
| 57 |
+
fill_scaler(cb_sum_scaler);
|
| 58 |
+
for (uint32_t r = r0; r < r0 + n; ++r) {
|
| 59 |
+
const uint32_t mr = r % Mt;
|
| 60 |
+
cb_reserve_back(cb_m, 2 * Wp);
|
| 61 |
+
{
|
| 62 |
+
uint32_t p = get_write_ptr(cb_m);
|
| 63 |
+
for (uint32_t j = 0; j < Wt; ++j) {
|
| 64 |
+
noc_async_read(mk.get_noc_addr(mr * Wt + j), p, TBM);
|
| 65 |
+
p += TBM;
|
| 66 |
+
}
|
| 67 |
+
}
|
| 68 |
+
noc_async_read_barrier();
|
| 69 |
+
cb_push_back(cb_m, 2 * Wp);
|
| 70 |
+
uint32_t t = r * Wt;
|
| 71 |
+
for (uint32_t pp = 0; pp < Wp; ++pp) {
|
| 72 |
+
cb_reserve_back(cb_s, 2);
|
| 73 |
+
const uint32_t p = get_write_ptr(cb_s);
|
| 74 |
+
noc_async_read(s.get_noc_addr(t), p, TB);
|
| 75 |
+
if (2 * pp + 1 < Wt) {
|
| 76 |
+
noc_async_read(s.get_noc_addr(t + 1), p + TB, TB);
|
| 77 |
+
}
|
| 78 |
+
noc_async_read_barrier();
|
| 79 |
+
cb_push_back(cb_s, 2);
|
| 80 |
+
t += 2;
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
}
|
code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
// Attention score scale + mask + softmax (tt/smsm_kernel.py), writer (RISCV_1): per tile row r of this core's range
|
| 3 |
+
// the Wt probability tiles in the compute's blocks of ndst (the last one clamped), then the row's pad tiles (pushed
|
| 4 |
+
// by the compute to keep cb_out aligned; not written).
|
| 5 |
+
// CT args: [0] Wt, [1] ndst, [2] per-core RT-arg count (P3), [3] out pad tiles per row, then the TensorAccessorArgs of
|
| 6 |
+
// out. Common RT args: [out_addr]. Per-core RT args: [r0, n].
|
| 7 |
+
#include <cstdint>
|
| 8 |
+
|
| 9 |
+
#include "api/dataflow/dataflow_api.h"
|
| 10 |
+
|
| 11 |
+
constexpr uint32_t Wt = get_compile_time_arg_val(0);
|
| 12 |
+
constexpr uint32_t ndst = get_compile_time_arg_val(1);
|
| 13 |
+
constexpr uint32_t out_pad = get_compile_time_arg_val(3);
|
| 14 |
+
constexpr auto o_args = TensorAccessorArgs<4>();
|
| 15 |
+
constexpr uint32_t cb_out = 16;
|
| 16 |
+
constexpr uint32_t TB = 4096;
|
| 17 |
+
|
| 18 |
+
void kernel_main() {
|
| 19 |
+
const uint32_t o_addr = get_common_arg_val<uint32_t>(0);
|
| 20 |
+
const uint32_t r0 = get_arg_val<uint32_t>(0);
|
| 21 |
+
const uint32_t n = get_arg_val<uint32_t>(1);
|
| 22 |
+
const auto o = TensorAccessor(o_args, o_addr, TB);
|
| 23 |
+
for (uint32_t r = r0; r < r0 + n; ++r) {
|
| 24 |
+
uint32_t t = r * Wt;
|
| 25 |
+
for (uint32_t j = 0; j < Wt; j += ndst) {
|
| 26 |
+
const uint32_t b = (j + ndst > Wt) ? (Wt - j) : ndst;
|
| 27 |
+
cb_wait_front(cb_out, b);
|
| 28 |
+
uint32_t p = get_read_ptr(cb_out);
|
| 29 |
+
for (uint32_t i = 0; i < b; ++i) {
|
| 30 |
+
noc_async_write(p, o.get_noc_addr(t + i), TB);
|
| 31 |
+
p += TB;
|
| 32 |
+
}
|
| 33 |
+
noc_async_writes_flushed();
|
| 34 |
+
cb_pop_front(cb_out, b);
|
| 35 |
+
t += b;
|
| 36 |
+
}
|
| 37 |
+
if constexpr (out_pad > 0) {
|
| 38 |
+
cb_wait_front(cb_out, out_pad);
|
| 39 |
+
cb_pop_front(cb_out, out_pad);
|
| 40 |
+
}
|
| 41 |
+
}
|
| 42 |
+
noc_async_write_barrier();
|
| 43 |
+
}
|
code/tt_diffusion_planner/tt/layers.py
CHANGED
|
@@ -19,8 +19,8 @@ from ..ttaw.tensors import to_device, ttnn_dtype
|
|
| 19 |
from . import config as T
|
| 20 |
from .params import Lin, Norm
|
| 21 |
|
| 22 |
-
__all__ = ["Build", "Linear", "SplitLinear", "make_linear", "LayerNorm", "Const", "ATTN", "policy",
|
| 23 |
-
"layer_norm_fp32"]
|
| 24 |
|
| 25 |
ATTN = "bfloat16" # dtype of the Q / K / V projections (SDPA takes bf16)
|
| 26 |
|
|
@@ -50,6 +50,30 @@ class Build:
|
|
| 50 |
self.split = tuple(T.globs(knobs.SPLIT_MATMUL) if split is None else split)
|
| 51 |
self.attn_fp32 = tuple(T.globs(knobs.ATTN_FP32_ACC) if attn_fp32_acc is None else attn_fp32_acc)
|
| 52 |
self.attn_mm = tuple(T.globs(knobs.ATTN_MATMUL) if attn_matmul is None else attn_matmul)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
self.uploaded_bytes = 0
|
| 54 |
|
| 55 |
@staticmethod
|
|
@@ -77,7 +101,16 @@ class Build:
|
|
| 77 |
|
| 78 |
def options(self) -> Dict[str, Any]:
|
| 79 |
return {"ln_fp32": list(self.ln_fp32), "hidden_fp32": list(self.hidden_fp32), "split": list(self.split),
|
| 80 |
-
"attn_fp32_acc": list(self.attn_fp32), "attn_matmul": list(self.attn_mm)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 81 |
|
| 82 |
def prec(self, module: str) -> Precision:
|
| 83 |
return self.policy.resolve(module)
|
|
@@ -89,6 +122,33 @@ class Build:
|
|
| 89 |
"""The residual-stream dtype of ``module``."""
|
| 90 |
return self.prec(module).activations
|
| 91 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 92 |
def upload(self, array: np.ndarray, dtype: str, *, shape4: bool = True):
|
| 93 |
"""Host float array -> DRAM TILE device tensor (rank padded to 4 with leading 1s)."""
|
| 94 |
a = np.ascontiguousarray(np.asarray(array, np.float32))
|
|
@@ -98,6 +158,251 @@ class Build:
|
|
| 98 |
return to_device(a, self.device, dtype)
|
| 99 |
|
| 100 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
class Linear:
|
| 102 |
"""``y = x @ w (+ b)`` with fused ``activation`` ("gelu" = exact erf, "gelu_tanh"); output dtype ``out``."""
|
| 103 |
|
|
@@ -110,12 +415,28 @@ class Linear:
|
|
| 110 |
self.b = None if lin.b is None else build.upload(np.asarray(lin.b).reshape(1, -1), wd)
|
| 111 |
self.cfg = build.cfg(module)
|
| 112 |
self.shape = tuple(lin.w.shape)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 113 |
|
| 114 |
def __call__(self, x):
|
| 115 |
import ttnn
|
| 116 |
|
| 117 |
-
|
| 118 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 119 |
|
| 120 |
|
| 121 |
def layer_norm_fp32(x, gamma=None, beta=None, *, eps: float = C.LN_EPS):
|
|
@@ -136,6 +457,26 @@ def layer_norm_fp32(x, gamma=None, beta=None, *, eps: float = C.LN_EPS):
|
|
| 136 |
return y
|
| 137 |
|
| 138 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 139 |
class SplitLinear:
|
| 140 |
"""``y = x @ w (+ b)`` to ~1e-5 relative from three device matmuls on split operands (the "bf16x3" scheme):
|
| 141 |
``x_hi @ w_hi + x_hi @ w_lo + x_lo @ w_hi`` with ``*_hi`` the bf16 roundings (bf16 x bf16 products are exact
|
|
@@ -158,27 +499,193 @@ class SplitLinear:
|
|
| 158 |
self.activation = activation
|
| 159 |
self.out = ttnn_dtype(out)
|
| 160 |
self.shape = tuple(w.shape)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 161 |
|
| 162 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 163 |
import ttnn
|
| 164 |
|
| 165 |
f32 = ttnn.float32
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 166 |
if x.dtype != f32: # a bf16 input is its own hi part: two matmuls
|
| 167 |
-
y = ttnn.matmul(x, self.w_hi, dtype=f32,
|
| 168 |
-
y = ttnn.add(y, ttnn.linear(x, self.w_lo, bias=self.b, dtype=f32,
|
| 169 |
else:
|
| 170 |
x_hi = ttnn.typecast(x, ttnn.bfloat16)
|
| 171 |
x_lo = ttnn.subtract(x, ttnn.typecast(x_hi, f32))
|
| 172 |
-
y = ttnn.matmul(x_hi, self.w_hi, dtype=f32,
|
| 173 |
-
y = ttnn.add(y, ttnn.matmul(x_hi, self.w_lo, dtype=f32,
|
| 174 |
-
y = ttnn.add(y, ttnn.linear(x_lo, self.w_hi, bias=self.b, dtype=f32,
|
| 175 |
if self.activation == "gelu":
|
| 176 |
y = ttnn.gelu(y, fast_and_approximate_mode=False)
|
| 177 |
elif self.activation == "gelu_tanh":
|
| 178 |
y = ttnn.gelu(y, variant=ttnn.GeluVariant.Tanh)
|
| 179 |
if self.out != f32:
|
| 180 |
y = ttnn.typecast(y, self.out)
|
| 181 |
-
return y
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 182 |
|
| 183 |
|
| 184 |
def make_linear(build: Build, lin: Lin, module: str, *, out: str, activation: Optional[str] = None):
|
|
@@ -196,26 +703,101 @@ class LayerNorm:
|
|
| 196 |
def __init__(self, build: Build, nrm: Optional[Norm], module: str):
|
| 197 |
self.cfg = build.cfg(module)
|
| 198 |
self.mode = build.ln_mode(module)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 199 |
self.g = self.b = None
|
| 200 |
if nrm is not None:
|
| 201 |
self.g = build.upload(np.asarray(nrm.gamma).reshape(1, -1), "float32")
|
| 202 |
self.b = build.upload(np.asarray(nrm.beta).reshape(1, -1), "float32")
|
| 203 |
|
| 204 |
-
def __call__(self, x, gamma=None, beta=None):
|
|
|
|
|
|
|
| 205 |
import ttnn
|
| 206 |
|
| 207 |
g = self.g if gamma is None else gamma
|
| 208 |
b = self.b if beta is None else beta
|
| 209 |
if self.mode == "fp32":
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 210 |
return layer_norm_fp32(x, g, b)
|
| 211 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
|
| 213 |
|
| 214 |
class Const:
|
| 215 |
"""A constant device tensor (upload once)."""
|
| 216 |
|
| 217 |
-
def __init__(self, build: Build, array: np.ndarray, dtype: str):
|
| 218 |
self.t = build.upload(array, dtype)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 219 |
|
| 220 |
def __call__(self):
|
| 221 |
return self.t
|
|
|
|
| 19 |
from . import config as T
|
| 20 |
from .params import Lin, Norm
|
| 21 |
|
| 22 |
+
__all__ = ["Build", "Linear", "SplitLinear", "make_linear", "chain2", "LayerNorm", "Const", "ATTN", "policy",
|
| 23 |
+
"layer_norm_fp32", "fold_batch"]
|
| 24 |
|
| 25 |
ATTN = "bfloat16" # dtype of the Q / K / V projections (SDPA takes bf16)
|
| 26 |
|
|
|
|
| 50 |
self.split = tuple(T.globs(knobs.SPLIT_MATMUL) if split is None else split)
|
| 51 |
self.attn_fp32 = tuple(T.globs(knobs.ATTN_FP32_ACC) if attn_fp32_acc is None else attn_fp32_acc)
|
| 52 |
self.attn_mm = tuple(T.globs(knobs.ATTN_MATMUL) if attn_matmul is None else attn_matmul)
|
| 53 |
+
self.ch2d = bool(knobs.ENC_CH2D)
|
| 54 |
+
self.attn_fast = int(knobs.ATTN_FAST)
|
| 55 |
+
self.dec_mmcfg = bool(knobs.DEC_MMCFG) and hasattr(device, "compute_with_storage_grid_size")
|
| 56 |
+
on_device = hasattr(device, "compute_with_storage_grid_size")
|
| 57 |
+
self.ln_kernel = int(knobs.LN_KERNEL) if on_device else 0
|
| 58 |
+
self.ln_resid = bool(knobs.LN_RESID)
|
| 59 |
+
self.split_kcat = int(knobs.SPLIT_KCAT)
|
| 60 |
+
self.attn_smask = bool(knobs.ATTN_SMASK)
|
| 61 |
+
self.attn_smsm = int(knobs.ATTN_SMSM)
|
| 62 |
+
self.attn_fused = bool(knobs.ATTN_FUSED)
|
| 63 |
+
self.kcat_emit = bool(knobs.KCAT_EMIT)
|
| 64 |
+
self.ln_tr = bool(knobs.LN_TR)
|
| 65 |
+
self.enc_kcat = bool(knobs.ENC_KCAT)
|
| 66 |
+
self.lin_act = bool(knobs.LIN_ACT)
|
| 67 |
+
self.kcat_act = bool(knobs.KCAT_ACT)
|
| 68 |
+
self.ln_split = bool(knobs.LN_SPLIT)
|
| 69 |
+
self.ln_sfpu_bcast = bool(knobs.LN_SFPU_BCAST)
|
| 70 |
+
self.kcat_l1 = bool(knobs.KCAT_L1)
|
| 71 |
+
self.kcat_act_once = bool(knobs.KCAT_ACT_ONCE)
|
| 72 |
+
self.attn_l1 = int(knobs.ATTN_L1) if on_device else 0
|
| 73 |
+
self.dec_l1 = int(knobs.DEC_L1) if on_device else 0
|
| 74 |
+
self.enc_l1 = bool(knobs.ENC_L1) and on_device
|
| 75 |
+
self.fus_l1 = bool(knobs.FUS_L1) and on_device
|
| 76 |
+
self.compact_enc = int(knobs.COMPACT) >= 2
|
| 77 |
self.uploaded_bytes = 0
|
| 78 |
|
| 79 |
@staticmethod
|
|
|
|
| 101 |
|
| 102 |
def options(self) -> Dict[str, Any]:
|
| 103 |
return {"ln_fp32": list(self.ln_fp32), "hidden_fp32": list(self.hidden_fp32), "split": list(self.split),
|
| 104 |
+
"attn_fp32_acc": list(self.attn_fp32), "attn_matmul": list(self.attn_mm), "enc_ch2d": self.ch2d,
|
| 105 |
+
"attn_fast": self.attn_fast,
|
| 106 |
+
"dec_mmcfg": self.dec_mmcfg, "ln_kernel": self.ln_kernel, "ln_resid": self.ln_resid,
|
| 107 |
+
"split_kcat": self.split_kcat, "attn_smask": self.attn_smask, "attn_smsm": self.attn_smsm, "attn_fused": self.attn_fused, "kcat_emit": self.kcat_emit, "ln_tr": self.ln_tr,
|
| 108 |
+
"enc_kcat": self.enc_kcat, "lin_act": self.lin_act,
|
| 109 |
+
"kcat_act": self.kcat_act, "ln_split": self.ln_split,
|
| 110 |
+
"ln_sfpu_bcast": self.ln_sfpu_bcast, "kcat_l1": self.kcat_l1,
|
| 111 |
+
"kcat_act_once": self.kcat_act_once, "attn_l1": self.attn_l1,
|
| 112 |
+
"dec_l1": self.dec_l1, "enc_l1": self.enc_l1,
|
| 113 |
+
"fus_l1": self.fus_l1}
|
| 114 |
|
| 115 |
def prec(self, module: str) -> Precision:
|
| 116 |
return self.policy.resolve(module)
|
|
|
|
| 122 |
"""The residual-stream dtype of ``module``."""
|
| 123 |
return self.prec(module).activations
|
| 124 |
|
| 125 |
+
def attn_mem(self):
|
| 126 |
+
"""``ATTN_L1``: the memory config of the fused attention's inputs (Q / K / V projections, hoisted cross K / V,
|
| 127 |
+
masks): L1 interleaved, else None (DRAM)."""
|
| 128 |
+
if not self.attn_l1:
|
| 129 |
+
return None
|
| 130 |
+
import ttnn
|
| 131 |
+
|
| 132 |
+
return ttnn.L1_MEMORY_CONFIG
|
| 133 |
+
|
| 134 |
+
def dec_mem(self):
|
| 135 |
+
"""``DEC_L1``: the memory config of the decoder blocks' intermediates (stream, LN operands, linear outputs,
|
| 136 |
+
attention outputs): L1 interleaved, else None (DRAM)."""
|
| 137 |
+
if not self.dec_l1:
|
| 138 |
+
return None
|
| 139 |
+
import ttnn
|
| 140 |
+
|
| 141 |
+
return ttnn.L1_MEMORY_CONFIG
|
| 142 |
+
|
| 143 |
+
def enc_mem(self):
|
| 144 |
+
"""``ENC_L1``: the memory config of the mixer blocks' intermediates (LN outputs and stream, token / channel
|
| 145 |
+
MLP outputs): L1 interleaved, else None (DRAM)."""
|
| 146 |
+
if not self.enc_l1:
|
| 147 |
+
return None
|
| 148 |
+
import ttnn
|
| 149 |
+
|
| 150 |
+
return ttnn.L1_MEMORY_CONFIG
|
| 151 |
+
|
| 152 |
def upload(self, array: np.ndarray, dtype: str, *, shape4: bool = True):
|
| 153 |
"""Host float array -> DRAM TILE device tensor (rank padded to 4 with leading 1s)."""
|
| 154 |
a = np.ascontiguousarray(np.asarray(array, np.float32))
|
|
|
|
| 158 |
return to_device(a, self.device, dtype)
|
| 159 |
|
| 160 |
|
| 161 |
+
# OPT round 1 item 2a: explicit 2-D multicast program configs for the decoder's 352-row matmuls, the fastest
|
| 162 |
+
# bit-identical candidate per (K, N, split pass) of the device sweep (logs/diffusion-planner/opt_r1/mm_sweep.json;
|
| 163 |
+
# "hh" = x_hi @ w_hi, "hl" = x_hi @ w_lo, "lh" = x_lo @ w_hi + bias; a plain Linear uses "hh").
|
| 164 |
+
# value: (transpose_mcast, per_core_M, per_core_N, in0_block_w, out_subblock_w); every candidate was bit-identical to
|
| 165 |
+
# the auto config (fp32 DEST accumulation over K either way).
|
| 166 |
+
DEC_ROWS = 352
|
| 167 |
+
DEC_MM_CONFIGS = {
|
| 168 |
+
(324, 512, "hh"): (True, 1, 3, 11, 1), (324, 512, "hl"): (True, 2, 2, 11, 2), (324, 512, "lh"): (True, 1, 3, 11, 1),
|
| 169 |
+
(512, 256, "hh"): (False, 2, 1, 8, 1), (512, 256, "hl"): (False, 2, 1, 8, 1), (512, 256, "lh"): (False, 2, 1, 8, 1),
|
| 170 |
+
(256, 768, "hh"): (False, 2, 3, 8, 1), (256, 768, "hl"): (False, 2, 2, 8, 2), (256, 768, "lh"): (True, 1, 3, 8, 1),
|
| 171 |
+
(256, 256, "hh"): (False, 2, 2, 8, 2), (256, 256, "hl"): (False, 2, 2, 8, 2), (256, 256, "lh"): (False, 2, 1, 8, 1),
|
| 172 |
+
(256, 1024, "hh"): (False, 2, 3, 8, 1), (256, 1024, "hl"): (False, 2, 3, 8, 1),
|
| 173 |
+
(256, 1024, "lh"): (False, 2, 3, 8, 1),
|
| 174 |
+
(1024, 256, "hh"): (True, 1, 1, 8, 1), (1024, 256, "hl"): (True, 1, 1, 8, 1), (1024, 256, "lh"): (False, 2, 1, 8, 1),
|
| 175 |
+
(1024, 324, "hh"): (False, 2, 1, 8, 1), (1024, 324, "hl"): (False, 2, 1, 8, 1),
|
| 176 |
+
(1024, 324, "lh"): (False, 2, 1, 8, 1),
|
| 177 |
+
}
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
def compact_rows(rows: int) -> bool:
|
| 181 |
+
"""``COMPACT`` (OPT round 5): a decoder row count other than :data:`DEC_ROWS` that the derived configs serve (a
|
| 182 |
+
tile-aligned agent bucket below the full capacity)."""
|
| 183 |
+
return rows % 32 == 0 and 32 <= rows < DEC_ROWS
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
def fit_pcm(tr: bool, pcm: int, rows: int, grid) -> int:
|
| 187 |
+
"""``per_core_M`` of a 2-D multicast config for ``rows`` rows: the swept value at :data:`DEC_ROWS`, else the
|
| 188 |
+
smallest value whose M blocks fit the grid axis that carries them (the in0 block width, ``per_core_N`` and the
|
| 189 |
+
subblocks stay as swept: the K accumulation order of every output element is unchanged)."""
|
| 190 |
+
if rows == DEC_ROWS:
|
| 191 |
+
return pcm
|
| 192 |
+
mt = -(-rows // 32)
|
| 193 |
+
lim = grid.x if tr else grid.y
|
| 194 |
+
p = 1
|
| 195 |
+
while -(-mt // p) > lim:
|
| 196 |
+
p += 1
|
| 197 |
+
return p
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
def mcast_fits(tr: bool, pcm: int, pcn: int, rows: int, n: int, grid) -> bool:
|
| 201 |
+
"""True when a 2-D multicast config's output blocks fit the grid: the swept configs assume the 12x10 ETH-dispatch
|
| 202 |
+
grid (N / per_core_N up to 12 blocks on x); on a smaller grid (WORKER dispatch, 11x10) the callers fall back to
|
| 203 |
+
the auto config instead of failing at capture (never hard-code the grid: CLAUDE.md)."""
|
| 204 |
+
mb, nb = -(-(-(-rows // 32)) // pcm), -(-(-(-n // 32)) // pcn)
|
| 205 |
+
gx, gy = (mb, nb) if tr else (nb, mb)
|
| 206 |
+
return gx <= grid.x and gy <= grid.y
|
| 207 |
+
|
| 208 |
+
|
| 209 |
+
class DecConfigs:
|
| 210 |
+
"""``DEC_MMCFG`` configs of one decoder linear: ``get(rows)`` -> ``{"hh" | "hl" | "lh": program config}`` for
|
| 211 |
+
:data:`DEC_ROWS` or (``COMPACT``) a smaller agent bucket, else None."""
|
| 212 |
+
|
| 213 |
+
def __init__(self, build: "Build", k: int, n: int):
|
| 214 |
+
self.build, self.k, self.n, self._cache = build, k, n, {}
|
| 215 |
+
|
| 216 |
+
def get(self, rows: int) -> Optional[Dict[str, Any]]:
|
| 217 |
+
rows = int(rows)
|
| 218 |
+
if rows != DEC_ROWS and not compact_rows(rows):
|
| 219 |
+
return None
|
| 220 |
+
if rows not in self._cache:
|
| 221 |
+
import ttnn
|
| 222 |
+
|
| 223 |
+
grid = self.build.device.compute_with_storage_grid_size()
|
| 224 |
+
out = {}
|
| 225 |
+
for ps in ("hh", "hl", "lh"):
|
| 226 |
+
tr, pcm, pcn, kb, sw = DEC_MM_CONFIGS[(self.k, self.n, ps)]
|
| 227 |
+
if not mcast_fits(tr, fit_pcm(tr, pcm, rows, grid), pcn, rows, self.n, grid):
|
| 228 |
+
out = None # a grid smaller than the swept one: the auto configs
|
| 229 |
+
break
|
| 230 |
+
out[ps] = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig(
|
| 231 |
+
compute_with_storage_grid_size=grid, in0_block_w=kb, out_subblock_h=1, out_subblock_w=sw,
|
| 232 |
+
per_core_M=fit_pcm(tr, pcm, rows, grid), per_core_N=pcn, transpose_mcast=tr,
|
| 233 |
+
fused_activation=None, fuse_batch=True)
|
| 234 |
+
self._cache[rows] = out
|
| 235 |
+
return self._cache[rows]
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
def dec_configs(build: "Build", module: str, shape) -> Optional[DecConfigs]:
|
| 239 |
+
"""The :class:`DecConfigs` of a decoder linear of weight ``shape`` (``DEC_MMCFG``), else None."""
|
| 240 |
+
if not (build.dec_mmcfg and module.startswith("dec.")):
|
| 241 |
+
return None
|
| 242 |
+
k, n = int(shape[0]), int(shape[1])
|
| 243 |
+
if (k, n, "hh") not in DEC_MM_CONFIGS:
|
| 244 |
+
return None
|
| 245 |
+
import ttnn
|
| 246 |
+
|
| 247 |
+
if not hasattr(ttnn, "MatmulMultiCoreReuseMultiCastProgramConfig"):
|
| 248 |
+
return None # the host fake ttnn
|
| 249 |
+
return DecConfigs(build, k, n)
|
| 250 |
+
|
| 251 |
+
|
| 252 |
+
# OPT round 2 item 2 (SPLIT_KCAT): X' row tiles 3 Kt + 1 rounded up to a multiple of kcat_pad(K) (the K block of
|
| 253 |
+
# the program config must divide them), and the fastest 2-D multicast config per (rows, K' tiles, N) of the device
|
| 254 |
+
# sweep (logs/diffusion-planner/opt_r2/kcat_sweep.json): value (transpose_mcast, per_core_M, per_core_N, in0_block_w,
|
| 255 |
+
# out_subblock_w); shapes not listed use the auto config.
|
| 256 |
+
def kcat_pad(k: int) -> int:
|
| 257 |
+
"""Sweep: K = 1024 (3 Kt + 1 = 97, prime) pads to 104 = 8 x 13; the other K keep 3 Kt + 1 (25 = 5 x 5, 49)."""
|
| 258 |
+
return 8 if k >= 1024 else 1
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
KCAT_CONFIGS: Dict[tuple, tuple] = {
|
| 262 |
+
(352, 49, 256): (False, 2, 1, 7, 1), (352, 25, 768): (False, 2, 2, 5, 2), (352, 25, 256): (False, 2, 1, 5, 1),
|
| 263 |
+
(352, 25, 1024): (False, 2, 3, 5, 1), (352, 104, 256): (False, 2, 1, 13, 1),
|
| 264 |
+
(352, 104, 324): (False, 2, 1, 13, 1), (576, 25, 256): (False, 2, 1, 5, 1),
|
| 265 |
+
}
|
| 266 |
+
|
| 267 |
+
|
| 268 |
+
def kcat_config(build: "Build", ktp: int, n: int, rows: int):
|
| 269 |
+
"""The program config of the K-concatenated matmul ``[rows, 32 ktp] @ [32 ktp, n]``: the swept entry, or
|
| 270 |
+
(``COMPACT``) the :data:`DEC_ROWS` entry with ``per_core_M`` fitted to a smaller agent bucket; None = auto."""
|
| 271 |
+
import ttnn
|
| 272 |
+
|
| 273 |
+
if not hasattr(ttnn, "MatmulMultiCoreReuseMultiCastProgramConfig") or not hasattr(
|
| 274 |
+
build.device, "compute_with_storage_grid_size"):
|
| 275 |
+
return None
|
| 276 |
+
grid = build.device.compute_with_storage_grid_size()
|
| 277 |
+
ent = KCAT_CONFIGS.get((rows, ktp, n))
|
| 278 |
+
if ent is None and compact_rows(rows):
|
| 279 |
+
ent = KCAT_CONFIGS.get((DEC_ROWS, ktp, n))
|
| 280 |
+
if ent is None:
|
| 281 |
+
return None
|
| 282 |
+
tr, pcm, pcn, kb, sw = ent
|
| 283 |
+
pcm = fit_pcm(tr, pcm, rows, grid) if (rows, ktp, n) not in KCAT_CONFIGS else pcm
|
| 284 |
+
if not mcast_fits(tr, pcm, pcn, rows, n, grid):
|
| 285 |
+
return None # a grid smaller than the swept one: the auto config
|
| 286 |
+
return ttnn.MatmulMultiCoreReuseMultiCastProgramConfig(
|
| 287 |
+
compute_with_storage_grid_size=grid, in0_block_w=kb, out_subblock_h=1, out_subblock_w=sw,
|
| 288 |
+
per_core_M=pcm, per_core_N=pcn,
|
| 289 |
+
transpose_mcast=tr, fused_activation=None, fuse_batch=True)
|
| 290 |
+
|
| 291 |
+
|
| 292 |
+
# OPT round 4 item 1 (KCAT_L1): the K = 1024 split operands X' [352, 32 Ktp] written L1 block-sharded by the operand
|
| 293 |
+
# build and read in place by a 2-D multicast matmul with in0 sharded (no 4.7 MB DRAM round trip). The K slices of
|
| 294 |
+
# X' are the matmul's N blocks, so Ktp is padded to a multiple of them (zero tiles, exact). Value per (rows, K, N):
|
| 295 |
+
# (transpose_mcast, per_core_M, per_core_N, pad, in0_block_w), the fastest bit-identical candidate of the device
|
| 296 |
+
# check (logs/diffusion-planner/opt_r4/kl1_real.log: operand + matmul 62.6 -> 41.3 us for N = 256, 62.3 -> 41.8 us
|
| 297 |
+
# for N = 324).
|
| 298 |
+
KCAT_L1_CONFIGS: Dict[tuple, tuple] = {
|
| 299 |
+
(352, 1024, 256): (False, 2, 1, 8, 13),
|
| 300 |
+
(352, 1024, 324): (False, 2, 1, 11, 9),
|
| 301 |
+
}
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
def kcat_l1_layout(build: "Build", rows: int, k: int, n: int, ktp: int):
|
| 305 |
+
"""``(memory_config, program_config)`` of the L1-sharded operand path for ``[rows, 32 ktp] @ [32 ktp, n]``."""
|
| 306 |
+
import ttnn
|
| 307 |
+
|
| 308 |
+
tr, pcm, pcn, _pad, kb = KCAT_L1_CONFIGS[(DEC_ROWS, k, n)]
|
| 309 |
+
grid = build.device.compute_with_storage_grid_size()
|
| 310 |
+
pcm = fit_pcm(tr, pcm, rows, grid) # COMPACT: a smaller agent bucket (the K blocking unchanged)
|
| 311 |
+
mt, nt = -(-rows // 32), -(-n // 32)
|
| 312 |
+
nb, mb = -(-nt // pcn), -(-mt // pcm)
|
| 313 |
+
assert ktp % nb == 0 and (ktp // nb) % kb == 0, (ktp, nb, kb)
|
| 314 |
+
gx, gy = (mb, nb) if tr else (nb, mb)
|
| 315 |
+
if gx > grid.x or gy > grid.y:
|
| 316 |
+
return None # a grid smaller than the swept one: the DRAM operand path
|
| 317 |
+
crs = ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))})
|
| 318 |
+
spec = ttnn.ShardSpec(crs, [32 * pcm, 32 * (ktp // nb)],
|
| 319 |
+
ttnn.ShardOrientation.COL_MAJOR if tr else ttnn.ShardOrientation.ROW_MAJOR)
|
| 320 |
+
mem = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.BLOCK_SHARDED, ttnn.BufferType.L1, spec)
|
| 321 |
+
pc = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig(
|
| 322 |
+
compute_with_storage_grid_size=grid, in0_block_w=kb, out_subblock_h=1,
|
| 323 |
+
out_subblock_w=max(d for d in (1, 2, 4) if pcn % d == 0), per_core_M=pcm, per_core_N=pcn,
|
| 324 |
+
transpose_mcast=tr, fused_activation=None, fuse_batch=True)
|
| 325 |
+
return mem, pc
|
| 326 |
+
|
| 327 |
+
|
| 328 |
+
def _activate(y, act):
|
| 329 |
+
"""The stock activation programs of the linears (``None`` / ``"gelu"`` / ``"gelu_tanh"``)."""
|
| 330 |
+
import ttnn
|
| 331 |
+
|
| 332 |
+
if act == "gelu":
|
| 333 |
+
return ttnn.gelu(y, fast_and_approximate_mode=False)
|
| 334 |
+
if act == "gelu_tanh":
|
| 335 |
+
return ttnn.gelu(y, variant=ttnn.GeluVariant.Tanh)
|
| 336 |
+
return y
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
def _identity(t):
|
| 340 |
+
return t
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
def fold_batch(x, n_out: int, enabled: bool):
|
| 344 |
+
"""Run a linear on a batched activation ``[B0, B1, T, K] @ [K, N]`` as one 2-D matmul over all ``B0*B1*T`` rows
|
| 345 |
+
(``ENC_CH2D``, OPT round 1 item 1). The stock auto-config tiles M per batch element (``fuse_batch=False``), which
|
| 346 |
+
puts the encoder's ``[1, E, T, C]`` channel MLPs on 4-8 cores (970 us at E = 320). Returns ``(x2, back, pc)``:
|
| 347 |
+
|
| 348 |
+
- ``T % 32 == 0`` (the MixerBlock channel MLPs, T = 64): the free TILE view ``[1, 1, B*T, K]`` (auto config,
|
| 349 |
+
2-D multicast over the grid), ``back`` reshapes the output to ``[B0, B1, T, N]``;
|
| 350 |
+
- otherwise (the pre-projections, T = 6 / 20 / 40 padded per element): an explicit 1-D in1-multicast program config
|
| 351 |
+
with ``fuse_batch=True`` (rows of tiles split over the grid, the full N per core, ``in0_block_w = Kt``);
|
| 352 |
+
``back`` is the identity.
|
| 353 |
+
|
| 354 |
+
The products and their fp32 DEST accumulation over K are the same; only the core assignment changes."""
|
| 355 |
+
import ttnn
|
| 356 |
+
|
| 357 |
+
shape = tuple(x.shape)
|
| 358 |
+
if not enabled or len(shape) != 4 or shape[0] * shape[1] == 1:
|
| 359 |
+
return x, _identity, None
|
| 360 |
+
b0, b1, t, k = shape
|
| 361 |
+
if t % 32 == 0:
|
| 362 |
+
x2 = ttnn.reshape(x, (1, 1, b0 * b1 * t, k))
|
| 363 |
+
return x2, (lambda y: ttnn.reshape(y, (b0, b1, t, y.shape[-1]))), None
|
| 364 |
+
cfg_cls = getattr(ttnn, "MatmulMultiCoreReuseMultiCast1DProgramConfig", None)
|
| 365 |
+
dev = x.device() if callable(getattr(x, "device", None)) else None
|
| 366 |
+
if cfg_cls is None or dev is None or not hasattr(dev, "compute_with_storage_grid_size"):
|
| 367 |
+
return x, _identity, None # the host fake: numerics are the same
|
| 368 |
+
grid = dev.compute_with_storage_grid_size()
|
| 369 |
+
cores = grid.x * grid.y
|
| 370 |
+
m_tiles = b0 * b1 * (-(-t // 32))
|
| 371 |
+
kt, nt = -(-k // 32), -(-n_out // 32)
|
| 372 |
+
per_core_m = -(-m_tiles // cores)
|
| 373 |
+
sub_w = max(d for d in (1, 2, 4) if nt % d == 0) # fp32 DEST: subblock h * w <= 4
|
| 374 |
+
pc = cfg_cls(compute_with_storage_grid_size=grid, in0_block_w=kt, out_subblock_h=1, out_subblock_w=sub_w,
|
| 375 |
+
per_core_M=per_core_m, per_core_N=nt, fuse_batch=True, fused_activation=None, mcast_in0=False)
|
| 376 |
+
return x, _identity, pc
|
| 377 |
+
|
| 378 |
+
|
| 379 |
+
def tall_config(x, wshape, fused_activation=None):
|
| 380 |
+
"""The stock auto config of a tall 2-D matmul ``[1, 1, M, K] @ [K, N]`` (M / 32 >= the grid, K, N <= 128): the
|
| 381 |
+
1-D in1-multicast config the profile shows for the mixer linears (``MatmulMultiCoreReuseMultiCast1DProgramConfig``,
|
| 382 |
+
``per_core_M = ceil(Mt / cores)``, ``in0_block_w = min(Kt, 2)``, the full N per core, ``mcast_in0 = False``), so a
|
| 383 |
+
fused activation can be added without changing the matmul (``LIN_ACT``). None for any other shape."""
|
| 384 |
+
import ttnn
|
| 385 |
+
|
| 386 |
+
shp = tuple(x.shape)
|
| 387 |
+
dev = x.device() if callable(getattr(x, "device", None)) else None
|
| 388 |
+
cls = getattr(ttnn, "MatmulMultiCoreReuseMultiCast1DProgramConfig", None)
|
| 389 |
+
if cls is None or dev is None or not hasattr(dev, "compute_with_storage_grid_size") or len(shp) != 4:
|
| 390 |
+
return None
|
| 391 |
+
if shp[0] * shp[1] != 1 or shp[-2] % 32 or wshape[0] % 32 or wshape[1] % 32:
|
| 392 |
+
return None
|
| 393 |
+
grid = dev.compute_with_storage_grid_size()
|
| 394 |
+
cores = grid.x * grid.y
|
| 395 |
+
mt, kt, nt = shp[-2] // 32, wshape[0] // 32, wshape[1] // 32
|
| 396 |
+
if mt < cores or nt > 4 or kt > 4:
|
| 397 |
+
return None
|
| 398 |
+
pcm = -(-mt // cores)
|
| 399 |
+
sw = max(d for d in (1, 2, 4) if nt % d == 0 and d <= 4)
|
| 400 |
+
sh = max(d for d in (1, 2, 4) if pcm % d == 0 and d * sw <= 4)
|
| 401 |
+
return cls(compute_with_storage_grid_size=grid, in0_block_w=min(kt, 2), out_subblock_h=sh, out_subblock_w=sw,
|
| 402 |
+
out_block_h=pcm, out_block_w=nt, per_core_M=pcm, per_core_N=nt, fuse_batch=False,
|
| 403 |
+
fused_activation=fused_activation, mcast_in0=False)
|
| 404 |
+
|
| 405 |
+
|
| 406 |
class Linear:
|
| 407 |
"""``y = x @ w (+ b)`` with fused ``activation`` ("gelu" = exact erf, "gelu_tanh"); output dtype ``out``."""
|
| 408 |
|
|
|
|
| 415 |
self.b = None if lin.b is None else build.upload(np.asarray(lin.b).reshape(1, -1), wd)
|
| 416 |
self.cfg = build.cfg(module)
|
| 417 |
self.shape = tuple(lin.w.shape)
|
| 418 |
+
self.fold = build.ch2d and module.startswith("enc.")
|
| 419 |
+
self.dec_pc = dec_configs(build, module, self.shape)
|
| 420 |
+
self.act_fuse = build.lin_act and module.startswith("enc.mixer.")
|
| 421 |
+
self.out_mem = None # ATTN_L1: output memory config (None = DRAM)
|
| 422 |
|
| 423 |
def __call__(self, x):
|
| 424 |
import ttnn
|
| 425 |
|
| 426 |
+
x, back, pc = fold_batch(x, self.shape[1], self.fold)
|
| 427 |
+
dpc = None if self.dec_pc is None else self.dec_pc.get(tuple(x.shape)[-2])
|
| 428 |
+
if dpc is not None and x.dtype == ttnn.bfloat16:
|
| 429 |
+
pc = dpc["hh"]
|
| 430 |
+
act = self.activation
|
| 431 |
+
if self.act_fuse and act == "gelu" and pc is None:
|
| 432 |
+
pc = tall_config(x, self.shape, ttnn.UnaryWithParam(ttnn.UnaryOpType.GELU, 0.0))
|
| 433 |
+
if pc is not None: # LIN_ACT: GELU in the matmul epilogue (fp32 DEST, before rounding)
|
| 434 |
+
act = None
|
| 435 |
+
kw = {} if pc is None else {"program_config": pc}
|
| 436 |
+
if self.out_mem is not None:
|
| 437 |
+
kw["memory_config"] = self.out_mem
|
| 438 |
+
return back(ttnn.linear(x, self.w, bias=self.b, activation=act, dtype=self.out,
|
| 439 |
+
compute_kernel_config=self.cfg, **kw))
|
| 440 |
|
| 441 |
|
| 442 |
def layer_norm_fp32(x, gamma=None, beta=None, *, eps: float = C.LN_EPS):
|
|
|
|
| 457 |
return y
|
| 458 |
|
| 459 |
|
| 460 |
+
class KcatOperand:
|
| 461 |
+
"""A split operand ``[x_hi | x_hi | x_lo | 1 | 0..]`` written by its producer (``KCAT_EMIT``: the split
|
| 462 |
+
LayerNorm, the fused attention) for a K-concatenated :class:`SplitLinear` (instead of x + a ``kcat_operand``
|
| 463 |
+
program). ``t``: the device tensor ``[..., M, 32 * ktp]``."""
|
| 464 |
+
|
| 465 |
+
def __init__(self, t, ktp: int):
|
| 466 |
+
self.t, self.ktp = t, int(ktp)
|
| 467 |
+
|
| 468 |
+
@property
|
| 469 |
+
def shape(self):
|
| 470 |
+
return self.t.shape
|
| 471 |
+
|
| 472 |
+
|
| 473 |
+
def operand_ktp(build: "Build", consumer) -> int:
|
| 474 |
+
"""Row tiles of the split operand ``consumer`` takes when its producer may emit it (``KCAT_EMIT``), else 0."""
|
| 475 |
+
if not getattr(build, "kcat_emit", False) or not isinstance(consumer, SplitLinear):
|
| 476 |
+
return 0
|
| 477 |
+
return consumer.operand_ktp()
|
| 478 |
+
|
| 479 |
+
|
| 480 |
class SplitLinear:
|
| 481 |
"""``y = x @ w (+ b)`` to ~1e-5 relative from three device matmuls on split operands (the "bf16x3" scheme):
|
| 482 |
``x_hi @ w_hi + x_hi @ w_lo + x_lo @ w_hi`` with ``*_hi`` the bf16 roundings (bf16 x bf16 products are exact
|
|
|
|
| 499 |
self.activation = activation
|
| 500 |
self.out = ttnn_dtype(out)
|
| 501 |
self.shape = tuple(w.shape)
|
| 502 |
+
self.fold = build.ch2d and module.startswith("enc.")
|
| 503 |
+
self.dec_pc = dec_configs(build, module, self.shape)
|
| 504 |
+
# OPT round 2 item 2 (SPLIT_KCAT): the three passes as ONE fp32 matmul over the concatenated K axis,
|
| 505 |
+
# [x_hi | x_hi | x_lo | ones] @ [w_hi; w_lo; w_hi; (b_hi, b_lo, 0...)]: the same exact / TF32-truncated
|
| 506 |
+
# products, accumulated in one fp32 DEST sum (the bias exact as two K rows instead of the packer epilogue).
|
| 507 |
+
# Decoder modules with a tile-aligned K only (the 352-row linears); others keep the 3-pass form.
|
| 508 |
+
# OPT round 3 item 2 (ENC_KCAT): the encoder's split linears (pre-projections, the ego / neighbour island)
|
| 509 |
+
# in the same form, any K: each weight block padded to whole tiles with zero rows (the operand keeps the
|
| 510 |
+
# input's tile padding, which meets those zero rows as in the 3-pass form).
|
| 511 |
+
self.kcat = None
|
| 512 |
+
self.out_mem = None # ATTN_L1: output memory config of the K-concatenated matmul
|
| 513 |
+
k = self.shape[0]
|
| 514 |
+
enc_kcat = build.enc_kcat and (module.startswith("enc.pre.") or module.startswith("enc.island."))
|
| 515 |
+
if build.split_kcat and ((module.startswith("dec.") and k % 32 == 0) or enc_kcat):
|
| 516 |
+
rows = np.zeros((32, self.shape[1]), np.float32)
|
| 517 |
+
if lin.b is not None:
|
| 518 |
+
b = np.asarray(lin.b, np.float32).reshape(-1)
|
| 519 |
+
b_hi = round_to_bf16(b)
|
| 520 |
+
rows[0], rows[1] = b_hi, b - b_hi
|
| 521 |
+
from .kcat_kernel import kcat_tiles
|
| 522 |
+
|
| 523 |
+
kp = -(-k // 32) * 32
|
| 524 |
+
self.kcat_kp = kp
|
| 525 |
+
|
| 526 |
+
def padk(a):
|
| 527 |
+
return np.concatenate([a, np.zeros((kp - k, self.shape[1]), np.float32)], axis=0)
|
| 528 |
+
|
| 529 |
+
l1_key = (DEC_ROWS, kp, self.shape[1])
|
| 530 |
+
use_l1 = (build.kcat_l1 and build.split_kcat == 2 and module.startswith("dec.")
|
| 531 |
+
and l1_key in KCAT_L1_CONFIGS and hasattr(build.device, "compute_with_storage_grid_size"))
|
| 532 |
+
self.kcat_padv = KCAT_L1_CONFIGS[l1_key][3] if use_l1 else kcat_pad(kp)
|
| 533 |
+
self.kcat_ktp = kcat_tiles(kp, self.kcat_padv)
|
| 534 |
+
# KCAT_L1: (memory config, program config) of the L1-sharded operand path for DEC_ROWS rows
|
| 535 |
+
self.kcat_l1 = {} if use_l1 else None # {rows: layout}, filled per row count on first use
|
| 536 |
+
pad_rows = np.zeros((32 * self.kcat_ktp - 3 * kp - 32, self.shape[1]), np.float32)
|
| 537 |
+
wk = np.concatenate([padk(w_hi), padk((w - w_hi).astype(np.float32)), padk(w_hi), rows, pad_rows],
|
| 538 |
+
axis=0)
|
| 539 |
+
self.kcat = build.upload(wk, "float32")
|
| 540 |
+
self.kcat_pc = {} # {rows: program config or None}
|
| 541 |
+
self.kcat_mode = build.split_kcat
|
| 542 |
+
self._act_once = build.kcat_act_once
|
| 543 |
+
self._ones = {}
|
| 544 |
+
self._build = build
|
| 545 |
+
|
| 546 |
+
def _ones_tile(self, x):
|
| 547 |
+
"""``[.., M, 32 (Kt' - 3 Kt)]`` fp32 with columns 0 and 1 = 1 (the bias rows of the concatenated weight), the
|
| 548 |
+
rest 0 (the K padding)."""
|
| 549 |
+
key = tuple(tuple(x.shape)[:-1])
|
| 550 |
+
if key not in self._ones:
|
| 551 |
+
o = np.zeros(key + (32 * self.kcat_ktp - 3 * self.kcat_kp,), np.float32)
|
| 552 |
+
o[..., :2] = 1.0
|
| 553 |
+
self._ones[key] = self._build.upload(o, "float32")
|
| 554 |
+
return self._ones[key]
|
| 555 |
+
|
| 556 |
+
def _call_kcat(self, x, pre_act=None):
|
| 557 |
+
import ttnn
|
| 558 |
|
| 559 |
+
f32 = ttnn.float32
|
| 560 |
+
if isinstance(x, KcatOperand): # KCAT_EMIT: the producer wrote the operand
|
| 561 |
+
assert pre_act is None and x.ktp == self.kcat_ktp
|
| 562 |
+
return self._kcat_mm(x.t, x.t)
|
| 563 |
+
if x.dtype != f32:
|
| 564 |
+
x = ttnn.typecast(x, f32)
|
| 565 |
+
from .kcat_kernel import kcat_operand, supported
|
| 566 |
+
|
| 567 |
+
if self.kcat_mode == 2 and supported(x): # else (the host fake ttnn) the stock build: the same values
|
| 568 |
+
rows = tuple(x.shape)[-2]
|
| 569 |
+
if self.kcat_l1 is not None and (rows == DEC_ROWS or compact_rows(rows)) and not self.fold:
|
| 570 |
+
if rows not in self.kcat_l1:
|
| 571 |
+
self.kcat_l1[rows] = kcat_l1_layout(self._build, rows, self.kcat_kp, self.shape[1],
|
| 572 |
+
self.kcat_ktp)
|
| 573 |
+
if self.kcat_l1 is not None and self.kcat_l1.get(rows) is not None and not self.fold:
|
| 574 |
+
mem, pc = self.kcat_l1[rows] # KCAT_L1: X' in L1, read in place by the sharded-in0 matmul
|
| 575 |
+
xk = kcat_operand(x, self.kcat_padv, pre_act, memory_config=mem, act_once=self._act_once)
|
| 576 |
+
y = ttnn.matmul(xk, self.kcat, dtype=f32, compute_kernel_config=self.cfg, program_config=pc,
|
| 577 |
+
memory_config=self.out_mem or ttnn.DRAM_MEMORY_CONFIG)
|
| 578 |
+
xk.deallocate()
|
| 579 |
+
return y
|
| 580 |
+
return self._kcat_mm(kcat_operand(x, self.kcat_padv, pre_act, act_once=self._act_once), x)
|
| 581 |
+
x = _activate(x, pre_act)
|
| 582 |
+
x_hi = ttnn.typecast(ttnn.typecast(x, ttnn.bfloat16), f32)
|
| 583 |
+
x_lo = ttnn.subtract(x, x_hi)
|
| 584 |
+
xk = ttnn.concat([x_hi, x_hi, x_lo, self._ones_tile(x)], dim=-1)
|
| 585 |
+
return self._kcat_mm(xk, x)
|
| 586 |
+
|
| 587 |
+
def _kcat_mm(self, xk, x):
|
| 588 |
+
import ttnn
|
| 589 |
+
|
| 590 |
+
if self.fold: # encoder [1, E, T, K'] operands: one 2-D matmul (ENC_CH2D)
|
| 591 |
+
xk, back, pc = fold_batch(xk, self.shape[1], True)
|
| 592 |
+
kw = {} if pc is None else {"program_config": pc}
|
| 593 |
+
if self.out_mem is not None:
|
| 594 |
+
kw["memory_config"] = self.out_mem
|
| 595 |
+
return back(ttnn.matmul(xk, self.kcat, dtype=ttnn.float32, compute_kernel_config=self.cfg, **kw))
|
| 596 |
+
kw = self._kcat_kw(x)
|
| 597 |
+
if self.out_mem is not None:
|
| 598 |
+
kw["memory_config"] = self.out_mem
|
| 599 |
+
return ttnn.matmul(xk, self.kcat, dtype=ttnn.float32, compute_kernel_config=self.cfg, **kw)
|
| 600 |
+
|
| 601 |
+
def _kcat_ok(self, x) -> bool:
|
| 602 |
+
"""The K-concatenated form applies: tile-aligned K (any operand build), or the generic_op operand build
|
| 603 |
+
(which keeps the input's tile padding for an unaligned K)."""
|
| 604 |
+
if self.kcat is None:
|
| 605 |
+
return False
|
| 606 |
+
if isinstance(x, KcatOperand) or self.kcat_kp == self.shape[0]:
|
| 607 |
+
return True
|
| 608 |
+
from .kcat_kernel import supported
|
| 609 |
+
|
| 610 |
+
return self.kcat_mode == 2 and supported(x)
|
| 611 |
+
|
| 612 |
+
def _kcat_kw(self, x):
|
| 613 |
+
rows = int(tuple(x.shape)[-2])
|
| 614 |
+
if rows not in self.kcat_pc:
|
| 615 |
+
self.kcat_pc[rows] = kcat_config(self._build, self.kcat_ktp, self.shape[1], rows)
|
| 616 |
+
pc = self.kcat_pc[rows]
|
| 617 |
+
return {} if pc is None else {"program_config": pc}
|
| 618 |
+
|
| 619 |
+
def operand_ktp(self) -> int:
|
| 620 |
+
"""Row tiles of the operand a producer may write for this linear (``KCAT_EMIT``): the generic_op form with a
|
| 621 |
+
tile-aligned K and no K-block padding of a different layout; else 0."""
|
| 622 |
+
if self.kcat is None or self.kcat_mode != 2 or self.kcat_kp != self.shape[0]:
|
| 623 |
+
return 0
|
| 624 |
+
return self.kcat_ktp
|
| 625 |
+
|
| 626 |
+
def fuses_input_act(self) -> bool:
|
| 627 |
+
"""True when this linear can take its input before the previous linear's activation (``KCAT_ACT``)."""
|
| 628 |
+
return self.kcat is not None and self.kcat_mode == 2 and self.kcat_kp == self.shape[0]
|
| 629 |
+
|
| 630 |
+
def __call__(self, x, *, pre_act=None, defer_act: bool = False):
|
| 631 |
+
"""``pre_act``: apply ``"gelu"`` / ``"gelu_tanh"`` to ``x`` first (fused into the operand build when
|
| 632 |
+
:meth:`fuses_input_act`); ``defer_act``: return the pre-activation fp32 output (the next linear applies
|
| 633 |
+
:attr:`activation` through ``pre_act``)."""
|
| 634 |
import ttnn
|
| 635 |
|
| 636 |
f32 = ttnn.float32
|
| 637 |
+
if self._kcat_ok(x):
|
| 638 |
+
y = self._call_kcat(x, pre_act)
|
| 639 |
+
if defer_act:
|
| 640 |
+
return y
|
| 641 |
+
y = _activate(y, self.activation)
|
| 642 |
+
if self.out != f32:
|
| 643 |
+
y = ttnn.typecast(y, self.out)
|
| 644 |
+
return y
|
| 645 |
+
x = _activate(x, pre_act)
|
| 646 |
+
if defer_act:
|
| 647 |
+
assert self.out == f32
|
| 648 |
+
act, self.activation = self.activation, None
|
| 649 |
+
try:
|
| 650 |
+
return self(x)
|
| 651 |
+
finally:
|
| 652 |
+
self.activation = act
|
| 653 |
+
x, back, pc = fold_batch(x, self.shape[1], self.fold)
|
| 654 |
+
pcs = {"hh": pc, "hl": pc, "lh": pc}
|
| 655 |
+
dpc = None if self.dec_pc is None else self.dec_pc.get(tuple(x.shape)[-2])
|
| 656 |
+
if dpc is not None:
|
| 657 |
+
pcs = dpc
|
| 658 |
+
|
| 659 |
+
def kw(ps):
|
| 660 |
+
p = pcs[ps]
|
| 661 |
+
return {"compute_kernel_config": self.cfg, **({} if p is None else {"program_config": p})}
|
| 662 |
if x.dtype != f32: # a bf16 input is its own hi part: two matmuls
|
| 663 |
+
y = ttnn.matmul(x, self.w_hi, dtype=f32, **kw("hh"))
|
| 664 |
+
y = ttnn.add(y, ttnn.linear(x, self.w_lo, bias=self.b, dtype=f32, **kw("hl")))
|
| 665 |
else:
|
| 666 |
x_hi = ttnn.typecast(x, ttnn.bfloat16)
|
| 667 |
x_lo = ttnn.subtract(x, ttnn.typecast(x_hi, f32))
|
| 668 |
+
y = ttnn.matmul(x_hi, self.w_hi, dtype=f32, **kw("hh"))
|
| 669 |
+
y = ttnn.add(y, ttnn.matmul(x_hi, self.w_lo, dtype=f32, **kw("hl")))
|
| 670 |
+
y = ttnn.add(y, ttnn.linear(x_lo, self.w_hi, bias=self.b, dtype=f32, **kw("lh")))
|
| 671 |
if self.activation == "gelu":
|
| 672 |
y = ttnn.gelu(y, fast_and_approximate_mode=False)
|
| 673 |
elif self.activation == "gelu_tanh":
|
| 674 |
y = ttnn.gelu(y, variant=ttnn.GeluVariant.Tanh)
|
| 675 |
if self.out != f32:
|
| 676 |
y = ttnn.typecast(y, self.out)
|
| 677 |
+
return back(y)
|
| 678 |
+
|
| 679 |
+
|
| 680 |
+
def chain2(first, second, x, enabled: bool):
|
| 681 |
+
"""``second(first(x))``; with ``enabled`` (``KCAT_ACT``) and a K-concatenated ``second``, ``first``'s activation
|
| 682 |
+
runs inside ``second``'s operand build instead of as its own program (same LLK, same values)."""
|
| 683 |
+
import ttnn
|
| 684 |
+
|
| 685 |
+
if (enabled and isinstance(first, SplitLinear) and isinstance(second, SplitLinear) and first.activation
|
| 686 |
+
and first.out == ttnn.float32 and second.fuses_input_act()):
|
| 687 |
+
return second(first(x, defer_act=True), pre_act=first.activation)
|
| 688 |
+
return second(first(x))
|
| 689 |
|
| 690 |
|
| 691 |
def make_linear(build: Build, lin: Lin, module: str, *, out: str, activation: Optional[str] = None):
|
|
|
|
| 703 |
def __init__(self, build: Build, nrm: Optional[Norm], module: str):
|
| 704 |
self.cfg = build.cfg(module)
|
| 705 |
self.mode = build.ln_mode(module)
|
| 706 |
+
self.fused = build.ln_kernel
|
| 707 |
+
self.resid = build.ln_resid
|
| 708 |
+
self.split = build.ln_split
|
| 709 |
+
self.sbc = build.ln_sfpu_bcast
|
| 710 |
+
self.tr = build.ln_tr
|
| 711 |
+
self.mem = None # DEC_L1: memory config of the split-row LN's outputs
|
| 712 |
self.g = self.b = None
|
| 713 |
if nrm is not None:
|
| 714 |
self.g = build.upload(np.asarray(nrm.gamma).reshape(1, -1), "float32")
|
| 715 |
self.b = build.upload(np.asarray(nrm.beta).reshape(1, -1), "float32")
|
| 716 |
|
| 717 |
+
def __call__(self, x, gamma=None, beta=None, *, kcat: int = 0):
|
| 718 |
+
"""``kcat`` (``KCAT_EMIT``, the consumer's :func:`operand_ktp`): return the consumer's split operand as a
|
| 719 |
+
:class:`KcatOperand` when the split-row kernel runs (else y as usual)."""
|
| 720 |
import ttnn
|
| 721 |
|
| 722 |
g = self.g if gamma is None else gamma
|
| 723 |
b = self.b if beta is None else beta
|
| 724 |
if self.mode == "fp32":
|
| 725 |
+
if self.fused:
|
| 726 |
+
from .ln_kernel import layer_norm_fp32_fused, supported
|
| 727 |
+
|
| 728 |
+
if self.split:
|
| 729 |
+
from .ln_kernel import layer_norm_fp32_split, split_supported
|
| 730 |
+
|
| 731 |
+
if split_supported(x):
|
| 732 |
+
y = layer_norm_fp32_split(x, g, b, eps=C.LN_EPS, sfpu_bcast=self.sbc, kcat_ktp=kcat,
|
| 733 |
+
memory_config=self.mem)
|
| 734 |
+
return KcatOperand(y, kcat) if kcat else y
|
| 735 |
+
if supported(x):
|
| 736 |
+
return layer_norm_fp32_fused(x, g, b, eps=C.LN_EPS, lean=self.fused == 2, sfpu_bcast=self.sbc,
|
| 737 |
+
memory_config=self.mem)
|
| 738 |
return layer_norm_fp32(x, g, b)
|
| 739 |
+
kw = {} if self.mem is None else {"memory_config": self.mem}
|
| 740 |
+
return ttnn.layer_norm(x, epsilon=C.LN_EPS, weight=g, bias=b, compute_kernel_config=self.cfg, **kw)
|
| 741 |
+
|
| 742 |
+
def can_transpose(self, x) -> bool:
|
| 743 |
+
"""``LN_TR``: the fused one-core-per-row kernel runs for ``x`` (not the split-row form), so it can read the
|
| 744 |
+
residual / write the output transposed per entity."""
|
| 745 |
+
if not (self.tr and self.mode == "fp32" and self.fused and self.resid):
|
| 746 |
+
return False
|
| 747 |
+
from .ln_kernel import split_supported, supported
|
| 748 |
+
|
| 749 |
+
return supported(x) and not (self.split and split_supported(x))
|
| 750 |
+
|
| 751 |
+
def transposed(self, x, *, res=None):
|
| 752 |
+
"""``LN_TR``: ``LN(x)`` (or, with ``res``, ``h = x + res`` and ``LN(h)``) with the output written as
|
| 753 |
+
``[.., E, W, T]`` (the per-entity transpose); returns ``y^T`` or ``(h, y^T)``."""
|
| 754 |
+
from .ln_kernel import layer_norm_fp32_fused
|
| 755 |
+
|
| 756 |
+
return layer_norm_fp32_fused(x, self.g, self.b, eps=C.LN_EPS, residual=res, lean=self.fused == 2,
|
| 757 |
+
sfpu_bcast=self.sbc, out_t=True, memory_config=self.mem)
|
| 758 |
+
|
| 759 |
+
def residual_t(self, x, res_t_tensor):
|
| 760 |
+
"""``LN_TR``: ``h = x + res^T`` (``res`` ``[.., E, W, T]``) and ``LN(h)`` -> ``(h, y)``."""
|
| 761 |
+
from .ln_kernel import layer_norm_fp32_fused
|
| 762 |
+
|
| 763 |
+
return layer_norm_fp32_fused(x, self.g, self.b, eps=C.LN_EPS, residual=res_t_tensor, lean=self.fused == 2,
|
| 764 |
+
sfpu_bcast=self.sbc, res_t=True, memory_config=self.mem)
|
| 765 |
+
|
| 766 |
+
def residual(self, x, res, rgate=None, gamma=None, beta=None, *, write_h: bool = True, kcat: int = 0):
|
| 767 |
+
"""``h = x + res (* rgate)`` then ``(h, LN(h))``: one fused program (``LN_RESID``, ``tt/ln_kernel.py``)
|
| 768 |
+
when this LN runs the fused fp32 kernel, else the stock ``add`` (+ ``multiply``) and :meth:`__call__`.
|
| 769 |
+
``h`` is None when ``write_h`` is False and the fused program runs."""
|
| 770 |
+
import ttnn
|
| 771 |
+
|
| 772 |
+
g = self.g if gamma is None else gamma
|
| 773 |
+
b = self.b if beta is None else beta
|
| 774 |
+
if self.mode == "fp32" and self.fused and self.resid:
|
| 775 |
+
from .ln_kernel import layer_norm_fp32_fused, supported
|
| 776 |
+
|
| 777 |
+
if supported(x) and supported(res) and tuple(x.shape) == tuple(res.shape):
|
| 778 |
+
if self.split:
|
| 779 |
+
from .ln_kernel import layer_norm_fp32_split, split_supported
|
| 780 |
+
|
| 781 |
+
if split_supported(x):
|
| 782 |
+
h, y = layer_norm_fp32_split(x, g, b, eps=C.LN_EPS, residual=res, rgate=rgate,
|
| 783 |
+
write_h=write_h, sfpu_bcast=self.sbc, kcat_ktp=kcat,
|
| 784 |
+
memory_config=self.mem)
|
| 785 |
+
return h, (KcatOperand(y, kcat) if kcat else y)
|
| 786 |
+
return layer_norm_fp32_fused(x, g, b, eps=C.LN_EPS, residual=res, rgate=rgate, write_h=write_h,
|
| 787 |
+
lean=self.fused == 2, sfpu_bcast=self.sbc, memory_config=self.mem)
|
| 788 |
+
h = ttnn.add(x, res if rgate is None else ttnn.multiply(res, rgate))
|
| 789 |
+
return h, self(h, gamma, beta)
|
| 790 |
|
| 791 |
|
| 792 |
class Const:
|
| 793 |
"""A constant device tensor (upload once)."""
|
| 794 |
|
| 795 |
+
def __init__(self, build: Build, array: np.ndarray, dtype: str, memory_config=None):
|
| 796 |
self.t = build.upload(array, dtype)
|
| 797 |
+
if memory_config is not None:
|
| 798 |
+
import ttnn
|
| 799 |
+
|
| 800 |
+
self.t = ttnn.to_memory_config(self.t, memory_config)
|
| 801 |
|
| 802 |
def __call__(self):
|
| 803 |
return self.t
|
code/tt_diffusion_planner/tt/ln_kernel.py
ADDED
|
@@ -0,0 +1,285 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""Fused fp32 LayerNorm as one ``ttnn.generic_op`` program (``LN_KERNEL``, OPT round 2 item 3).
|
| 3 |
+
|
| 4 |
+
:func:`layer_norm_fp32_fused` computes exactly what :func:`tt.layers.layer_norm_fp32` computes with 7-9 stock
|
| 5 |
+
programs (mean, subtract, square, mean, add eps, rsqrt, multiply, gamma, beta), in one program: the compute kernel
|
| 6 |
+
(``kernels/ln32_compute.cpp``) issues the same SFPU LLK calls in the same order as the stock kernels, and every
|
| 7 |
+
intermediate the stock graph writes to an fp32 DRAM tensor stays fp32 in L1 / DST, so the output is meant to be
|
| 8 |
+
bit-identical (checked on the device: ``code/scripts/ln_kernel_check.py``).
|
| 9 |
+
|
| 10 |
+
Layout: x is an fp32 TILE DRAM-interleaved tensor ``[..., R, W]`` (R a multiple of 32, W = 32 * Wt); the tile rows
|
| 11 |
+
are split in contiguous blocks over ``min(rows, grid)`` cores (one core per tile row up to the grid). gamma / beta
|
| 12 |
+
are ``[1, 1, 1, W]`` fp32 TILE rows (row 0 valid) or None.
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import os
|
| 17 |
+
import struct
|
| 18 |
+
from typing import Any, Optional
|
| 19 |
+
|
| 20 |
+
__all__ = ["layer_norm_fp32_fused", "supported"]
|
| 21 |
+
|
| 22 |
+
_KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
|
| 23 |
+
TB = 4096 # fp32 tile bytes
|
| 24 |
+
N_RT = 2 # per-core RT args of every kernel ([row0, n_rows] / [n_rows] + pad): a CT arg (probe P3)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def _bits(v: float) -> int:
|
| 28 |
+
return int.from_bytes(struct.pack("<f", float(v)), "little")
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def supported(x: Any) -> bool:
|
| 32 |
+
"""True when ``x`` is a device fp32 TILE interleaved tensor with tile-aligned rows and columns."""
|
| 33 |
+
import ttnn
|
| 34 |
+
|
| 35 |
+
try:
|
| 36 |
+
shp = list(x.padded_shape)
|
| 37 |
+
return (x.dtype == ttnn.float32 and x.layout == ttnn.TILE_LAYOUT and not x.is_sharded()
|
| 38 |
+
and shp[-1] % 32 == 0 and shp[-2] % 32 == 0 and int(x.shape[-1]) == shp[-1]
|
| 39 |
+
and hasattr(ttnn, "generic_op"))
|
| 40 |
+
except Exception: # noqa: BLE001 - the host fake ttnn
|
| 41 |
+
return False
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def _cores(n: int, g):
|
| 45 |
+
"""The first ``n`` cores of the grid in row-major order: a CoreRangeSet and the coordinate list."""
|
| 46 |
+
import ttnn
|
| 47 |
+
|
| 48 |
+
cs = [(i % g.x, i // g.x) for i in range(n)]
|
| 49 |
+
full, rem = divmod(n, g.x)
|
| 50 |
+
rs = []
|
| 51 |
+
if full:
|
| 52 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(g.x - 1, full - 1)))
|
| 53 |
+
if rem:
|
| 54 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, full), ttnn.CoreCoord(rem - 1, full)))
|
| 55 |
+
return ttnn.CoreRangeSet(set(rs)), cs
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def layer_norm_fp32_fused(x, gamma=None, beta=None, *, eps: float, residual=None, rgate=None, write_h: bool = True,
|
| 59 |
+
lean: bool = True, sfpu_bcast: bool = False, res_t: bool = False, out_t: bool = False,
|
| 60 |
+
memory_config=None):
|
| 61 |
+
"""LayerNorm over the last dim of fp32 ``x`` -> fp32 tensor of x's shape (see the module docstring).
|
| 62 |
+
|
| 63 |
+
With ``residual`` (``LN_RESID``): ``h = x + residual (* rgate)`` first (the stock ``ttnn.add(x,
|
| 64 |
+
ttnn.multiply(residual, rgate))``, fp32 SFPU ops), then ``LN(h)``; returns ``(h, LN(h))`` (``h`` is None when
|
| 65 |
+
``write_h`` is False). ``residual``: fp32 like ``x``; ``rgate``: a ``[1, 1, 1, W]`` fp32 row.
|
| 66 |
+
``lean``: one unpacker / SFPU-binary init per phase instead of one per tile (same LLK math calls).
|
| 67 |
+
``sfpu_bcast`` (``LN_SFPU_BCAST``): the SFPU row reduce writes the row statistics to every column
|
| 68 |
+
(``kernels/ln32_sfpu.h``) instead of the writer's RISC-V column fill (same values).
|
| 69 |
+
``res_t`` / ``out_t`` (``LN_TR``, x ``[.., E, T, W]``): the residual comes as ``[.., E, W, T]`` (its per-entity
|
| 70 |
+
transpose) / the output is written as ``[.., E, W, T]``, the tiles transposed in the kernel with the stock
|
| 71 |
+
``ttnn.transpose`` LLK (exact): the two transposes around the mixer's token-mixing MLP. ``memory_config``: of
|
| 72 |
+
the outputs (default DRAM; interleaved L1 for ``ENC_L1``)."""
|
| 73 |
+
import ttnn
|
| 74 |
+
|
| 75 |
+
dev = x.device()
|
| 76 |
+
shp = list(x.padded_shape)
|
| 77 |
+
W = shp[-1]
|
| 78 |
+
Wt = W // 32
|
| 79 |
+
rows = 1
|
| 80 |
+
for d in shp[:-1]:
|
| 81 |
+
rows *= d
|
| 82 |
+
rows //= 32
|
| 83 |
+
hr = int(residual is not None)
|
| 84 |
+
hrg = int(hr and rgate is not None)
|
| 85 |
+
wh = int(hr and write_h)
|
| 86 |
+
Tt = shp[-2] // 32
|
| 87 |
+
assert not (res_t and hrg), "LN_TR: no residual gate"
|
| 88 |
+
oshape = x.shape if not out_t else ttnn.Shape(list(x.shape)[:-2] + [shp[-1], shp[-2]])
|
| 89 |
+
omem = memory_config or ttnn.DRAM_MEMORY_CONFIG
|
| 90 |
+
out = ttnn.allocate_tensor_on_device(oshape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem)
|
| 91 |
+
h = (ttnn.allocate_tensor_on_device(x.shape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem)
|
| 92 |
+
if wh else None)
|
| 93 |
+
g = dev.compute_with_storage_grid_size()
|
| 94 |
+
n = min(rows, g.x * g.y)
|
| 95 |
+
crs, cs = _cores(n, g)
|
| 96 |
+
base, extra = divmod(rows, n)
|
| 97 |
+
rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
|
| 98 |
+
r0 = 0
|
| 99 |
+
for i, (cx, cy) in enumerate(cs):
|
| 100 |
+
k = base + (1 if i < extra else 0)
|
| 101 |
+
rd[cx][cy] = [r0, k]
|
| 102 |
+
wr[cx][cy] = [r0, k]
|
| 103 |
+
cp[cx][cy] = [k, 0]
|
| 104 |
+
r0 += k
|
| 105 |
+
hg, hb = int(gamma is not None), int(beta is not None)
|
| 106 |
+
|
| 107 |
+
def acc(t):
|
| 108 |
+
return list(ttnn.TensorAccessorArgs(t).get_compile_time_args())
|
| 109 |
+
|
| 110 |
+
def cb(idx, pages):
|
| 111 |
+
return ttnn.CBDescriptor(total_size=pages * TB, core_ranges=crs, format_descriptors=[
|
| 112 |
+
ttnn.CBFormatDescriptor(buffer_index=idx, data_format=ttnn.float32, page_size=TB)])
|
| 113 |
+
|
| 114 |
+
cbs = [cb(0, 2 * Wt), cb(3, 1), cb(4, 2), cb(5, 2), cb(6, Wt), cb(16, 2 * Wt if Wt <= 4 else Wt)]
|
| 115 |
+
if hg:
|
| 116 |
+
cbs.append(cb(1, Wt))
|
| 117 |
+
if hb:
|
| 118 |
+
cbs.append(cb(2, Wt))
|
| 119 |
+
if hr:
|
| 120 |
+
cbs += [cb(7, Wt), cb(9, Wt)]
|
| 121 |
+
if hrg:
|
| 122 |
+
cbs.append(cb(8, Wt))
|
| 123 |
+
if wh:
|
| 124 |
+
cbs.append(cb(17, 2))
|
| 125 |
+
if out_t:
|
| 126 |
+
cbs.append(cb(18, 1))
|
| 127 |
+
um = [ttnn.UnpackToDestMode.Default] * 64
|
| 128 |
+
for i in (0, 1, 2, 3, 5, 6, 7, 8, 9, 18):
|
| 129 |
+
um[i] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 130 |
+
ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True,
|
| 131 |
+
math_approx_mode=False)
|
| 132 |
+
ccfg.unpack_to_dest_mode = um
|
| 133 |
+
g_t = gamma if hg else x
|
| 134 |
+
b_t = beta if hb else x
|
| 135 |
+
r_t = residual if hr else x
|
| 136 |
+
rg_t = rgate if hrg else x
|
| 137 |
+
h_t = h if wh else out
|
| 138 |
+
reader_ct = [Wt, hg, hb, _bits(eps), N_RT, hr, hrg, int(bool(res_t)), Tt] + acc(x) + acc(g_t) + acc(b_t) + acc(
|
| 139 |
+
r_t) + acc(rg_t)
|
| 140 |
+
writer_ct = [Wt, N_RT, wh, int(bool(sfpu_bcast)), int(bool(out_t)), Tt] + acc(out) + acc(h_t)
|
| 141 |
+
compute_ct = [Wt, hg, hb, _bits(1.0 / W), N_RT, hr, hrg, wh, int(bool(lean)), int(bool(sfpu_bcast)),
|
| 142 |
+
int(bool(res_t)), int(bool(out_t))]
|
| 143 |
+
SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH
|
| 144 |
+
|
| 145 |
+
def kd(name, ct, rt, common, config):
|
| 146 |
+
return ttnn.KernelDescriptor(kernel_source=os.path.join(_KDIR, name), source_type=SRC, core_ranges=crs,
|
| 147 |
+
compile_time_args=ct, defines=[], runtime_args=rt, common_runtime_args=common,
|
| 148 |
+
config=config, compiler_include_paths=[_KDIR])
|
| 149 |
+
|
| 150 |
+
ks = [kd("ln32_reader.cpp", reader_ct, rd, [x.buffer_address(), g_t.buffer_address(), b_t.buffer_address(),
|
| 151 |
+
r_t.buffer_address(), rg_t.buffer_address()],
|
| 152 |
+
ttnn.ReaderConfigDescriptor()),
|
| 153 |
+
kd("ln32_writer.cpp", writer_ct, wr, [out.buffer_address(), h_t.buffer_address()],
|
| 154 |
+
ttnn.WriterConfigDescriptor()),
|
| 155 |
+
kd("ln32_compute.cpp", compute_ct, cp, [], ccfg)]
|
| 156 |
+
ins = [x] + ([gamma] if hg else []) + ([beta] if hb else []) + ([residual] if hr else []) + (
|
| 157 |
+
[rgate] if hrg else []) + ([h] if wh else [])
|
| 158 |
+
ttnn.generic_op(ins + [out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))
|
| 159 |
+
if hr:
|
| 160 |
+
return h, out
|
| 161 |
+
return out
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
def reference_decomposition(x, gamma: Optional[Any] = None, beta: Optional[Any] = None, *, eps: float):
|
| 165 |
+
"""The stock decomposition (``tt.layers.layer_norm_fp32``), for the device check."""
|
| 166 |
+
from .layers import layer_norm_fp32
|
| 167 |
+
|
| 168 |
+
return layer_norm_fp32(x, gamma, beta, eps=eps)
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
def split_supported(x: Any) -> bool:
|
| 172 |
+
"""The split-row form needs every tile row's Wt column tiles on distinct cores (rows * Wt <= grid cores)."""
|
| 173 |
+
if not supported(x):
|
| 174 |
+
return False
|
| 175 |
+
shp = list(x.padded_shape)
|
| 176 |
+
rows = 1
|
| 177 |
+
for d in shp[:-1]:
|
| 178 |
+
rows *= d
|
| 179 |
+
rows //= 32
|
| 180 |
+
g = x.device().compute_with_storage_grid_size()
|
| 181 |
+
wt = shp[-1] // 32
|
| 182 |
+
return 2 <= wt <= g.x * g.y and rows * wt <= g.x * g.y
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
def layer_norm_fp32_split(x, gamma=None, beta=None, *, eps: float, residual=None, rgate=None, write_h: bool = True,
|
| 186 |
+
sfpu_bcast: bool = False, kcat_ktp: int = 0, memory_config=None):
|
| 187 |
+
""":func:`layer_norm_fp32_fused` with each tile row spread over Wt cores (``LN_SPLIT``, ``kernels/ln32s_*.cpp``):
|
| 188 |
+
member j owns column tile j; the root (member 0) folds the gathered tiles in order and broadcasts the mean and
|
| 189 |
+
rstd, so the result is the same bit for bit. For few rows (the decoder's 11 tile rows: 88 cores instead of 11).
|
| 190 |
+
``kcat_ktp`` (``KCAT_EMIT``): instead of y, write the split operand ``[y_hi | y_hi | y_lo | 1 | 0..]`` of the next
|
| 191 |
+
K-concatenated linear (``[..., R, 32 * kcat_ktp]``, the ``kcat_operand`` layout and LLK calls).
|
| 192 |
+
``memory_config``: of the outputs (default DRAM; interleaved L1 for ``DEC_L1``)."""
|
| 193 |
+
import ttnn
|
| 194 |
+
|
| 195 |
+
dev = x.device()
|
| 196 |
+
shp = list(x.padded_shape)
|
| 197 |
+
W = shp[-1]
|
| 198 |
+
Wt = W // 32
|
| 199 |
+
omem = memory_config or ttnn.DRAM_MEMORY_CONFIG
|
| 200 |
+
rows = 1
|
| 201 |
+
for d in shp[:-1]:
|
| 202 |
+
rows *= d
|
| 203 |
+
rows //= 32
|
| 204 |
+
hg, hb = int(gamma is not None), int(beta is not None)
|
| 205 |
+
hr = int(residual is not None)
|
| 206 |
+
hrg = int(hr and rgate is not None)
|
| 207 |
+
wh = int(hr and write_h)
|
| 208 |
+
oshape = x.shape if not kcat_ktp else ttnn.Shape(list(x.shape)[:-1] + [32 * kcat_ktp])
|
| 209 |
+
out = ttnn.allocate_tensor_on_device(oshape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem)
|
| 210 |
+
h = (ttnn.allocate_tensor_on_device(x.shape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem)
|
| 211 |
+
if wh else None)
|
| 212 |
+
g = dev.compute_with_storage_grid_size()
|
| 213 |
+
G = min(rows, (g.x * g.y) // Wt)
|
| 214 |
+
n = G * Wt
|
| 215 |
+
crs, cs = _cores(n, g)
|
| 216 |
+
phys = [dev.worker_core_from_logical_core(ttnn.CoreCoord(cx, cy)) for cx, cy in cs]
|
| 217 |
+
rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
|
| 218 |
+
n_rt = 6 + 2 * Wt
|
| 219 |
+
for gi in range(G):
|
| 220 |
+
n_rows = len(range(gi, rows, G))
|
| 221 |
+
root = phys[gi * Wt]
|
| 222 |
+
members = []
|
| 223 |
+
for m in range(Wt):
|
| 224 |
+
members += [phys[gi * Wt + m].x, phys[gi * Wt + m].y]
|
| 225 |
+
for j in range(Wt):
|
| 226 |
+
cx, cy = cs[gi * Wt + j]
|
| 227 |
+
args = [gi, n_rows, G, j, root.x, root.y] + members
|
| 228 |
+
rd[cx][cy] = args
|
| 229 |
+
wr[cx][cy] = args
|
| 230 |
+
cp[cx][cy] = [n_rows, int(j == 0)]
|
| 231 |
+
|
| 232 |
+
def acc(t):
|
| 233 |
+
return list(ttnn.TensorAccessorArgs(t).get_compile_time_args())
|
| 234 |
+
|
| 235 |
+
def cb(idx, pages):
|
| 236 |
+
return ttnn.CBDescriptor(total_size=pages * TB, core_ranges=crs, format_descriptors=[
|
| 237 |
+
ttnn.CBFormatDescriptor(buffer_index=idx, data_format=ttnn.float32, page_size=TB)])
|
| 238 |
+
|
| 239 |
+
cbs = [cb(0, 2), cb(3, 1), cb(4, 2), cb(5, 2), cb(6, 2), cb(10, Wt), cb(11, 2), cb(12, 1), cb(16, 2)]
|
| 240 |
+
if hg:
|
| 241 |
+
cbs.append(cb(1, 1))
|
| 242 |
+
if hb:
|
| 243 |
+
cbs.append(cb(2, 1))
|
| 244 |
+
if hr:
|
| 245 |
+
cbs += [cb(7, 2), cb(9, 2)]
|
| 246 |
+
if hrg:
|
| 247 |
+
cbs.append(cb(8, 1))
|
| 248 |
+
if wh:
|
| 249 |
+
cbs.append(cb(17, 2))
|
| 250 |
+
if kcat_ktp:
|
| 251 |
+
cbs += [cb(18, 1), cb(19, 2), cb(20, 1)] + ([cb(21, 1)] if kcat_ktp > 3 * Wt + 1 else [])
|
| 252 |
+
um = [ttnn.UnpackToDestMode.Default] * 64
|
| 253 |
+
for i in (0, 1, 2, 3, 5, 6, 7, 8, 9, 10, 18):
|
| 254 |
+
um[i] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 255 |
+
ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True,
|
| 256 |
+
math_approx_mode=False)
|
| 257 |
+
ccfg.unpack_to_dest_mode = um
|
| 258 |
+
g_t = gamma if hg else x
|
| 259 |
+
b_t = beta if hb else x
|
| 260 |
+
r_t = residual if hr else x
|
| 261 |
+
rg_t = rgate if hrg else x
|
| 262 |
+
h_t = h if wh else out
|
| 263 |
+
reader_ct = [Wt, hg, hb, _bits(eps), n_rt, hr, hrg] + acc(x) + acc(g_t) + acc(b_t) + acc(r_t) + acc(rg_t)
|
| 264 |
+
writer_ct = [Wt, n_rt, wh, int(bool(sfpu_bcast)), int(bool(kcat_ktp)), int(kcat_ktp)] + acc(out) + acc(h_t)
|
| 265 |
+
compute_ct = [Wt, hg, hb, _bits(1.0 / W), 2, hr, hrg, wh, int(bool(sfpu_bcast)), int(bool(kcat_ktp))]
|
| 266 |
+
SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH
|
| 267 |
+
|
| 268 |
+
def kd(name, ct, rt, common, config):
|
| 269 |
+
return ttnn.KernelDescriptor(kernel_source=os.path.join(_KDIR, name), source_type=SRC, core_ranges=crs,
|
| 270 |
+
compile_time_args=ct, defines=[], runtime_args=rt, common_runtime_args=common,
|
| 271 |
+
config=config, compiler_include_paths=[_KDIR])
|
| 272 |
+
|
| 273 |
+
ks = [kd("ln32s_reader.cpp", reader_ct, rd, [x.buffer_address(), g_t.buffer_address(), b_t.buffer_address(),
|
| 274 |
+
r_t.buffer_address(), rg_t.buffer_address()],
|
| 275 |
+
ttnn.ReaderConfigDescriptor()),
|
| 276 |
+
kd("ln32s_writer.cpp", writer_ct, wr, [out.buffer_address(), h_t.buffer_address()],
|
| 277 |
+
ttnn.WriterConfigDescriptor()),
|
| 278 |
+
kd("ln32s_compute.cpp", compute_ct, cp, [], ccfg)]
|
| 279 |
+
sems = [ttnn.SemaphoreDescriptor(id=i, core_ranges=crs, initial_value=0) for i in range(2)]
|
| 280 |
+
ins = [x] + ([gamma] if hg else []) + ([beta] if hb else []) + ([residual] if hr else []) + (
|
| 281 |
+
[rgate] if hrg else []) + ([h] if wh else [])
|
| 282 |
+
ttnn.generic_op(ins + [out], ttnn.ProgramDescriptor(kernels=ks, semaphores=sems, cbs=cbs))
|
| 283 |
+
if hr:
|
| 284 |
+
return h, out
|
| 285 |
+
return out
|
code/tt_diffusion_planner/tt/model.py
CHANGED
|
@@ -33,7 +33,39 @@ from .decoder import STEP_KEYS, TtDecoder, TtTurnHead
|
|
| 33 |
from .encoder import TtEncoder
|
| 34 |
from .layers import Build, policy
|
| 35 |
|
| 36 |
-
__all__ = ["TtDiffusionPlanner"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
|
| 39 |
class TtDiffusionPlanner:
|
|
@@ -57,15 +89,24 @@ class TtDiffusionPlanner:
|
|
| 57 |
attn_fp32_acc=attn_fp32_acc, attn_matmul=attn_matmul)
|
| 58 |
p = weights.params
|
| 59 |
self.tables = P.step_tables(p, steps)
|
| 60 |
-
self.
|
| 61 |
-
self.
|
|
|
|
| 62 |
self.turn = TtTurnHead(self.build, p)
|
| 63 |
self.runner = TraceRunner(device, num_command_queues=num_command_queues, name="diffusion-planner")
|
| 64 |
warm = I.warmup_inputs()
|
| 65 |
for name, shape in I.INPUT_SPECS.items():
|
| 66 |
dtype = "bfloat16" if name in I.BF16_INPUTS else "float32"
|
| 67 |
self.runner.add_input(name, init=warm[name], dtype=dtype, layout=ttnn.TILE_LAYOUT)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 68 |
self.runner.add_variant("plan", self._plan)
|
|
|
|
|
|
|
| 69 |
self.debug = bool(debug)
|
| 70 |
if self.debug:
|
| 71 |
self.runner.add_input("dbg_x", init=np.zeros((1, 1, T.AGENTS, T.STATE_COLS), np.float32),
|
|
@@ -80,11 +121,22 @@ class TtDiffusionPlanner:
|
|
| 80 |
self.build_ms = (time.perf_counter() - t0) * 1e3
|
| 81 |
|
| 82 |
# ---- traced functions --------------------------------------------------------------------------------------
|
| 83 |
-
def _plan(self, ctx):
|
| 84 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 85 |
kv = self.decoder.cross_kv(enc)
|
| 86 |
-
|
| 87 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 88 |
logit = self.turn(final, enc)
|
| 89 |
return pack_outputs({"final_x0": final, "logit": logit, "ego_steps": ego})
|
| 90 |
|
|
@@ -108,15 +160,54 @@ class TtDiffusionPlanner:
|
|
| 108 |
def capture(self) -> None:
|
| 109 |
self.runner.capture()
|
| 110 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 111 |
def _run(self, variant: str, inputs: Dict[str, Any], eager: bool):
|
|
|
|
| 112 |
return self.runner.run_eager(variant, inputs=inputs) if eager else self.runner(variant, inputs=inputs)
|
| 113 |
|
| 114 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
"""One plan: upload, replay, read (``eager=True``: the same graph without the trace, for bring-up and
|
| 116 |
replay-vs-eager checks). ``final_x0`` ``[321, 81, 4]`` (normalised, prefix-constrained), ``logit`` ``[5]``,
|
| 117 |
-
``denoising_steps``: the 11 iterates' ego rows as ``[1, 81, 4]`` arrays.
|
| 118 |
-
|
| 119 |
-
|
|
|
|
| 120 |
res = {"final_x0": final.reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM).astype(np.float32),
|
| 121 |
"logit": out["logit"].reshape(-1)[:C.TURN_INDICATOR_OUTPUT_DIM].astype(np.float32)}
|
| 122 |
if ego_steps:
|
|
@@ -180,7 +271,7 @@ class TtDiffusionPlanner:
|
|
| 180 |
"precision": self.build.policy.describe(),
|
| 181 |
"options": self.build.options(),
|
| 182 |
"uploaded_mb": round(self.build.uploaded_bytes / 2 ** 20, 2), "build_ms": round(self.build_ms, 1),
|
| 183 |
-
"tokens": T.TOKENS, "agents": T.AGENTS, "nfe": self.tables.nfe}
|
| 184 |
|
| 185 |
def release(self) -> None:
|
| 186 |
self.runner.release()
|
|
|
|
| 33 |
from .encoder import TtEncoder
|
| 34 |
from .layers import Build, policy
|
| 35 |
|
| 36 |
+
__all__ = ["TtDiffusionPlanner", "agent_buckets", "needed_rows", "bucket_rows"]
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def agent_buckets() -> tuple:
|
| 40 |
+
"""``COMPACT``: the agent buckets of ``AGENT_BUCKETS`` (multiples of 32 below 352, sorted); () when off."""
|
| 41 |
+
k = T.KNOBS.read()
|
| 42 |
+
if not k.COMPACT:
|
| 43 |
+
return ()
|
| 44 |
+
out = sorted({int(v) for v in str(k.AGENT_BUCKETS).replace(" ", "").split(",") if v})
|
| 45 |
+
bad = [v for v in out if v % T.TILE or not 0 < v < T.AGENTS]
|
| 46 |
+
if bad:
|
| 47 |
+
raise ValueError(f"DIFFUSION_PLANNER_AGENT_BUCKETS: {bad} are not multiples of {T.TILE} below {T.AGENTS}")
|
| 48 |
+
return tuple(out)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def needed_rows(prepared: Any) -> int:
|
| 52 |
+
"""1 + the last decoder row a plan reads or attends to: the ego (row 0), the valid self-attention keys
|
| 53 |
+
(``agent_valid``) and the emitted neighbours (``neighbor_rows`` + 1). The rows past it are masked keys whose
|
| 54 |
+
outputs are never read, so dropping them leaves the needed rows' values unchanged."""
|
| 55 |
+
valid = np.flatnonzero(np.asarray(prepared.decoder.agent_valid, bool))
|
| 56 |
+
last = int(valid.max()) if valid.size else 0
|
| 57 |
+
tok = np.flatnonzero(np.asarray(prepared.features.valid["neighbor"], bool)) # valid neighbour tokens
|
| 58 |
+
if tok.size:
|
| 59 |
+
last = max(last, int(tok.max()) + 1)
|
| 60 |
+
emitted = np.asarray(prepared.neighbor_rows)
|
| 61 |
+
if emitted.size:
|
| 62 |
+
last = max(last, int(emitted.max()) + 1)
|
| 63 |
+
return last + 1
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def bucket_rows(prepared: Any, buckets: Sequence[int]) -> int:
|
| 67 |
+
n = needed_rows(prepared)
|
| 68 |
+
return next((int(b) for b in sorted(buckets) if b >= n), T.AGENTS)
|
| 69 |
|
| 70 |
|
| 71 |
class TtDiffusionPlanner:
|
|
|
|
| 89 |
attn_fp32_acc=attn_fp32_acc, attn_matmul=attn_matmul)
|
| 90 |
p = weights.params
|
| 91 |
self.tables = P.step_tables(p, steps)
|
| 92 |
+
self.buckets = agent_buckets() if hasattr(device, "compute_with_storage_grid_size") else ()
|
| 93 |
+
self.encoder = TtEncoder(self.build, p, nb_rows=self.buckets if self.build.compact_enc else ())
|
| 94 |
+
self.decoder = TtDecoder(self.build, p, self.tables, rows=(T.AGENTS,) + self.buckets)
|
| 95 |
self.turn = TtTurnHead(self.build, p)
|
| 96 |
self.runner = TraceRunner(device, num_command_queues=num_command_queues, name="diffusion-planner")
|
| 97 |
warm = I.warmup_inputs()
|
| 98 |
for name, shape in I.INPUT_SPECS.items():
|
| 99 |
dtype = "bfloat16" if name in I.BF16_INPUTS else "float32"
|
| 100 |
self.runner.add_input(name, init=warm[name], dtype=dtype, layout=ttnn.TILE_LAYOUT)
|
| 101 |
+
self.cs_pad = {} # INPUT_TRIM: zero columns that widen cs to STATE_COLS, per row count
|
| 102 |
+
if I.CS_COLS < T.STATE_COLS:
|
| 103 |
+
for r in (T.AGENTS,) + self.buckets:
|
| 104 |
+
self.cs_pad[r] = self.build.upload(np.zeros((r, T.STATE_COLS - I.CS_COLS), np.float32), "float32")
|
| 105 |
+
self.trim = I.CS_COLS < T.STATE_COLS
|
| 106 |
+
self._zero_inputs: set = set() # INPUT_TRIM: inputs whose device buffer holds +0.0 (our last upload)
|
| 107 |
self.runner.add_variant("plan", self._plan)
|
| 108 |
+
for r in self.buckets: # COMPACT: one trace per agent bucket (the decoder on r rows)
|
| 109 |
+
self.runner.add_variant(f"plan_r{r}", lambda ctx, r=r: self._plan(ctx, r))
|
| 110 |
self.debug = bool(debug)
|
| 111 |
if self.debug:
|
| 112 |
self.runner.add_input("dbg_x", init=np.zeros((1, 1, T.AGENTS, T.STATE_COLS), np.float32),
|
|
|
|
| 121 |
self.build_ms = (time.perf_counter() - t0) * 1e3
|
| 122 |
|
| 123 |
# ---- traced functions --------------------------------------------------------------------------------------
|
| 124 |
+
def _plan(self, ctx, rows: int = T.AGENTS):
|
| 125 |
+
"""The whole plan; ``rows`` < 352 (``COMPACT``): the decoder on the first ``rows`` agents (the inputs'
|
| 126 |
+
leading rows, sliced in the trace)."""
|
| 127 |
+
import ttnn
|
| 128 |
+
|
| 129 |
+
enc = self.encoder.forward(ctx, nb=rows if rows != T.AGENTS else None)
|
| 130 |
kv = self.decoder.cross_kv(enc)
|
| 131 |
+
y0, cs, key = ctx["y0"], ctx["cs"], ctx["agent_key_row"]
|
| 132 |
+
if rows != T.AGENTS:
|
| 133 |
+
y0 = ttnn.slice(y0, [0, 0, 0, 0], [1, 1, rows, T.STATE_COLS])
|
| 134 |
+
cs = ttnn.slice(cs, [0, 0, 0, 0], [1, 1, rows, int(cs.shape[-1])])
|
| 135 |
+
key = ttnn.slice(key, [0, 0, 0, 0], [1, 1, 1, rows])
|
| 136 |
+
if int(cs.shape[-1]) < T.STATE_COLS: # INPUT_TRIM: widen the one-tile cs (tile-aligned concat)
|
| 137 |
+
cs = ttnn.concat([cs, self.cs_pad[rows]], dim=-1)
|
| 138 |
+
self_mask = A.expand_key_bias(key, rows, memory_config=self.build.attn_mem())
|
| 139 |
+
final, ego = self.decoder.solve(y0, cs, kv, self_mask)
|
| 140 |
logit = self.turn(final, enc)
|
| 141 |
return pack_outputs({"final_x0": final, "logit": logit, "ego_steps": ego})
|
| 142 |
|
|
|
|
| 160 |
def capture(self) -> None:
|
| 161 |
self.runner.capture()
|
| 162 |
|
| 163 |
+
def filter_inputs(self, inputs: Dict[str, Any]) -> Dict[str, Any]:
|
| 164 |
+
"""``INPUT_TRIM``: drop the inputs that are all +0.0 while their device buffer already holds the zeros of this
|
| 165 |
+
model's previous upload (a value test, not a cache of the previous request: any non-zero array is always
|
| 166 |
+
uploaded). Every input passed on is recorded as uploaded."""
|
| 167 |
+
if not self.trim:
|
| 168 |
+
return inputs
|
| 169 |
+
out = {}
|
| 170 |
+
for name, a in inputs.items():
|
| 171 |
+
arr = a if isinstance(a, np.ndarray) else None
|
| 172 |
+
zero = arr is not None and arr.dtype == np.float32 and not arr.view(np.uint32).any()
|
| 173 |
+
if zero and name in self._zero_inputs:
|
| 174 |
+
continue
|
| 175 |
+
out[name] = a
|
| 176 |
+
if zero:
|
| 177 |
+
self._zero_inputs.add(name)
|
| 178 |
+
else:
|
| 179 |
+
self._zero_inputs.discard(name)
|
| 180 |
+
return out
|
| 181 |
+
|
| 182 |
def _run(self, variant: str, inputs: Dict[str, Any], eager: bool):
|
| 183 |
+
inputs = self.filter_inputs(inputs)
|
| 184 |
return self.runner.run_eager(variant, inputs=inputs) if eager else self.runner(variant, inputs=inputs)
|
| 185 |
|
| 186 |
+
def rows_for(self, prepared: Any) -> int:
|
| 187 |
+
"""The decoder rows a plan needs (``COMPACT``): the smallest agent bucket holding the ego, every valid
|
| 188 |
+
self-attention key and every emitted neighbour row; 352 without buckets or when none is large enough."""
|
| 189 |
+
return bucket_rows(prepared, self.buckets)
|
| 190 |
+
|
| 191 |
+
def variant_for(self, prepared: Any) -> str:
|
| 192 |
+
r = self.rows_for(prepared)
|
| 193 |
+
return "plan" if r == T.AGENTS else f"plan_r{r}"
|
| 194 |
+
|
| 195 |
+
def unpack(self, out: Dict[str, Any]) -> np.ndarray:
|
| 196 |
+
"""The readback's ``final_x0`` (``[R, 324]``) -> ``[321, 324]``; the rows past R (not computed: masked keys,
|
| 197 |
+
never emitted) are zero."""
|
| 198 |
+
f = np.asarray(out["final_x0"], np.float32).reshape(-1, T.STATE_COLS)
|
| 199 |
+
if f.shape[0] < T.AGENTS:
|
| 200 |
+
f = np.concatenate([f, np.zeros((T.AGENTS - f.shape[0], T.STATE_COLS), np.float32)], 0)
|
| 201 |
+
return f[:C.MAX_NUM_AGENTS]
|
| 202 |
+
|
| 203 |
+
def forward(self, prepared: Any, *, ego_steps: bool = True, eager: bool = False,
|
| 204 |
+
variant: Optional[str] = None) -> Dict[str, Any]:
|
| 205 |
"""One plan: upload, replay, read (``eager=True``: the same graph without the trace, for bring-up and
|
| 206 |
replay-vs-eager checks). ``final_x0`` ``[321, 81, 4]`` (normalised, prefix-constrained), ``logit`` ``[5]``,
|
| 207 |
+
``denoising_steps``: the 11 iterates' ego rows as ``[1, 81, 4]`` arrays. ``variant``: force a plan variant
|
| 208 |
+
(default :meth:`variant_for`)."""
|
| 209 |
+
out = self._run(variant or self.variant_for(prepared), I.plan_inputs(prepared), eager)
|
| 210 |
+
final = self.unpack(out)
|
| 211 |
res = {"final_x0": final.reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM).astype(np.float32),
|
| 212 |
"logit": out["logit"].reshape(-1)[:C.TURN_INDICATOR_OUTPUT_DIM].astype(np.float32)}
|
| 213 |
if ego_steps:
|
|
|
|
| 271 |
"precision": self.build.policy.describe(),
|
| 272 |
"options": self.build.options(),
|
| 273 |
"uploaded_mb": round(self.build.uploaded_bytes / 2 ** 20, 2), "build_ms": round(self.build_ms, 1),
|
| 274 |
+
"tokens": T.TOKENS, "agents": T.AGENTS, "agent_buckets": list(self.buckets), "nfe": self.tables.nfe}
|
| 275 |
|
| 276 |
def release(self) -> None:
|
| 277 |
self.runner.release()
|
code/tt_diffusion_planner/tt/smask_kernel.py
ADDED
|
@@ -0,0 +1,98 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""Attention score scale + additive mask as one ``ttnn.generic_op`` (``ATTN_SMASK``, OPT round 2 item 4b).
|
| 3 |
+
|
| 4 |
+
``scale_mask(s, scale, mask)``: fp32 TILE scores ``[1, H, Sq, Sk]`` and a ``[1, 1, Sq, Sk]`` mask (bf16 or fp32,
|
| 5 |
+
broadcast over the heads) -> ``s * scale + mask`` fp32, the two stock ``binary_ng`` programs of
|
| 6 |
+
``tt/attention.py`` (``ttnn.multiply(s, scale)``, ``ttnn.add(., mask)``) in one pass over the scores
|
| 7 |
+
(``kernels/smask_*.cpp``): bit-identical, half the DRAM traffic and one program less per attention. A bf16 mask
|
| 8 |
+
sends the stock fp32 + bf16 add down binary_ng's FPU path, which truncates the scores to TF32 (measured); the
|
| 9 |
+
kernel reproduces that truncation (the mask holds 0 / -inf only, so the add itself is exact).
|
| 10 |
+
The mask-plane tiles are split in contiguous ranges over the grid; each core streams the H score tiles of its mask
|
| 11 |
+
tiles.
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import os
|
| 16 |
+
import struct
|
| 17 |
+
from typing import Any
|
| 18 |
+
|
| 19 |
+
__all__ = ["scale_mask", "supported"]
|
| 20 |
+
|
| 21 |
+
_KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
|
| 22 |
+
TB = 4096
|
| 23 |
+
N_RT = 2
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def supported(s: Any, mask: Any) -> bool:
|
| 27 |
+
import ttnn
|
| 28 |
+
|
| 29 |
+
try:
|
| 30 |
+
ss, ms = list(s.padded_shape), list(mask.padded_shape)
|
| 31 |
+
return (hasattr(ttnn, "generic_op") and s.dtype == ttnn.float32 and s.layout == ttnn.TILE_LAYOUT
|
| 32 |
+
and mask.layout == ttnn.TILE_LAYOUT and mask.dtype in (ttnn.bfloat16, ttnn.float32)
|
| 33 |
+
and not s.is_sharded() and not mask.is_sharded() and len(ss) == 4 and len(ms) == 4
|
| 34 |
+
and ss[0] == 1 and ss[1] % 2 == 0 and ms[0] == 1 and ms[1] == 1 and ss[2:] == ms[2:])
|
| 35 |
+
except Exception: # noqa: BLE001 - the host fake ttnn
|
| 36 |
+
return False
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def scale_mask(s: Any, scale: float, mask: Any):
|
| 40 |
+
import ttnn
|
| 41 |
+
|
| 42 |
+
dev = s.device()
|
| 43 |
+
_, H, Sq, Sk = list(s.padded_shape)
|
| 44 |
+
plane = (Sq // 32) * (Sk // 32)
|
| 45 |
+
out = ttnn.allocate_tensor_on_device(s.shape, ttnn.float32, ttnn.TILE_LAYOUT, dev, ttnn.DRAM_MEMORY_CONFIG)
|
| 46 |
+
g = dev.compute_with_storage_grid_size()
|
| 47 |
+
n = min(plane, g.x * g.y)
|
| 48 |
+
cs = [(i % g.x, i // g.x) for i in range(n)]
|
| 49 |
+
full, rem = divmod(n, g.x)
|
| 50 |
+
rs = []
|
| 51 |
+
if full:
|
| 52 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(g.x - 1, full - 1)))
|
| 53 |
+
if rem:
|
| 54 |
+
rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, full), ttnn.CoreCoord(rem - 1, full)))
|
| 55 |
+
crs = ttnn.CoreRangeSet(set(rs))
|
| 56 |
+
base, extra = divmod(plane, n)
|
| 57 |
+
rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
|
| 58 |
+
m0 = 0
|
| 59 |
+
for i, (cx, cy) in enumerate(cs):
|
| 60 |
+
k = base + (1 if i < extra else 0)
|
| 61 |
+
rd[cx][cy] = [m0, k]
|
| 62 |
+
wr[cx][cy] = [m0, k]
|
| 63 |
+
cp[cx][cy] = [k, 0]
|
| 64 |
+
m0 += k
|
| 65 |
+
mbf = mask.dtype == ttnn.bfloat16
|
| 66 |
+
tbm = 2048 if mbf else 4096
|
| 67 |
+
|
| 68 |
+
def cb(idx, pages, dt, tb):
|
| 69 |
+
return ttnn.CBDescriptor(total_size=pages * tb, core_ranges=crs, format_descriptors=[
|
| 70 |
+
ttnn.CBFormatDescriptor(buffer_index=idx, data_format=dt, page_size=tb)])
|
| 71 |
+
|
| 72 |
+
cbs = [cb(0, 2 * H, ttnn.float32, TB), cb(1, 2, mask.dtype, tbm), cb(2, 1, ttnn.float32, TB),
|
| 73 |
+
cb(16, 4, ttnn.float32, TB)]
|
| 74 |
+
um = [ttnn.UnpackToDestMode.Default] * 64
|
| 75 |
+
um[0] = um[2] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 76 |
+
if not mbf:
|
| 77 |
+
um[1] = ttnn.UnpackToDestMode.UnpackToDestFp32
|
| 78 |
+
ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True,
|
| 79 |
+
math_approx_mode=False)
|
| 80 |
+
ccfg.unpack_to_dest_mode = um
|
| 81 |
+
bits = int.from_bytes(struct.pack("<f", float(scale)), "little")
|
| 82 |
+
SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH
|
| 83 |
+
|
| 84 |
+
def acc(t):
|
| 85 |
+
return list(ttnn.TensorAccessorArgs(t).get_compile_time_args())
|
| 86 |
+
|
| 87 |
+
def kd(name, ct, rt, common, config):
|
| 88 |
+
return ttnn.KernelDescriptor(kernel_source=os.path.join(_KDIR, name), source_type=SRC, core_ranges=crs,
|
| 89 |
+
compile_time_args=ct, defines=[], runtime_args=rt, common_runtime_args=common,
|
| 90 |
+
config=config)
|
| 91 |
+
|
| 92 |
+
ks = [kd("smask_reader.cpp", [H, plane, bits, N_RT, tbm] + acc(s) + acc(mask), rd,
|
| 93 |
+
[s.buffer_address(), mask.buffer_address()], ttnn.ReaderConfigDescriptor()),
|
| 94 |
+
kd("smask_writer.cpp", [H, plane, N_RT] + acc(out), wr, [out.buffer_address()],
|
| 95 |
+
ttnn.WriterConfigDescriptor()),
|
| 96 |
+
kd("smask_compute.cpp", [H, N_RT, int(mbf)], cp, [], ccfg)]
|
| 97 |
+
ttnn.generic_op([s, mask, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))
|
| 98 |
+
return out
|