changh95 commited on
Commit
be62f78
·
verified ·
1 Parent(s): 460a357

tt-model push diffusion-planner-p150 (container)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +13 -0
  2. OPT_BASELINE.md +26 -0
  3. OPT_REPORT.md +579 -30
  4. PYTHON.md +20 -19
  5. README.md +46 -44
  6. SERVING.md +20 -18
  7. VERIFICATION_OPT_2026-10-11.md +157 -0
  8. build_info.json +6 -6
  9. code/PYTHON.md +20 -19
  10. code/scripts/bench.py +24 -9
  11. code/scripts/profile_ops.py +7 -4
  12. code/tt_diffusion_planner/__init__.py +1 -1
  13. code/tt_diffusion_planner/host/features.py +37 -10
  14. code/tt_diffusion_planner/host/normalize.py +63 -1
  15. code/tt_diffusion_planner/host/pipeline.py +10 -3
  16. code/tt_diffusion_planner/host/postprocess.py +68 -1
  17. code/tt_diffusion_planner/tests/test_bundle_host.py +5 -1
  18. code/tt_diffusion_planner/tests/test_grid_fit_host.py +23 -0
  19. code/tt_diffusion_planner/tests/test_tt_params_host.py +28 -0
  20. code/tt_diffusion_planner/tt/attention.py +96 -0
  21. code/tt_diffusion_planner/tt/config.py +102 -0
  22. code/tt_diffusion_planner/tt/decoder.py +85 -10
  23. code/tt_diffusion_planner/tt/encoder.py +94 -13
  24. code/tt_diffusion_planner/tt/fattn_kernel.py +140 -0
  25. code/tt_diffusion_planner/tt/inputs.py +6 -3
  26. code/tt_diffusion_planner/tt/kcat_kernel.py +109 -0
  27. code/tt_diffusion_planner/tt/kernels/README.md +30 -8
  28. code/tt_diffusion_planner/tt/kernels/fattn_compute.cpp +222 -0
  29. code/tt_diffusion_planner/tt/kernels/fattn_reader.cpp +95 -0
  30. code/tt_diffusion_planner/tt/kernels/fattn_writer.cpp +82 -0
  31. code/tt_diffusion_planner/tt/kernels/kcat_compute.cpp +74 -0
  32. code/tt_diffusion_planner/tt/kernels/kcat_reader.cpp +24 -0
  33. code/tt_diffusion_planner/tt/kernels/kcat_writer.cpp +66 -0
  34. code/tt_diffusion_planner/tt/kernels/ln32_compute.cpp +282 -0
  35. code/tt_diffusion_planner/tt/kernels/ln32_reader.cpp +123 -0
  36. code/tt_diffusion_planner/tt/kernels/ln32_sfpu.h +54 -0
  37. code/tt_diffusion_planner/tt/kernels/ln32_writer.cpp +84 -0
  38. code/tt_diffusion_planner/tt/kernels/ln32s_compute.cpp +214 -0
  39. code/tt_diffusion_planner/tt/kernels/ln32s_reader.cpp +130 -0
  40. code/tt_diffusion_planner/tt/kernels/ln32s_writer.cpp +161 -0
  41. code/tt_diffusion_planner/tt/kernels/smask_compute.cpp +69 -0
  42. code/tt_diffusion_planner/tt/kernels/smask_reader.cpp +50 -0
  43. code/tt_diffusion_planner/tt/kernels/smask_writer.cpp +31 -0
  44. code/tt_diffusion_planner/tt/kernels/smsm_compute.cpp +179 -0
  45. code/tt_diffusion_planner/tt/kernels/smsm_reader.cpp +83 -0
  46. code/tt_diffusion_planner/tt/kernels/smsm_writer.cpp +43 -0
  47. code/tt_diffusion_planner/tt/layers.py +597 -15
  48. code/tt_diffusion_planner/tt/ln_kernel.py +285 -0
  49. code/tt_diffusion_planner/tt/model.py +103 -12
  50. 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: baseline port; optimization pending.** This is the first public release: the functional port, measured
4
- once (`OPT_BASELINE.md`, baseline commit `5541833`, 2026-10-08) and not optimized yet. No optimization round has run,
5
- so there is no step, no rejected attempt and no hang to report, and every number below is the baseline. The whole
6
- plan (encoder, the 11 DiT evaluations with the DPM-Solver++(2M) updates, the turn-indicator head) runs as one metal
7
- trace of 6,282 programs: 102.04 ms per back-to-back replay (9.80 plans/s of device throughput), 117.9 ms per
8
- synchronous `model()` call on the shipped sample. The trace is kernel-bound (op-to-op gaps 4.4 ms of a 103.7 ms span).
9
- Two sinks of similar size come first: the **numerics defaults** that the end-to-end gates need (split hi / lo matmuls,
10
- fp32 LayerNorm, fp32 matmul attention) cost **33.4 ms** over the first device round, and the encoder's channel-MLP and
11
- pre-projection matmuls run on only 4-8 cores (**25.8 ms**, unrelated to precision).
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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) | **final (= baseline: no round yet)** |
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 | same |
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) | same |
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
- Not started. The optimization phase follows the publication of every model of the collection (`research/PLAN.md`
53
- §4.1, §5): each step is one commit with one `DIFFUSION_PLANNER_*` A/B knob, keeps the frozen gates and is re-checked on
54
- all 99 end-to-end scenes (the plan's sensitivity to device numerics is chaotic per scene: PORT_LOG known issues). The
55
- precision policy (PLAN.md §9.3) may be relaxed per module only with gate evidence; the precision cost is to be
56
- recovered with fused kernels that keep fp32-level accuracy, not by dropping precision.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- None yet. The candidates (D19) are in the backlog (item 6): the persistent DiT-evaluation kernel first.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 or the release docs | – |
86
 
87
  ## Rejected / not kept
88
 
89
- None yet. Configurations that were measured and are **not** the default, for accuracy (PORT_LOG decision 11,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 trace. 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,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 one metal trace (`warmup_variants="default"`: the variant `plan`) 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-09): the load takes 315 s with an empty JIT cache and 8.6 s with a warm one (build 0.53 s: the ONNX initializers read and 48.1 MB of weights and constants uploaded; warm-up and capture 3.7 s; the rest is the device open). The first call then takes 120 ms and the second 120 ms (the stage bench's steady state: 118 ms p50 for decoded arrays, 125 ms for an `.npz` path). The trace holds 74.6 MB of DRAM (`trace_region_size` 192 MiB).
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 trace 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,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; the numbers of `OPT_BASELINE.md`, 2026-10-08, on a shared host; p50, with p99 in brackets):
156
 
157
- | stage | shipped sample `kashiwanoha_dense` |
158
  |---|---:|
159
- | `.npz` decode + schema check (path inputs only) | 7.02 (18.20) ms |
160
- | host pre-processing (the node's normalization, the encoder's host features, masks, `x_T`) | 4.75 (12.55) ms |
161
- | packing the 19 persistent trace inputs · ttnn host tensors · H2D | 0.45 · 2.41 · 1.14 ms |
162
- | **device trace, one blocking plan** | **102.13** (104.81) ms |
163
- | D2H (one packed read, 460 KB) | 0.58 ms |
164
- | host post-processing (trajectory, predicted paths, turn decision) | 4.32 (11.54) ms |
165
- | **`model(inputs=arrays)` end to end** | **117.90** (134.56) ms |
166
- | `model(inputs=<.npz path>)` | 124.76 (147.74) ms |
167
- | back-to-back replays (device time per plan) | 102.04 ms = 9.80 plans/s |
168
-
169
- The device time does not depend on the scene: every plan computes the full capacities (re-checked on c0d84f9: kashiwanoha_dense 102.10, straight_road 102.11, a nuScenes instant 102.09 ms). Throughput above one plan per ~118 ms needs pipelining of the host work of neighbouring requests (not implemented); 2 CQs do 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 trace; every plan computes
176
- the full capacities, so the device time does not depend on the scene.
 
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
- - tt-nn
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, fp32 matmul attention in the fusion encoder and the decoder: the configuration the end-to-end accuracy gates need (Caveats). 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. 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 trace. The first load compiles the kernels (315 s with an empty JIT cache, firmware and every kernel compiled); later loads take about 9 s (device open, reading and uploading the weights, warm-up and trace capture).
60
- - The trace is captured during the load, so the first call is as fast as the later calls and no call compiles anything.
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": 3.8, "device": 105.3, "postprocess": 3.2, "total": 122.8, "decode": 8.5, "model_call": 112.5},
99
  "num_poses": 80,
100
  "columns": ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"],
101
- "trajectory": [[0.365, 0.01, -0.017, 0.9962, -0.0169, 3.8059, -0.0431], [0.7631, -0.0022, -0.0414, 0.9966, -0.0413, 3.8016, -0.5485], [1.1524, -0.0282, -0.0675, 0.9942, -0.0674, 3.7468, -0.4708], "... 77 more rows"],
102
- "turn_indicator": {"command": 1, "command_name": "DISABLE", "keep_selected": true, "held": false, "logits": [-15.960084915161133, -5.125061511993408, -5.230083465576172, -0.7658059597015381, 5.504596710205078], "probabilities": [1.651601411190029e-09, 8.38487030705437e-05, 7.548934809165075e-05, 0.006556871347129345, 0.993184506893158]},
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; the stage bench 2026-10-08, the served rows, load times and the re-check 2026-10-09. Latency: the stage bench of [`OPT_BASELINE.md`](OPT_BASELINE.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; the device time is the same for every scene: every plan computes the full capacities); the served rows 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 several ms; the device rows repeat to ±0.05 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.
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.3 cm** / mean 1.1 cm; turn command identical; neighbours: median per-agent max 8.6 cm |
140
- | Agreement, shipped sample `straight_road` | ego max 4.9 cm / mean 1.8 cm; turn command identical; neighbours 6.3 cm |
141
- | Agreement, all 99 gated scenes (2 samples, 5 research scenes, 92 nuScenes-mini instants) | worst ego max **0.313 m** / mean **0.143 m** (gates 1.0 / 0.3 m; nuscenes/scene-0103_kf14); ego mean: median 1.3 cm, 95th percentile 6.2 cm; turn command identical **99 / 99**; neighbours: median per-agent max ≤ 0.086 m (gate 1.5 m) |
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.313 m / mean 0.143 m, turn command identical 92 / 92; sequence: worst ego max 0.118 m / mean 0.049 m, turn 33 / 33; every instant within the gates |
143
- | Module PCC vs the fp32 reference (encoder categories, encoding, teacher-forced decoder evaluation; gate 0.999) | ≥ 0.999952 / 0.999978 / ≥ 0.9999995 |
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) | **117.9 ms p50** (p99 134.6) · 8.5 plans/s |
146
- | Served `/predict` `timing_ms.total` (uvicorn on the host, the shipped sample) | **123.1 ms median** (min 122.5; of which decode 8.6) |
147
- | Served client round trip, loopback (base64 `.npz` request, 0.15 MB) | 127.5 ms median |
148
- | Device trace, one blocking plan (encoder + 11 DiT evaluations + 10 solver updates + turn head) | **102.13 ms** |
149
- | Back-to-back trace replays | **102.04 ms per plan** · 9.80 plans/s |
150
- | Host pre-processing · pack · host tensors · H2D · D2H · host post-processing | 4.75 · 0.45 · 2.41 · 1.14 · 0.58 · 4.32 ms |
151
- | `from_pretrained` load: empty JIT cache / warm cache | 315 s / 8.6 s |
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), [`OPT_BASELINE.md`](OPT_BASELINE.md), [`OPT_REPORT.md`](OPT_REPORT.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
- - First release: **baseline port, optimization pending.** The plan is one metal trace and kernel-bound (6,282 programs; op-to-op gaps 4.4 ms of a 103.7 ms span). The **first optimization target is the precision cost**: the numerics defaults that the end-to-end gates need (split hi / lo matmuls, fp32 LayerNorm, fp32 matmul attention) cost 33.4 ms per plan (102.0 ms vs 68.6 ms with the first device round's defaults, which fail the gates); fused kernels are to recover it without dropping precision. A second sink of similar size is unrelated to precision: the encoder's channel-MLP and pre-projection matmuls run on 4-8 cores (25.8 ms). [`OPT_REPORT.md`](OPT_REPORT.md) ranks what comes next.
 
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 decomposition 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; 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.
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` (container smoke, served p150 output vs the stored CPU reference) the per-agent max displacement is 8.6 cm median and 0.26 m at the 90th percentile; the worst of the 88 agents differs by 4.31 m. The ego plan is gated on its max and mean; neighbour paths only on that median.
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 (101.87 ms per replay, the same as ETH within 0.2 %); 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.
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 trace, and every plan computes them in full, so the device time does not depend on the scene.
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` modifies tt-metal (Apache-2.0).
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; dirty tree: the image includes the patch) |
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.20.0, vendored as `code/tt_diffusion_planner/ttaw` from the Autoware ports' shared `common` repository at commit `89dec49` (`code/tt_diffusion_planner/ttaw/VENDORED.json`: version, commit and per-file sha256) |
192
- | `code/` digest (image) | `c0e7abb7888098a9` (sha256, first 16 hex digits; `built.code_sha256` of `tt_kernel_manifest.json`) |
193
- | image | `tt-model/diffusion-planner-p150:3b96d8ea7190` (`sha256:3b96d8ea71902fe6a00f1792dd41290839b7758070431464cd613cde3f6bf909`) |
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-09T04:37:19+00:00 by tt-model 0.1.0 |
 
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.20.0 (`common` @ `89dec49`) |
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 plan's trace holds 74.6 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; 6,282 device programs |
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 (315 s measured for
49
- `from_pretrained`); later boots take about 10 s to READY (measured on the host, warm JIT cache). Startup failures raise and uvicorn exits non-zero (no CPU
50
- fallback). SIGTERM: the lifespan releases the trace 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,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": 3.8, "device": 105.3, "postprocess": 3.2, "total": 122.8, "decode": 8.5, "model_call": 112.5},
148
  "num_poses": 80,
149
  "columns": ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"],
150
- "trajectory": [[0.365, 0.01, -0.017, 0.9962, -0.0169, 3.8059, -0.0431], [0.7631, -0.0022, -0.0414, 0.9966, -0.0413, 3.8016, -0.5485], [1.1524, -0.0282, -0.0675, 0.9942, -0.0674, 3.7468, -0.4708], "... 77 more rows"],
151
- "turn_indicator": {"command": 1, "command_name": "DISABLE", "keep_selected": true, "held": false, "logits": [-15.960084915161133, -5.125061511993408, -5.230083465576172, -0.7658059597015381, 5.504596710205078], "probabilities": [1.651601411190029e-09, 8.38487030705437e-05, 7.548934809165075e-05, 0.006556871347129345, 0.993184506893158]},
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 replay ~10 % slower on ETH, OPT_BASELINE.md) |
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 19 persistent device inputs ->
201
- `execute_trace` of the whole plan -> one D2H -> the node's host post-processing) under one lock -> `Output.to_dict()`.
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 trace; every plan computes
232
- the full capacities whatever the scene holds, so the device time does not depend on the scene.
 
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 trace inside the lifespan, so READY means warm; the first cold boot pays the ttnn JIT.
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-09T04:48:03Z",
5
  "image": {
6
- "tag": "tt-model/diffusion-planner-p150:3b96d8ea7190",
7
- "digest": "sha256:3b96d8ea71902fe6a00f1792dd41290839b7758070431464cd613cde3f6bf909",
8
- "built_at": "2026-10-09T04:37:19+00:00",
9
- "code_sha256": "c0e7abb7888098a9319ab5c66a10c4fd4009fc1542326b26507554c04f05a481",
10
- "size_bytes": 4069025450,
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 trace. 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,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 one metal trace (`warmup_variants="default"`: the variant `plan`) 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-09): the load takes 315 s with an empty JIT cache and 8.6 s with a warm one (build 0.53 s: the ONNX initializers read and 48.1 MB of weights and constants uploaded; warm-up and capture 3.7 s; the rest is the device open). The first call then takes 120 ms and the second 120 ms (the stage bench's steady state: 118 ms p50 for decoded arrays, 125 ms for an `.npz` path). The trace holds 74.6 MB of DRAM (`trace_region_size` 192 MiB).
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 trace 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,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; the numbers of `OPT_BASELINE.md`, 2026-10-08, on a shared host; p50, with p99 in brackets):
156
 
157
- | stage | shipped sample `kashiwanoha_dense` |
158
  |---|---:|
159
- | `.npz` decode + schema check (path inputs only) | 7.02 (18.20) ms |
160
- | host pre-processing (the node's normalization, the encoder's host features, masks, `x_T`) | 4.75 (12.55) ms |
161
- | packing the 19 persistent trace inputs · ttnn host tensors · H2D | 0.45 · 2.41 · 1.14 ms |
162
- | **device trace, one blocking plan** | **102.13** (104.81) ms |
163
- | D2H (one packed read, 460 KB) | 0.58 ms |
164
- | host post-processing (trajectory, predicted paths, turn decision) | 4.32 (11.54) ms |
165
- | **`model(inputs=arrays)` end to end** | **117.90** (134.56) ms |
166
- | `model(inputs=<.npz path>)` | 124.76 (147.74) ms |
167
- | back-to-back replays (device time per plan) | 102.04 ms = 9.80 plans/s |
168
-
169
- The device time does not depend on the scene: every plan computes the full capacities (re-checked on c0d84f9: kashiwanoha_dense 102.10, straight_road 102.11, a nuScenes instant 102.09 ms). Throughput above one plan per ~118 ms needs pipelining of the host work of neighbouring requests (not implemented); 2 CQs do 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 trace; every plan computes
176
- the full capacities, so the device time does not depend on the scene.
 
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 = out["final_x0"].reshape(T.AGENTS, T.STATE_COLS)[:C.MAX_NUM_AGENTS]
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("plan")
132
  sync()
133
  with bench.stage("d2h"):
134
- out = runner.read("plan")
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("plan"), sync, n=a.b2b_iters, warmup=3)
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
- inputs = I.plan_inputs(hp.prepare(raw, model.normalization.observation))
 
 
 
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("plan", inputs=inputs)
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("plan", n=a.replays)
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("plan")
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.1.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
- nb = np.zeros_like(nb_raw)
133
- nb[:, C.NEIGHBOR_HISTORY_KEEP] = nb_raw[:, C.NEIGHBOR_HISTORY_KEEP]
134
- nb_type = nb[:, -1, 8:11].copy()
135
- x8 = nb[..., :8]
136
- step_valid = _any_nonzero(x8, -1) # [320, 31]
137
- nb_valid = step_valid.any(axis=-1) # [320]
138
- feat = np.concatenate([x8, step_valid[..., None].astype(np.float32)], axis=-1)
139
- feat[..., 4:6] = 0.0 # velocities zeroed after the validity test
140
- feat[~nb_valid] = 0.0
141
- pos["neighbor"] = _onehot_pos(x8[:, -1, :4], C.POS_CLASS["neighbor"])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 normalize_array(value: np.ndarray, mean: np.ndarray, std: np.ndarray) -> np.ndarray:
 
 
 
 
 
 
 
 
 
 
 
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
- denorm = denormalize(x0, mean, std) # [321, 80, 4]
 
 
 
 
 
 
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
- rows = prepared.neighbor_rows
94
- paths = predicted_paths(denorm, C.MAX_NUM_NEIGHBORS)[rows] if rows.size else np.zeros(
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
- def trajectory_from_poses(poses_xycs: np.ndarray, base_position: Tuple[float, float, float] = (0.0, 0.0, 0.0), *,
 
 
 
 
 
 
 
 
 
 
 
 
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 .layers import ATTN, Build, Const, LayerNorm, make_linear
 
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.agent_rows = Const(build, P.agent_rows(p, T.AGENTS), self.stream)
 
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, (T.AGENTS, T.TOKENS)), "bfloat16")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- return [(A.split_heads(k(enc), C.NUM_HEADS), A.split_heads(v(enc), C.NUM_HEADS))
114
- for k, v in zip(self.k_lin, self.v_lin)]
 
 
 
 
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(A.attention_matmul(q, k, v, scale=self.scale, attn_mask=mask))
 
 
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, 352, 324]`` fp32 (prefix-constrained) -> ``m * mask0`` fp32."""
 
124
  import ttnn
125
 
126
- h = ttnn.add(self.pre2(self.pre1(x)), self.agent_rows()) # [1, 1, 352, 256]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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, self.E * C.MIXER_CHANNELS, x.shape[-1]))
93
 
94
  def _entities_view(self, x):
95
  import ttnn
96
 
97
- return ttnn.reshape(x, (1, self.E, C.MIXER_CHANNELS, x.shape[-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, self.E, C.MIXER_CHANNELS))
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 forward(self, ctx: Mapping[str, Any], taps: Optional[Dict[str, Any]] = None):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- x0 = tap(f"enc.{cat}.pre", tr.pre(ctx[f"{cat}_x"]))
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) # [1, 1, 576, 576], once per plan
 
253
  for i, b in enumerate(self.blocks):
254
  kv = b["kv"](x) # K | V from the un-normalised x
255
- q, k, v = A.split_q_kv(b["q"](b["n1"](x)), kv, C.NUM_HEADS)
256
- if self.attn_mm:
257
- a = A.merge_heads(A.attention_matmul(q, k, v, scale=self.scale, attn_mask=mask))
 
 
 
 
 
 
258
  else:
 
259
  a = A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32)
260
- x = ttnn.add(x, b["out"](a))
261
- x = ttnn.add(x, b["fc2"](b["fc1"](b["n2"](x))))
 
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, T.STATE_COLS),
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
- None for the functional port: PLAN.md section 2.12 maps every op of the v5.0 graphs to stock ttnn (matmul / linear,
4
- layer_norm, gelu, SDPA through the shared C20 wrapper, eltwise). Fused kernels (entity-parallel MLP-Mixer, fusion
5
- block, DiT step + solver update) and the single-megakernel attempt (decision D19) are optimization-phase work; when
6
- they land, their `.cpp` sources go here (package data in the repo-root `pyproject.toml`, asserted by a `verify:` line of
7
- `tt-model.yaml`), each with a numpy / torch oracle test and the hang protocol of PLAN.md section 4.4.
8
-
9
- | file | op | used by | status |
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
- return ttnn.linear(x, self.w, bias=self.b, activation=self.activation, dtype=self.out,
118
- compute_kernel_config=self.cfg)
 
 
 
 
 
 
 
 
 
 
 
 
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
- def __call__(self, x):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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, compute_kernel_config=self.cfg)
168
- y = ttnn.add(y, ttnn.linear(x, self.w_lo, bias=self.b, dtype=f32, compute_kernel_config=self.cfg))
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, compute_kernel_config=self.cfg)
173
- y = ttnn.add(y, ttnn.matmul(x_hi, self.w_lo, dtype=f32, compute_kernel_config=self.cfg))
174
- y = ttnn.add(y, ttnn.linear(x_lo, self.w_hi, bias=self.b, dtype=f32, compute_kernel_config=self.cfg))
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
- return ttnn.layer_norm(x, epsilon=C.LN_EPS, weight=g, bias=b, compute_kernel_config=self.cfg)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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.encoder = TtEncoder(self.build, p)
61
- self.decoder = TtDecoder(self.build, p, self.tables)
 
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
- enc = self.encoder.forward(ctx)
 
 
 
 
85
  kv = self.decoder.cross_kv(enc)
86
- self_mask = A.expand_key_bias(ctx["agent_key_row"], T.AGENTS)
87
- final, ego = self.decoder.solve(ctx["y0"], ctx["cs"], kv, self_mask)
 
 
 
 
 
 
 
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 forward(self, prepared: Any, *, ego_steps: bool = True, eager: bool = False) -> Dict[str, Any]:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- out = self._run("plan", I.plan_inputs(prepared), eager)
119
- final = out["final_x0"].reshape(T.AGENTS, T.STATE_COLS)[:C.MAX_NUM_AGENTS]
 
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