raahemnabeel commited on
Commit
88e0a5c
·
verified ·
1 Parent(s): 767698c

Add files using upload-large-folder tool

Browse files
Files changed (49) hide show
  1. .gitattributes +14 -0
  2. code/models/demos/blackhole/qwen3_coder_30b_a3b/README.md +143 -0
  3. code/models/demos/blackhole/qwen3_coder_30b_a3b/__init__.py +0 -0
  4. code/models/demos/blackhole/qwen3_coder_30b_a3b/config/context_contract.json +360 -0
  5. code/models/demos/blackhole/qwen3_coder_30b_a3b/config/selected_precision_config.json +22 -0
  6. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/__init__.py +0 -0
  7. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/reference.py +145 -0
  8. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_attention.py +117 -0
  9. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_attention_decode.py +140 -0
  10. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decode_compaction_fifo.py +596 -0
  11. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decoder_layer.py +176 -0
  12. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decoder_layer_decode.py +155 -0
  13. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_determinism.py +167 -0
  14. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_full_model.py +976 -0
  15. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_moe.py +232 -0
  16. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_multichip_decoder.py +1057 -0
  17. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_optimized_decoder.py +503 -0
  18. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_perf.py +611 -0
  19. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_precision_config.py +369 -0
  20. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_reference.py +123 -0
  21. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_rmsnorm.py +99 -0
  22. code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_trace.py +200 -0
  23. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt-model-localplugin.yaml +152 -0
  24. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt-model.yaml +242 -0
  25. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/__init__.py +0 -0
  26. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/functional_decoder.py +1169 -0
  27. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator.py +1637 -0
  28. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator_vllm.py +1568 -0
  29. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/model.py +1723 -0
  30. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/multichip_decoder.py +1982 -0
  31. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/optimized_decoder.py +1093 -0
  32. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/precision.py +378 -0
  33. code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/weight_mapping.py +184 -0
  34. code/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle/qwen3_coder_30b_a3b_instruct/tt_qwen3_coder_30b_a3b_instruct.py +35 -0
  35. code/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle/qwen3_coder_30b_a3b_instruct/vllm_metadata.json +8 -0
  36. image/blobs/sha256/0926a8eb0e608a5c6888d1cd5594184bdf3ed3aa311dba5b42a547caefdc6f2e +3 -0
  37. image/blobs/sha256/24cba7375920bef8d4cc4f0ce4294f8f70c65b9a37c4ed6f6c2d63405c76ba3c +3 -0
  38. image/blobs/sha256/3de1f5eb93e54b4561afa733cf844d1914a9ba260f49e86fe600f344c5dd025c +3 -0
  39. image/blobs/sha256/530b0e35f44c6f963e06fdaacdbfecb2021d4c55d49fe4c0019e161d94c18de3 +3 -0
  40. image/blobs/sha256/540cf00275e913a9bccc49fe7beba58037696661bde6ae083a3ac84c5b160e67 +3 -0
  41. image/blobs/sha256/8753e0cfbd424e962ccaf50aaaf02fd06ff2efb8677219657e751a53922efa9f +3 -0
  42. image/blobs/sha256/b6df468b82a4b2f9ee3ca3a79a6bbe99b5cda02ddc63b7e3c89e7bd08ef41706 +3 -0
  43. image/blobs/sha256/c00314a02c644cf2aea5a8ae4e3ee5ee383f4072d3e35248fe20dea57a90c46c +3 -0
  44. image/blobs/sha256/c18d0f3c8022bcd5a8059f67fc0e5cfd53ef36a44997335d6cd0aa6b19db140d +3 -0
  45. image/blobs/sha256/c3bc7b373b4523cabcdd9c64ab10d32510ef61bace60873a279ccc4902738989 +3 -0
  46. image/blobs/sha256/ca0b072b65f8c21199f96e7498f2ccbc252490ce39060800647332751f857287 +3 -0
  47. image/blobs/sha256/cbb77a738c7df827819a8b8bf87682eeca8bdb434d41411a9682c38717f2f187 +3 -0
  48. image/blobs/sha256/e19b6f1fb65dd2888d9003ef9513a21d129041a328bb8a9a4164d29ef0382b16 +3 -0
  49. image/blobs/sha256/fdc1ed79ffd24d66f8be3754ec4dc80b1ab0fcc8c8165a6007748c464c94f897 +3 -0
.gitattributes CHANGED
@@ -33,3 +33,17 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ image/blobs/sha256/3de1f5eb93e54b4561afa733cf844d1914a9ba260f49e86fe600f344c5dd025c filter=lfs diff=lfs merge=lfs -text
37
+ image/blobs/sha256/ca0b072b65f8c21199f96e7498f2ccbc252490ce39060800647332751f857287 filter=lfs diff=lfs merge=lfs -text
38
+ image/blobs/sha256/cbb77a738c7df827819a8b8bf87682eeca8bdb434d41411a9682c38717f2f187 filter=lfs diff=lfs merge=lfs -text
39
+ image/blobs/sha256/530b0e35f44c6f963e06fdaacdbfecb2021d4c55d49fe4c0019e161d94c18de3 filter=lfs diff=lfs merge=lfs -text
40
+ image/blobs/sha256/540cf00275e913a9bccc49fe7beba58037696661bde6ae083a3ac84c5b160e67 filter=lfs diff=lfs merge=lfs -text
41
+ image/blobs/sha256/b6df468b82a4b2f9ee3ca3a79a6bbe99b5cda02ddc63b7e3c89e7bd08ef41706 filter=lfs diff=lfs merge=lfs -text
42
+ image/blobs/sha256/0926a8eb0e608a5c6888d1cd5594184bdf3ed3aa311dba5b42a547caefdc6f2e filter=lfs diff=lfs merge=lfs -text
43
+ image/blobs/sha256/c3bc7b373b4523cabcdd9c64ab10d32510ef61bace60873a279ccc4902738989 filter=lfs diff=lfs merge=lfs -text
44
+ image/blobs/sha256/fdc1ed79ffd24d66f8be3754ec4dc80b1ab0fcc8c8165a6007748c464c94f897 filter=lfs diff=lfs merge=lfs -text
45
+ image/blobs/sha256/c18d0f3c8022bcd5a8059f67fc0e5cfd53ef36a44997335d6cd0aa6b19db140d filter=lfs diff=lfs merge=lfs -text
46
+ image/blobs/sha256/8753e0cfbd424e962ccaf50aaaf02fd06ff2efb8677219657e751a53922efa9f filter=lfs diff=lfs merge=lfs -text
47
+ image/blobs/sha256/24cba7375920bef8d4cc4f0ce4294f8f70c65b9a37c4ed6f6c2d63405c76ba3c filter=lfs diff=lfs merge=lfs -text
48
+ image/blobs/sha256/e19b6f1fb65dd2888d9003ef9513a21d129041a328bb8a9a4164d29ef0382b16 filter=lfs diff=lfs merge=lfs -text
49
+ image/blobs/sha256/c00314a02c644cf2aea5a8ae4e3ee5ee383f4072d3e35248fe20dea57a90c46c filter=lfs diff=lfs merge=lfs -text
code/models/demos/blackhole/qwen3_coder_30b_a3b/README.md ADDED
@@ -0,0 +1,143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Qwen3-Coder-30B-A3B-Instruct on Blackhole
2
+
3
+ This directory implements Tenstorrent Blackhole inference for
4
+ **`Qwen/Qwen3-Coder-30B-A3B-Instruct`** — a 48-layer `Qwen3MoeForCausalLM`
5
+ sparse mixture-of-experts model (30B total / ~3B active parameters) with an
6
+ advertised context of 262144 tokens.
7
+
8
+ | Model | `HF_MODEL` | Mesh / `--mesh-device` | Parallelism |
9
+ | ----- | ---------- | ---------------------- | ----------- |
10
+ | Qwen3-Coder-30B-A3B-Instruct | `Qwen/Qwen3-Coder-30B-A3B-Instruct` | 4 Blackhole dies — `P300x2` (a `(1, 4)` mesh) | 4-way tensor parallel |
11
+
12
+ The 4-die path needs `FABRIC_1D_RING` for the cross-device collectives and a
13
+ trace region for the captured decode and chunked-prefill traces
14
+ (`DEFAULT_TRACE_REGION_SIZE = 300_000_000`, see [tt/model.py](tt/model.py)).
15
+
16
+ ## Architecture
17
+
18
+ Assembly: `tok_embeddings → 48 × decoder layer → RMSNorm → LM head → on-device
19
+ sampling`.
20
+
21
+ Each decoder layer is GQA attention (rotary, per-head QK-norm) followed by a
22
+ sparse MoE block — a router plus top-k expert MLPs — replacing the dense MLP.
23
+ Everything shape-related (layer count, expert count, top-k, head dims, vocab,
24
+ rope base) is read from the parsed HF config, so the code follows the
25
+ checkpoint rather than hard-coding it.
26
+
27
+ | File | Role |
28
+ | ---- | ---- |
29
+ | [tt/model.py](tt/model.py) | Full 48-layer model, weight load, KV cache, trace capture, LM head + sampling |
30
+ | [tt/functional_decoder.py](tt/functional_decoder.py) | Reference-shaped single-layer decoder (attention + MoE) — the correctness baseline |
31
+ | [tt/optimized_decoder.py](tt/optimized_decoder.py) | Single-device optimized layer (fused QKV, sharded matmuls, program configs) |
32
+ | [tt/multichip_decoder.py](tt/multichip_decoder.py) | Tensor-parallel layer and the CCL schedule for the 4-die mesh |
33
+ | [tt/precision.py](tt/precision.py) | `PrecisionConfig` — per-tensor dtypes and math fidelity, overridable via `QWEN3_PRECISION_CONFIG` |
34
+ | [config/](config/) | Runtime policy the serving path reads: the selected precision config and the served-context contract |
35
+ | [tt/weight_mapping.py](tt/weight_mapping.py) | HF checkpoint → device tensor layout (QKV permutes, expert stacking) |
36
+ | [tt/generator.py](tt/generator.py) | `build_generator()` + the high-level `generate()` loop (owns KV cache and page table) |
37
+ | [tt/generator_vllm.py](tt/generator_vllm.py) | vLLM adapter — `prefill_forward` / `decode_forward` against a caller-owned cache |
38
+ | [vllm_bundle/](vllm_bundle/) | `EXTRA_MODELS_DIR` bundle that registers the adapter with the TT vLLM plugin |
39
+
40
+ The generator implements the `Generator` ABC in
41
+ [models/common/readiness_check/contract.py](../../../common/readiness_check/contract.py),
42
+ so the same object serves the host-side demo/readiness path and vLLM.
43
+
44
+ ## Precision
45
+
46
+ A 29-row datatype sweep over the attention, MoE, KV, CCL, norm and LM-head
47
+ tensors selected [config/selected_precision_config.json](config/selected_precision_config.json),
48
+ which the vLLM path loads on every serve so that serving and readiness cannot
49
+ run different numerics. `DEFAULT_PRECISION` in [tt/precision.py](tt/precision.py)
50
+ is the equivalent in-code default. Override for experiments with
51
+ `QWEN3_PRECISION_CONFIG=<path-to-json>`.
52
+
53
+ Served context is capped by
54
+ [config/context_contract.json](config/context_contract.json) rather than by the
55
+ `--max-model-len` you pass, so a request for more context than has been
56
+ validated fails loudly instead of serving a quietly-clipped model.
57
+
58
+ ## Running the tests
59
+
60
+ All device tests target the 4-die mesh. From the repository root:
61
+
62
+ ```bash
63
+ source python_env/bin/activate
64
+ export HF_MODEL=Qwen/Qwen3-Coder-30B-A3B-Instruct
65
+ D=models/demos/blackhole/qwen3_coder_30b_a3b
66
+
67
+ # module + model correctness (excludes the perf-only tests)
68
+ pytest $D/tests/ -m "not models_performance_bare_metal" -q
69
+
70
+ # perf tests (decode/prefill timings; writes CSVs under doc/)
71
+ pytest $D/tests/test_perf.py -q
72
+ ```
73
+
74
+ `test_full_model.py` runs a 2-layer model by default so it stays cheap; set
75
+ `QWEN3_FULL_MODEL_LAYERS=48` for the complete model.
76
+
77
+ Host-only tests (no device): `tests/test_reference.py`,
78
+ `tests/test_precision_config.py`.
79
+
80
+ ## Readiness checks and vLLM serving
81
+
82
+ The shared harness in [models/common/readiness_check/](../../../common/readiness_check/)
83
+ drives this model through `tt/generator.py`:
84
+
85
+ ```bash
86
+ # teacher-forced accuracy against a reference completion
87
+ python -m models.common.readiness_check.run_prefill_check \
88
+ --model-dir models/demos/blackhole/qwen3_coder_30b_a3b \
89
+ --reference <reference.refpt> \
90
+ --mesh-device P300X2 --fabric-config FABRIC_1D_RING --trace-region-size 300000000
91
+ ```
92
+
93
+ vLLM serving registers through the plugin's `EXTRA_MODELS_DIR` hook — no edit
94
+ to the vLLM checkout is required:
95
+
96
+ ```bash
97
+ export EXTRA_MODELS_DIR=$PWD/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle
98
+
99
+ python -m models.common.readiness_check.run_vllm_server \
100
+ --model-dir <output-dir> --hf-model Qwen/Qwen3-Coder-30B-A3B-Instruct \
101
+ --mesh-device P300x2 --max-num-seqs 1 --max-model-len 262144 \
102
+ --block-size 32 --port 8100 --stages serve \
103
+ --tt-config '{"trace_region_size": 50331648, "fabric_config": "FABRIC_1D_RING"}' \
104
+ --additional-server-args "--generation-config vllm"
105
+ ```
106
+
107
+ `--generation-config vllm` matters: this checkpoint's `generation_config.json`
108
+ injects `repetition_penalty=1.05` into every request that does not override it,
109
+ and a penalised request costs ~14% TPOT because the penalty operands are staged
110
+ per step on the host.
111
+
112
+ ## Measured performance
113
+
114
+ 4 Blackhole dies, 48 layers, traced decode, on-device sampling, greedy,
115
+ 128-token input / 128-token output, batch 1.
116
+
117
+ | Path | TTFT | Decode |
118
+ | ---- | ---- | ------ |
119
+ | Standalone traced (`generate()`) | 129.9 ms | 19.213 ms — **52.05 t/s/u** |
120
+ | Through vLLM (`max_num_seqs=1`, `--max-concurrency 1`) | 307–312 ms | 19.78 ms — **50.3–50.6 t/s/u** |
121
+
122
+ vLLM adds ~0.57 ms per decoded token (2.9%) over the standalone traced path.
123
+ The TTFT gap is request handling, tokenisation and detokenisation, not decode.
124
+
125
+ Evaluated end to end by tt-inference-server 0.20.0 against a live server:
126
+ mbpp 77.2%, humaneval 92.7%, ifeval 81.1%/87.1%, gpqa_diamond_cot 56.1%.
127
+
128
+ ## Known limitations
129
+
130
+ - **Long-prefill scaling is bad.** A 131072-token prefill completes and returns
131
+ valid output but takes 94.4 minutes — 3.05× worse per token than 65536, and
132
+ well off what the single-layer sweep predicts. The suspected cause is
133
+ per-chunk tensor accumulation in the MoE prefill path; this is an unproven
134
+ hypothesis, not a measured root cause.
135
+ - **No full-model 262144 prefill has been verified.** 262144 is allocated,
136
+ page-tabled and served, and prefills through a single layer, but the largest
137
+ 48-layer prefill measured end to end is 131072. The advertised context is
138
+ left at 262144.
139
+ - **Not yet registered in the tiered models CI.** No entry exists in
140
+ `models/model_ci_tiers.md`, `tests/pipeline_reorg/models_*_tests.yaml`,
141
+ `models/model_targets.yaml` or the vLLM test registry, and there is no
142
+ `demo/` entry point in the shape those pipelines invoke. See
143
+ [models/MIGRATING_TO_TIERED_CI.md](../../../MIGRATING_TO_TIERED_CI.md).
code/models/demos/blackhole/qwen3_coder_30b_a3b/__init__.py ADDED
File without changes
code/models/demos/blackhole/qwen3_coder_30b_a3b/config/context_contract.json ADDED
@@ -0,0 +1,360 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "hf_model": "Qwen/Qwen3-Coder-30B-A3B-Instruct",
3
+ "hf_advertised_context": 262144,
4
+ "current_supported_context": 256000,
5
+ "capability_reduction": true,
6
+ "limiting_reason": "Serving prefill. TWO independent limits, both above this bound. (1) An UNRESOLVED prefill cliff: 259,000 tokens completes in 1130.8 s, tracking the O(S^2) curve to within 0.3 %, but 262,136 tokens does not complete in 3600 s -- one continuous blocking prefill, no watchdog throw, engine healthy. The cliff is confined to the last ~3,100 tokens of the advertised context and its mechanism was deliberately not chased (see stage11_serving_context). (2) LATENCY: prefill at this length is minutes, not seconds -- 253,000 tokens costs 1049.3 s (17.5 min) of TTFT on the path that WORKS. The usable interactive range is bounded far below this figure by latency, not by the cliff.",
7
+ "stage": "optimized_full_model",
8
+ "scope": "the complete 48-layer causal LM on the full 4-die P300_X2 mesh (1x4), batch 1, AFTER the stage-06 optimizations: embeddings, 48 stage-04 multichip decoder layers, final norm, column-parallel lm_head and on-device traced sampling with a distributed argmax. The single-layer stage-02/03/04 numbers and the stage-05 full-model numbers below are retained unchanged; the stage06_* fields are the shipped ones.",
9
+ "device": {
10
+ "arch": "blackhole",
11
+ "board": "p300 x2 (ClusterType.P300_X2), 4 dies",
12
+ "mesh": "1x4",
13
+ "dram_per_die_gb": 34.18,
14
+ "dram_per_die_note": "34.18 GB is what the TTNN allocator reports as total DRAM per die on this host (doc/multichip_decoder/footprint_probe.log, first line). Stages 01 and 02 recorded the nominal board figure of 32 GB; nothing they concluded changes.",
15
+ "topology": "ring, 2 ethernet links per hop, FABRIC_1D_RING"
16
+ },
17
+ "measured": {
18
+ "decode_context_tokens": 262144,
19
+ "decode_probe": "paged KV cache allocated to depth 262144 (block_size=32); decode step executed at position 262143; output finite; 2.42 s",
20
+ "prefill_context_tokens": 262144,
21
+ "prefill_probe": "single-shot prefill of a 262144-token sequence through the full decoder layer; output finite; 192.23 s",
22
+ "prefill_sweep_seconds": {
23
+ "512": 0.51,
24
+ "1024": 4.7,
25
+ "2048": 5.35,
26
+ "4096": 7.92,
27
+ "8192": 10.24,
28
+ "16384": 14.87,
29
+ "32768": 23.79,
30
+ "65536": 42.05,
31
+ "131072": 85.85,
32
+ "262144": 192.23
33
+ },
34
+ "decode_sweep_seconds": {
35
+ "4096": 1.19,
36
+ "16384": 1.21,
37
+ "65536": 1.45,
38
+ "131072": 1.77,
39
+ "262144": 2.42
40
+ },
41
+ "prefill_sweep_seconds_PROVENANCE": "SINGLE-LAYER, PRE-OPTIMISATION. Not comparable to the shipped 48-layer model and not to be divided into any full-model figure. This file's own `scope` says 'The single-layer stage-02/03/04 numbers ... are retained unchanged', and `prefill_probe` describes 'a single-shot prefill ... through the full decoder LAYER' -- one complete layer, not the stack. The evidence list points at doc/functional_decoder/ and doc/multichip_decoder/, i.e. stages 02-04, before the stage-04/06 optimisations. Measured on the current tree, the whole 48-layer model prefills 32,768 tokens in 37.638 s, i.e. 0.78 s per layer, against 23.79 s for the single layer recorded here -- ~30x apart, so the two are measuring different things on different code. In particular the 262144 -> 192.23 s entry does NOT describe the shipped model, and no current-tree measurement of that length exists.",
42
+ "prefill_same_tree_seconds": {
43
+ "note": "Full 48-layer model on the current tree. Standalone and served at a MATCHED length, which is the comparison that retires the suspected serving penalty.",
44
+ "32768": {
45
+ "standalone": 37.638,
46
+ "served": 37.778,
47
+ "served_over_standalone": 1.004
48
+ },
49
+ "65536": {
50
+ "served": 98.557,
51
+ "standalone": null,
52
+ "warning": "Do NOT divide 98.557 by the 42.05 s single-layer entry above; that repeats the exact error this block exists to correct."
53
+ },
54
+ "evidence": [
55
+ "doc/batch_scaling/probes/expert_chunk_sweep_v3.json (standalone 32,768)",
56
+ "doc/batch_scaling/probes/long_prompt_gap_32768.json",
57
+ "doc/batch_scaling/probes/long_prompt_gap_65536.json"
58
+ ]
59
+ }
60
+ },
61
+ "kv_cache": {
62
+ "dtype": "bfloat16",
63
+ "bytes_per_token_per_layer": 2048,
64
+ "formula": "num_key_value_heads(4) * head_dim(128) * 2 bytes * 2 tensors (K and V)",
65
+ "bytes_at_full_context_per_layer": 536870912,
66
+ "paged": true,
67
+ "default_block_size": 32
68
+ },
69
+ "pcc_validated_context_tokens": 512,
70
+ "pcc_validation_note": "PCC against the HuggingFace reference is verified up to 512 tokens; longer lengths are capacity/liveness probes only, because a full-length torch reference for 262144 tokens does not fit in host RAM. This is a limit of the reference, not of the TTNN decoder.",
71
+ "forward_looking_note": "DISCHARGED by stage 03. The stage-01 note read: 'the full 48-layer model would need ~24 GiB on a 32 GiB die, so the full-model stage will have to weigh KV dtype, paging across dice, or a served-context cap.' None of those trade-offs is needed. TP=4 across the four dies puts the whole 48-layer model plus a full 262144-token paged KV cache in 11.759 GB per die with 22.119 GB free -- measured in stage 05 by allocating the real 48-layer model, and the figure this note now carries (see full_model_measured and comparison_to_stage03_prediction). Stage 03 predicted 11.829 / 22.350 from a synthetic allocation of the same shapes; the prediction and the measurement agree to within 0.6%. KV dtype is unchanged (bfloat16), paging is unchanged (block_size 32), and the served context is unchanged. The note is kept rather than deleted because it is the question this stage answers.",
72
+ "evidence": [
73
+ "doc/functional_decoder/work_log.md (context probe section)",
74
+ "doc/multichip_decoder/README.md and work_log.md",
75
+ "doc/multichip_decoder/footprint_probe.log (the allocation this file quotes)",
76
+ "doc/multichip_decoder/probes/footprint_probe.py",
77
+ "doc/multichip_decoder/pcc_log.txt",
78
+ "doc/optimized_multichip_decoder/README.md and work_log.md",
79
+ "doc/optimized_multichip_decoder/pcc_log.txt",
80
+ "doc/full_model/README.md and work_log.md",
81
+ "doc/full_model/probes/footprint_262144.json (the allocation this file quotes)",
82
+ "doc/full_model/probes/perf_full_model.csv and .json",
83
+ "doc/full_model/run_prefill_check.log, run_teacher_forcing.log, run_autoregressive.log",
84
+ "doc/optimized_full_model/README.md and work_log.md",
85
+ "doc/optimized_full_model/profile_48layer_work_log.md (the lever analysis; its op-level figures are the pre-adoption ones)",
86
+ "doc/optimized_full_model/ops_perf_full_model_48layer_decode.csv.gz (the verified one-iteration decode window)",
87
+ "doc/optimized_full_model/ops_perf_full_model_48layer_prefill_s128.csv.gz (the 48-layer prefill window)",
88
+ "doc/optimized_full_model/probes/profile_summary_decode.json and profile_summary_prefill.json",
89
+ "doc/optimized_full_model/probes/perf_full_model_p{128,1024,4096}_{before,after,argmaxrows}.json",
90
+ "doc/optimized_full_model/probes/footprint_262144.json (the allocation this file quotes)",
91
+ "doc/optimized_full_model/probes/runtime_fallback_audit.json",
92
+ "doc/optimized_full_model/probes/sdpa_sweep_confirm_bf16.json and sdpa_hf_pcc_at_depth.json",
93
+ "doc/optimized_full_model/logs/run_prefill_check_argmaxrows.log, run_teacher_forcing_argmaxrows.log, run_autoregressive_argmaxrows.log, check_degenerate_argmaxrows.log"
94
+ ],
95
+ "expert_weight_dtype": "bfloat4_b",
96
+ "optimized_note": "Stage 02 changed expert weight dtype to bfloat4_b, attention projection dtype to bfloat8_b, the expert matmul block widths, expert math fidelity to LoFi, and the decode attention projections to a DRAM-sharded program config. None of these touches the KV cache layout, dtype, or paging, so the measured context limits carry over unchanged. Per-layer weight DRAM: expert weights shrink ~3.6x (bf16 to bfloat4_b, which is 0.5625 B/elem once each 16-element block's exponent byte is counted). Decode additionally needs a second, DRAM-width-sharded copy of wqkv/wo; at bfloat8_b (1.0625 B/elem) each copy is 20.05 MB, so the pair is 40.11 MB against the single bf16 copy stage 01 held at 37.75 MB -- a 2.36 MB increase, not a wash, and an earlier revision of this file called it one by rounding bfloat8_b to 1 B/elem. Against 24 GB of usable DRAM it changes nothing: no capability reduction.",
97
+ "attention_weight_dtype": "bfloat8_b",
98
+ "parallelism": {
99
+ "attention": "TP=4 (8 Q heads, 1 K head, 1 V head per die)",
100
+ "experts": "EP=4 (32 of 128 experts per die)",
101
+ "router": "replicated",
102
+ "norms": "replicated",
103
+ "residual": "replicated [1,1,B,2048]",
104
+ "collectives_per_layer": "2 all-reduces, each reduce-scatter then all-gather on dim 3",
105
+ "padding": "none -- 2048/4, 32/4, 4/4, 128/4 and 151936/4 are all exact"
106
+ },
107
+ "multichip_measured": {
108
+ "method": "doc/multichip_decoder/probes/footprint_probe.py -- allocates the real per-die shapes rather than computing their size. Output archived verbatim at doc/multichip_decoder/footprint_probe.log.",
109
+ "decoder_weights_48_layers_gb_per_die": 4.596,
110
+ "plus_embed_and_lm_head_gb_per_die": 5.374,
111
+ "embed_lm_head_note": "embed_tokens replicated (bf16), lm_head column-parallel, per-die N = 37984 = 151936/4.",
112
+ "kv_bytes_per_token_per_layer_per_die": 512,
113
+ "kv_formula": "1 local kv head * head_dim(128) * 2 bytes * 2 tensors (K and V); one quarter of the single-die 2048",
114
+ "total_at_full_context_batch1_gb_per_die": 11.829,
115
+ "total_note": "MEASURED: 48 layers of sharded weights + embed + lm_head + 48 paged KV caches at 262144 tokens, all allocated simultaneously on the mesh. 22.350 GB/die free afterwards. The arithmetic prediction in mesh_plan.md section 8 was 11.80 GB; the allocator says 11.829.",
116
+ "headroom_gb_per_die": 22.35,
117
+ "single_die_equivalent_gb": 45.28,
118
+ "single_die_note": "One die would need 19.52 GB of weights plus 25.77 GB of KV at the advertised context. That does not fit in 34.18 GB, so the mesh is a capability requirement here and not only a speed one. This is what discharges the stage-01 forward_looking_note.",
119
+ "single_die_convention": "Recomputed from the probe's own per-die tensor list (footprint_probe.log, GB = 1e9 as the probe reports it). Convention: tensors that are sharded across the mesh are counted 4x -- the two expert tensors, both wqkv copies, both wo copies, and lm_head -- and tensors that are replicated are counted ONCE, since a single die would hold exactly one of each: the router (2048x128 bf16), the two RMSNorm vectors, and embed_tokens. That gives 48 * (4 * 94.962 MB sharded + 0.786 MB replicated) = 18.270 GB of decoder weights, + 0.622 GB embed + 4 * 0.156 GB lm_head = 19.515 GB, i.e. 19.52. Counting the replicated router/norms 4x as well would read 19.63; the difference is 0.11 GB and changes nothing. The 19.50 this field previously carried came from mesh_plan.md section 8, whose per-layer table omits the two tile-padded RMSNorm vectors (0.26 MB/layer, 0.013 GB over 48 layers) and rounds each row; recomputing from the probe's tensor list closes that. KV is 4 kv heads * 128 * 2 B * 2 tensors = 2048 B/token/layer * 48 layers * 262144 tokens = 25.770 GB. Total 45.28 GB against 34.18 GB of DRAM -- the conclusion is unaffected by which convention is used."
120
+ },
121
+ "multichip_largest_feasible": {
122
+ "context_at_batch1": 262144,
123
+ "context_at_batch1_basis": "the HF-advertised context, allocated and held on the mesh; not a limit reached here",
124
+ "batch_at_full_context": 4,
125
+ "batch_at_full_context_basis": "floor((34.18 - 5.374 weights - ~2 GB trace/activation) / 6.442 GB of KV per user) = floor(26.806 / 6.442) = floor(4.16) = 4. KV per user is 48 layers * 512 B/token/layer/die * 262144 tokens = 6.442 GB, from the measured per-die per-token figure above. Arithmetic on the measured per-user KV, not allocated; an earlier revision of this file recorded 3, which its own formula does not give.",
126
+ "batch_ceiling_from_ops": 32,
127
+ "batch_ceiling_basis": "nlp_create_qkv_heads_decode_device_operation.cpp:51 asserts num_users <= 32. A TTNN op limit, unchanged by TP, and reached long before DRAM is.",
128
+ "context_at_batch32": 28788,
129
+ "context_at_batch32_basis": "~22.6 GB usable / (32 users * 48 layers * 512 B/token); arithmetic"
130
+ },
131
+ "multichip_pcc_note": "Multichip PCC is validated against the single-chip TTNN optimized decoder run replicated on the same mesh, 0.99962-0.99997 across prefill S = 32/33/100/128/257/512 and decode in both cache modes, and against HF at 0.9990 (prefill) and 0.9988-0.9993 (decode). See doc/multichip_decoder/pcc_log.txt. The 512-token PCC validation ceiling of the stage-01/02 entry above is unchanged and has the same cause: a full-length torch reference does not fit in host RAM. Stage 04 re-ran every one of those gates on the optimized path: 112 passed, 0 failed (stage 03's 111 plus the host-only test_meta_rope_weights_match_hf), with prefill 0.999620-0.999868 and decode 0.999935-0.999941 against the single-chip TTNN baseline. The one movement is decode vs single-chip, 0.99997 -> 0.99994, and it is the right sign: the reference is optimized_decoder.py, whose numerics stage 04 left unchanged (its only edit is an optional rope= seam that defaults to the shipped op) and which still accumulates its RMSNorm sum of squares in bf16, while the multichip norm now measures 1.686e-02 against a torch fp64 reference where the old one measured 6.711e-02. Against HF the four decode steps move by at most 8e-05 and in both directions (0.9992853->0.9992450, 0.9987817->0.9987388, 0.9993126->0.9993928, 0.9989592->0.9989421), i.e. three are marginally lower and one is higher; that is bf16 rounding under a changed norm, not a regression, and the earlier wording 'nothing regressed' was literally false. See doc/optimized_multichip_decoder/pcc_log.txt.",
132
+ "stage04_note": "Stage 04 optimized the multichip decoder in place and changed NOTHING this file measures. No capability reduction. What it changed: both residual RMSNorms now run width-sharded over 8 L1 cores instead of one core, the router projection reads that L1 shard instead of DRAM-interleaved, and the two collectives use caller-owned persistent buffers. KV cache dtype, layout, paging and bytes/token are untouched, and so is the activation sharding of the cache. Two things grew: each layer carries a second, ROW_MAJOR copy of each residual RMSNorm vector (the layout the sharded rms_norm program factory reads) at 4 KB each, i.e. 8 KB per layer and 0.384 MB per die over all 48 layers, against 4.596 GB of sharded decoder weights; and the mesh context owns two sets of persistent reduce-scatter/all-gather buffers of about 0.5 MB each, once per mesh rather than once per layer. Both are far below the 22.350 GB/die of measured headroom, so the 262144-token contract, the batch-4-at-full-context figure and the single-die comparison all stand unchanged. Evidence: doc/optimized_multichip_decoder/README.md and work_log.md section 9.",
133
+ "full_model_measured": {
134
+ "method": "doc/full_model/probes/footprint_probe.py --context 262144. Unlike the stage-03 probe this builds the REAL model -- real weights, real embed_tokens, real lm_head, the real paged KV cache at the advertised context, the real RoPE tables -- captures both decode traces and runs a token through it. Output archived at doc/full_model/probes/footprint_262144.json and .log. GB = 1e9, as the allocator reports it.",
135
+ "weights_embed_lm_head_rope_gb_per_die": 5.311,
136
+ "kv_cache_at_262144_batch1_gb_per_die": 6.443,
137
+ "traces_and_persistent_buffers_gb_per_die": 0.006,
138
+ "total_gb_per_die": 11.759,
139
+ "headroom_gb_per_die": 22.119,
140
+ "dram_per_die_gb_reported": 33.879,
141
+ "dram_per_die_note": "33.879 rather than the 34.18 recorded for stage 03/04 because this probe opens the mesh with trace_region_size=300 MB, which the allocator carves out of DRAM before reporting. It is the same hardware; the difference is 0.30 GB and is the trace region the decode traces live in.",
142
+ "sum_note": "the three stage rows in doc/full_model/probes/footprint_262144.json sum to 11.759415296, BIT-IDENTICAL to total_gb_per_die -- there is no residual and nothing is omitted. The 0.001 that appears if the 3dp values in this file or in the README's table are added up is introduced by that display rounding alone. An earlier revision of this field attributed it to 'the allocator's rounding across the three separate reads', which the raw JSON contradicts; doc/full_model/probes/check_published_figures.py now asserts the raw rows equal the raw total exactly.",
143
+ "comparison_to_stage03_prediction": "stage 03's footprint probe ALLOCATED 11.829 GB/die for the same contents and never ran through them. Executing it measures 11.759, 0.070 GB lower: the running model carries one shared pair of RoPE tables and one shared set of persistent collective buffers where the allocation probe's per-layer tensor list double-counted nothing, and the two paged caches are allocated once per layer at exactly ceil(262144/32) blocks. The stage-03 figure remains the honest allocation-time answer and is not restated here as wrong.",
144
+ "weight_load_seconds": 182.66,
145
+ "weight_load_note": "48 layers streamed one at a time out of the sharded safetensors checkpoint; the 61 GB checkpoint is never fully materialised on the host.",
146
+ "rope_table_note": "the probe builds the model with the default rope_cache_len of 8192, so the RoPE tables inside weights_embed_lm_head_rope_gb_per_die are the 8192-row pair (0.004 GB/die). Growing them to the full 262144 context costs 0.134 GB/die (262144 * 128 * 2 B * 2 tables), taking the total to 11.889 GB/die and the headroom to 21.985. No capability reduction; the tables grow on demand via Qwen3CoderModel.ensure_rope_capacity, and tt/generator.py's decode_forward takes decode_horizon= so a low-level caller grows them once before any trace is captured."
147
+ },
148
+ "full_model_context_note": "NO REDUCTION. The advertised 262144 context is allocated and held by the running 48-layer model with 22.119 GB/die free. What has NOT been done at the full stack is a 262144-token prefill: prefill is single-shot (no internal chunking), and the longest prompt actually pushed through all 48 layers in this stage is 1000 tokens. The 262144-token single-layer prefill probe in measured.prefill_probe above still stands, and the KV cache, page tables and positions are sized and exercised for the full context. This is a coverage gap in the evidence, not a capability reduction, and it is named as limitation 1 of doc/full_model/README.md. STAGE 06: still no reduction, and the 262144 context is now usable rather than only allocatable -- decode is nearly flat in context (1.04x from 128 to 4096, against 1.96x before). See stage06_context_flatness for exactly how deep that was measured and what is still unmeasured.",
149
+ "full_model_batch": {
150
+ "primary": 1,
151
+ "primary_note": "every performance figure in doc/full_model/README.md is batch 1.",
152
+ "largest_tested": 4,
153
+ "largest_tested_note": "tests/test_full_model.py::test_mixed_length_batch_prefill_and_decode runs four users at prompt lengths 7 / 33 / 64 / 129 through one prefill and one decode with disjoint physical cache pages. Batch was not pushed to 32 at the full stack in this stage.",
154
+ "hard_ceiling": 32,
155
+ "hard_ceiling_basis": "nlp_create_qkv_heads_decode_device_operation.cpp:51 asserts num_users <= 32. A TTNN op limit, unchanged by TP, and reached long before DRAM is."
156
+ },
157
+ "full_model_performance": {
158
+ "workload": "prompt 128 / generate 128 / batch 1, 48 layers, 1x4 P300_X2, FABRIC_1D_RING",
159
+ "source": "doc/full_model/probes/perf_full_model.csv and .json",
160
+ "ttft_ms_warmed": 126.695,
161
+ "ttft_ms_cold": 1115.157,
162
+ "decode_logits_only_ms": 20.211,
163
+ "decode_logits_only_tps_user": 49.479,
164
+ "decode_token_out_ms": 22.079,
165
+ "decode_token_out_tps_user": 45.292,
166
+ "decode_token_out_with_readback_ms": 22.748,
167
+ "decode_token_out_with_readback_tps_user": 43.959,
168
+ "layer_stack_lower_bound_ms": 20.573,
169
+ "layer_stack_lower_bound_basis": "48 x the stage-04 traced decode layer at ctx128, 0.4286 ms",
170
+ "ttft_cold_note": "first pass through a freshly opened mesh, so dominated by JIT kernel compilation and by the state of ~/.cache/ttnn. An earlier revision recorded 188.866 from a run with a warmer cache. Warmed TTFT is the served figure.",
171
+ "sampler_note": "greedy routes to Sampling1D's force-argmax strategy, measured at 1.125 ms against the split path's 6.155 ms on the same logits, both returning token 16. The 1.125 figure is after tt/model.py's _WatcherCleanSampling1D stopped pinning num_workers_per_link=1 on the argmax gather; it was 1.859 with the upstream spelling, which also tripped a device ASSERT under the watcher. See doc/full_model/README.md and watcher_ab.log.",
172
+ "watcher": "clean. TT_METAL_WATCHER=10 TT_METAL_WATCHER_DISABLE_ETH=1 pytest tests/ -m 'not models_performance_bare_metal' -q is 145 passed with zero tripped asserts (doc/full_model/pytest_watcher_clean.log.gz)."
173
+ },
174
+ "full_model_accuracy": {
175
+ "reference": "readiness_aime24_chat.refpt, AIME24 prompt 0, HF chat template, 158 prompt tokens, gen_len 100, top_k 100, generated fresh by stage 05",
176
+ "prefill": {
177
+ "top1": 0.98,
178
+ "top5": 1.0,
179
+ "top100": 1.0,
180
+ "log": "doc/full_model/run_prefill_check.log"
181
+ },
182
+ "decode_teacher_forced": {
183
+ "top1": 0.99,
184
+ "top5": 1.0,
185
+ "top100": 1.0,
186
+ "log": "doc/full_model/run_teacher_forcing.log"
187
+ },
188
+ "bar": "top5 >= 0.98 and top100 == 1.00; both met at 1.000"
189
+ },
190
+ "full_model_policy_note": "Stage 05 changed NOTHING about the decoder layer's dtype, fidelity, KV, activation or CCL policy, the paged-cache dtype split, the rejection ledger, or the inter-layer residual layout. The one edit to tt/multichip_decoder.py is an optional rope= parameter defaulting to the op it already called, and the stage-04 suite re-runs at 112 passed / 0 failed on this tree. What the full model adds is a replicated bf16 embedding (0.622 GB/die, no collective), a replicated final norm, a column-parallel bfloat8_b lm_head with zero vocabulary padding (151936 = 4 x 37984), growable replicated RoPE tables, and two captured traces plus their persistent inputs -- 0.006 GB/die of traces and buffers in total. Decode rotary moved from ttnn.experimental.rotary_embedding (Python-int position, unreplayable in a trace) to ttnn.experimental.rotary_embedding_hf over a device-gathered cos/sin pair; both are HF rotate_half and the swap is bit-identical at max|diff| 0.000e+00 and PCC 1.0 at every position tested, so the KV cache channel convention is untouched.",
191
+ "stage06_note": "Stage 06 changed NOTHING this file measures about capability. No capability reduction. What it changed, all three measured: (1) the greedy sampler stopped all-gathering the 151936-wide logit row and now reduces per die and all-gathers four candidate values and indices (tt/model.py, _WatcherCleanSampling1D._sample_argmax); (2) the PAGED SDPA-decode call got the program config it never had -- q_chunk_size 32, k_chunk_size min(256, per-user cache depth), max_cores_per_head_batch 16, memoised on the compute grid (tt/multichip_decoder.py, _sdpa_program_config); (3) that same reduction now runs over the live user rows instead of the 32 fixed sampler slots. KV-cache dtype, layout, paging and bytes/token are untouched; so are expert and attention weight dtypes, math fidelity, activation memory configs, the CCL policy (Topology.Ring, no num_workers_per_link pinned, 2 links prefill / 1 decode) and the inter-layer residual layout. tt/functional_decoder.py gained an sdpa_program_config= seam on attention_prefill that defaults to None and is passed None -- the prefill lever is built, measured and NOT adopted. Evidence: doc/optimized_full_model/README.md and work_log.md.",
192
+ "stage06_performance": {
193
+ "workload": "prompt 128 / generate 128 / batch 1, 48 layers, 1x4 P300_X2, FABRIC_1D_RING",
194
+ "source": "doc/optimized_full_model/probes/perf_full_model_p128_argmaxrows.json (128 timed reps, median)",
195
+ "ttft_ms_warmed": 125.431,
196
+ "ttft_ms_cold": 221.494,
197
+ "decode_logits_only_ms": 19.567,
198
+ "decode_logits_only_tps_user": 51.107,
199
+ "decode_token_out_ms": 19.693,
200
+ "decode_token_out_tps_user": 50.781,
201
+ "decode_token_out_with_readback_ms": 19.71,
202
+ "decode_token_out_with_readback_tps_user": 50.735,
203
+ "teacher_forcing_decode_tps_user": 42.25,
204
+ "teacher_forcing_note": "run_teacher_forcing uploads a forced token and reads the prediction back every step, so it is a correctness gate that prints a rate and is NOT the same measurement as token-out. It moved 38.50 -> 42.25 t/s/u (doc/full_model/run_teacher_forcing.log, doc/optimized_full_model/logs/run_teacher_forcing_argmaxrows.log).",
205
+ "against_stage05": {
206
+ "source": "doc/full_model/probes/perf_full_model.json",
207
+ "decode_token_out_ms": 22.079,
208
+ "decode_token_out_tps_user": 45.292,
209
+ "ttft_ms_warmed": 126.695,
210
+ "token_out_speedup": 1.1212,
211
+ "note": "the stage-05 run allocated a 4096-position KV cache and the stage-06 runs allocate 8192. SDPA-decode cost is independent of ALLOCATED depth and linear in cur_pos, measured directly at doc/optimized_full_model/probes/sdpa_depth_probe.json, so the comparison holds; the stage-06 before/after legs are all at 8192 and are like-for-like throughout."
212
+ },
213
+ "layer_stack_lower_bound_ms": 18.47,
214
+ "layer_stack_lower_bound_basis": "48 x the OPTIMIZED IN-MODEL per-layer device-kernel time, 384.791 us, from the verified one-iteration window at doc/optimized_full_model/ops_perf_full_model_48layer_decode.csv.gz (summary: probes/profile_summary_decode.json, regions_us.layer_stack). This REPLACES the stage-05 basis of 48 x 0.4286 ms = 20.573 ms, which multiplied a WALL figure for a one-layer traced model and so charged 48 layers for one iteration's dispatch overhead -- which is why stage 05 appeared to be under its own lower bound. The stage-04 layer's isolated device-kernel content is 362.83 us (doc/optimized_multichip_decoder/window_decode.txt); the in-model layer is 6.1% dearer than that.",
215
+ "bound_plus_terminal_ms": 18.889,
216
+ "gap_to_token_out_ms": 0.803,
217
+ "gap_to_token_out_percent": 4.08,
218
+ "gap_note": "dispatch and op-to-op gap across 3512 device ops -- 0.23 us each. The stage goal flags >10-15% as needing action; this is 4.08%.",
219
+ "prefill_profile": {
220
+ "source": "doc/optimized_full_model/probes/profile_summary_prefill.json",
221
+ "note": "stage 05 shipped with prefill unprofiled and disclosed it as a gap; stage 06 closes it. One verified 48-layer prefill of a 128-token prompt, boundary-checked by requiring the preceding pass to be the identical sequence of op codes row for row on all four devices, plus 56 per-device tallies.",
222
+ "device_kernel_ms": 122.921,
223
+ "share_of_ttft_percent": 98.0,
224
+ "expert_sparse_matmul_percent": 61.44,
225
+ "collectives_percent": 3.13,
226
+ "sdpa_percent": 0.58
227
+ },
228
+ "watcher": "clean. TT_METAL_WATCHER=10 TT_METAL_WATCHER_DISABLE_ETH=1 pytest tests/ -m 'not models_performance_bare_metal' -q is 145 passed with zero tripped asserts (doc/optimized_full_model/logs/watcher_argmaxrows.log.gz)."
229
+ },
230
+ "stage06_context_flatness": {
231
+ "claim": "The advertised 262144-token context is now USABLE, not merely allocatable. Before stage 06 decode cost grew 1.96x between a 128-token context and a 4096-token one; after it, 1.04x. That is the difference between a context you can hold and a context you can serve.",
232
+ "token_out_ms": {
233
+ "128": 19.6925,
234
+ "1024": 19.9787,
235
+ "4096": 20.505
236
+ },
237
+ "token_out_tps_user": {
238
+ "128": 50.781,
239
+ "1024": 50.053,
240
+ "4096": 48.768
241
+ },
242
+ "token_out_ms_before_stage06": {
243
+ "128": 21.4776,
244
+ "1024": 26.1432,
245
+ "4096": 42.0623
246
+ },
247
+ "ratio_4096_over_128": 1.0413,
248
+ "ratio_4096_over_128_before_stage06": 1.9584,
249
+ "cause": "the PAGED SDPA-decode call ran at the op default, whose cost is linear in cur_pos. With k_chunk_size=256 / max_cores_per_head_batch=16 it is nearly flat. Sources: doc/optimized_full_model/probes/sdpa_sweep_confirm_bf16.json (the op) and the three perf_full_model_p*_{before,argmaxrows}.json pairs (the model).",
250
+ "measured_to_context_tokens": 4096,
251
+ "measured_to_context_basis": "END TO END, this is how deep it was actually measured: prompt 4096, generate 128, batch 1, through the real 48-layer model, 128 timed reps (doc/optimized_full_model/probes/perf_full_model_p4096_argmaxrows.json). Decode beyond a 4096-token context has NOT been run end to end at 48 layers. The 262144 claim beyond that point rests on the three narrower measurements below, and this field exists so nobody reads the flatness claim as an end-to-end measurement at the advertised context.",
252
+ "op_level_evidence_to_cur_pos": 32767,
253
+ "op_level_evidence_basis": "the shipped SDPA-decode configuration measured at the real per-die decode shapes and the real bfloat16 cache dtype out to cur_pos 32767: 74.13 us against the op default's 3545.05 us, a 47.8x, with the configured leg's PCC holding at 0.9997 where the default's has decayed to 0.9897. doc/optimized_full_model/probes/sdpa_sweep_confirm_bf16.json. This is the op, not the model.",
254
+ "in_model_pcc_to_cur_pos": 16383,
255
+ "in_model_pcc_basis": "the real multichip layer against a HuggingFace reference with a prefill-primed paged cache: PCC 0.999363 at cur_pos 16383 on the adopted path (0.999294 on the default). doc/optimized_full_model/logs/sdpa_hf_pcc_at_depth_deep.log; the 128/1024/4096 legs are in probes/sdpa_hf_pcc_at_depth.json. 32768 was attempted and the HOST was OOM-killed building the torch reference, not the device -- the same reference-side ceiling pcc_validation_note records.",
256
+ "capacity_evidence": "the full 262144-token paged KV cache is allocated and held by the running 48-layer model, with both decode traces captured and a token run through it: doc/optimized_full_model/probes/footprint_262144.json.",
257
+ "still_not_measured": "a 262144-token PREFILL through all 48 layers (prefill is single-shot, and the longest 48-layer prompt run is 4096); end-to-end decode at contexts above 4096; and any of this at batch > 1. These are coverage gaps in the evidence, not capability reductions, and they are named as limitations in doc/optimized_full_model/README.md."
258
+ },
259
+ "stage06_measured": {
260
+ "method": "doc/optimized_full_model/probes/footprint_probe.py --context 262144, re-run on the SHIPPED tree. Same probe as stage 05's, copied rather than imported so it writes beside this stage's evidence. It builds the real model -- real weights, real embed_tokens, real lm_head, the real paged KV cache at the advertised context, the real RoPE tables -- captures both decode traces and runs a token through it. GB = 1e9, as the allocator reports it.",
261
+ "weights_embed_lm_head_rope_gb_per_die": 5.311,
262
+ "kv_cache_at_262144_batch1_gb_per_die": 6.443,
263
+ "traces_and_persistent_buffers_gb_per_die": 0.006,
264
+ "total_gb_per_die": 11.76,
265
+ "headroom_gb_per_die": 22.119,
266
+ "dram_per_die_gb_reported": 33.879,
267
+ "sum_note": "the three stage rows in doc/optimized_full_model/probes/footprint_262144.json sum to 11.759906816, bit-identical to total_gb_per_die. Any residual visible when the 3dp values above are added by hand is this file's display rounding and nothing else; doc/optimized_full_model/probes/check_published_figures.py asserts the raw rows equal the raw total exactly.",
268
+ "kv_formula_check": "1 local kv head * head_dim(128) * 2 bytes * 2 tensors (K and V) = 512 B/token/layer/die, x 48 layers x 262144 tokens = 6.442450944 GB/die of KV at batch 1 -- which is what the allocator reports, so the paged cache carries no per-block overhead worth a digit. Unchanged by stage 06: the KV dtype, the block size and the head split are all untouched.",
269
+ "comparison_to_stage05": "stage 05 measured 11.759 GB/die total and 22.119 free for the same contents. Stage 06 adds one small constant to the model -- the distributed argmax's per-die vocabulary offset, a [1,1,rows,4] int32 sharded to one tile-padded column per die -- and no new activation, cache or collective buffer. Any difference between the two totals is that constant plus allocator placement, and it is far below the headroom either way.",
270
+ "rope_table_note": "the probe builds the model with the default rope_cache_len of 8192, so the RoPE tables inside weights_embed_lm_head_rope_gb_per_die are the 8192-row pair (0.004 GB/die). Growing them to the full 262144 context costs 0.134 GB/die (262144 * 128 * 2 B * 2 tables). No capability reduction; the tables grow on demand via Qwen3CoderModel.ensure_rope_capacity, and tt/generator.py's decode_forward takes decode_horizon= so a low-level caller grows them once before any trace is captured."
271
+ },
272
+ "stage06_runtime_audit": {
273
+ "source": "doc/optimized_full_model/probes/runtime_fallback_audit.json",
274
+ "host_logit_readback_on_token_out_path": false,
275
+ "host_argmax_on_token_out_path": false,
276
+ "sampling_greedy": "Sampling1D force-argmax, distributed: per-die untilize/argmax/gather -> all-gather 4 candidates -> masked-min, traced, writes tt_out_tok",
277
+ "sdpa_decode_program_config": "SDPAProgramConfig(compute_with_storage_grid_size=(11, 10), q_chunk_size=32, k_chunk_size=256, max_cores_per_head_batch=16)",
278
+ "sdpa_decode_k_chunk_clamped_at_shipped_context": false,
279
+ "sdpa_prefill_program_config_passed": "None",
280
+ "steady_state_only_replays_moved": true,
281
+ "note": "runtime_fallback_audit() itself does not yet carry the two properties stage 06 introduced -- the paged SDPA program config and the sampler's live-row count -- because adding fields would change a dict tests/test_full_model.py::test_runtime_fallback_audit_is_clean pins field by field. They are recorded here and in the JSON above instead, read off the modules that own them."
282
+ },
283
+ "stage07_note": "Stage 07 (the datatype sweep) changed NOTHING this file measures about capability. No capability reduction; advertised and supported context both remain 262144. The selected precision config moves only the two expert matmul inner block widths (experts_gate_up_in0_block_w 16 -> 64, experts_down_in0_block_w 12 -> 24). A block width is a program-config field, not a tensor dtype or shape: device_expert_bytes_per_die is identical (84 934 656 B) before and after, no allocation moves, and the KV cache dtype, layout, paging and bytes/token are untouched. Evidence: doc/datatype_sweep/README.md and work_log.md, doc/datatype_sweep/sweep_results.json.",
284
+ "stage07_kv_bfp8_candidate": {
285
+ "status": "measured, not selected -- the shipped kv_cache_dtype stays bfloat16",
286
+ "why_recorded": "This is the only candidate in the sweep that could move capacity, so the contract records what it would buy even though it was not taken. It moves capacity UPWARD (fewer bytes per token), so it could never have forced a reduction. Its first measurement was invalid -- a bfloat8_b paged cache filled from this model's bfloat16 K/V reads back as NaN, so the row scored at chance -- and the prefill writer now casts K/V to the cache dtype (tt/functional_decoder.match_cache_dtype). The numbers below are from the post-fix runs.",
287
+ "kv_cache_dtype": "bfloat8_b",
288
+ "bytes_per_token_per_layer": 1088,
289
+ "bytes_per_token_per_layer_per_die": 272,
290
+ "bytes_per_token_per_layer_per_die_bfloat16": 512,
291
+ "formula": "num_key_value_heads_per_die(1) * head_dim(128) * 1.0625 B/elem * 2 tensors (K and V); 1.0625 is bfloat8_b once each 16-element block's shared exponent byte is counted",
292
+ "device_kv_bytes_at_full_context_per_die": 3422552064,
293
+ "device_kv_bytes_at_full_context_per_die_bfloat16": 6442450944,
294
+ "device_kv_bytes_saved_per_die": 3019898880,
295
+ "device_kv_gb_at_full_context_per_die": 3.423,
296
+ "device_kv_gb_at_full_context_per_die_bfloat16": 6.442,
297
+ "share_of_the_11_759_gb_per_die_footprint_bfloat16": 0.548,
298
+ "measured_rows": {
299
+ "R19_kv_bfp8": "bfp8 KV as a delta from the stage-06 baseline: top-1 0.980 / top-5 1.000 / top-100 1.000, decode 42.29 t/s/u (-0.12% vs 42.34, inside the 0.368% band), TTFT 7842.57 ms against 3250.90",
300
+ "R28_kv_bfp8_bw64_24": "bfp8 KV on top of the selected block widths: top-1 0.980 / top-5 1.000 / top-100 1.000, decode 43.45 t/s/u (-0.21% vs the selected 43.54, inside the band), TTFT 7792.06 ms against 3287.61"
301
+ },
302
+ "verdict": "Not selected. bfp8 KV buys 3.020 GB/die of headroom at 262144 tokens -- real capacity -- but it buys NO decode throughput (both rows land inside the run-to-run band), costs one top-1 point, and costs 2.4x TTFT (7.8 s against 3.3 s on the 158-token gate prompt). The stage ranks on decode and the advertised context is already met at bfloat16 with 22.119 GB/die free, so there is no capacity problem for this to solve. It is recorded here so that a future stage needing more KV headroom -- a larger batch, or a longer served context -- has the price already measured. The TTFT cost would need explaining first; it is not the cast (it was present before the cast existed) and is unattributed.",
303
+ "evidence": [
304
+ "doc/datatype_sweep/sweep_results.json (rows R19_kv_bfp8, R28_kv_bfp8_bw64_24)",
305
+ "doc/datatype_sweep/probes/kv_bfp8_diagnosis.json (both cache writers, six cache/input combinations)",
306
+ "doc/datatype_sweep/logs/rows/R19_kv_bfp8.log, R28_kv_bfp8_bw64_24.log"
307
+ ]
308
+ },
309
+ "stage11_note": "Stage 11 (batch scaling, doc/batch_scaling/) added the variable-width decode ladder and re-validated the context bound under the SHIPPED SERVING configuration -- vLLM KV pool 263168, max_num_seqs 32, QWEN3_DECODE_WIDTHS 1,2,4,8,16,32 -- which is a different setup from the standalone batch-1 probe recorded in `measured` above. current_supported_context is unchanged at 262144 and there is still no capability reduction: `get_max_tokens_all_users` serves min(max_model_len, 262144). The fields below record exactly how far that bound has been exercised in serving, so that this file -- which tt/generator_vllm.py reads as the source of truth -- and doc/batch_scaling/README.md's limitations state the same thing rather than each carrying half of it.",
310
+ "stage11_serving_context": {
311
+ "bound_tokens": 256000,
312
+ "bound_blocks": 8000,
313
+ "bound_rationale": "Block-aligned (8000 blocks x 32) and 3,000 tokens below the largest length actually validated in serving (259,000). 259,000 is not block-aligned (8093.75 blocks); the largest aligned value at or below it is 258,976, but that leaves NO margin under a cliff whose onset is not located and whose mechanism is unresolved. 256,000 is a clean figure for a catalog, sits 6,136 tokens below the known-bad 262,136, and costs 1.2 % of the validated range.",
314
+ "validated_by": "a real vLLM server, prompt sent as exact token ids",
315
+ "evidence": [
316
+ "doc/batch_scaling/probes/long_prompt_t120_253k.json",
317
+ "doc/batch_scaling/probes/long_prompt_t120_259k.json",
318
+ "doc/batch_scaling/probes/long_prompt_t120_full.json",
319
+ "doc/batch_scaling/logs/bd_engine_len259.log"
320
+ ],
321
+ "serving_measurements": {
322
+ "253000": {
323
+ "elapsed_s": 1049.3,
324
+ "ok": true
325
+ },
326
+ "259000": {
327
+ "elapsed_s": 1130.8,
328
+ "ok": true,
329
+ "note": "O(S^2) predicted 1127 s; 0.3 % error"
330
+ },
331
+ "262136": {
332
+ "elapsed_s": null,
333
+ "ok": false,
334
+ "note": "did not complete in 3600 s; one blocking prefill; no watchdog throw"
335
+ }
336
+ },
337
+ "conditional_on_env": {
338
+ "TT_METAL_OPERATION_TIMEOUT_SECONDS": "120.0",
339
+ "why": "This bound is only reachable with the tt-metal dispatch STALL watchdog set generously. It is not an elapsed-time limit: loop_and_wait_with_timeout resets its clock whenever the device dispatch progress counter moves, and fires only after that counter has been still for the whole duration. tt-metal's own default is 0.0 (off). tt-inference-server/vllm-tt-metal/src/run_vllm_api_server.py:521 sets 5.0 unconditionally, and at 5.0 a 253,000-token prompt KILLED a deployment and required tt-smi -r. Measured: 259,000 passes at 120.0 and at 30.0, and FAILS at 10.0 (fired at 17.7 s). Recommended 120.0 -- see timeout_recommendation."
340
+ },
341
+ "timeout_recommendation": {
342
+ "value": 120.0,
343
+ "largest_legitimate_gap_bracketed_between_s": [
344
+ 10.0,
345
+ 30.0
346
+ ],
347
+ "rationale": "A genuine hang NEVER advances dispatch progress, so ANY finite timeout detects it -- a larger value costs detection LATENCY, not detection ability. The costs are therefore asymmetric: too low misdiagnoses a legitimate long op as a hang, kills the server and requires tt-smi -r (observed); too high only delays detection of a real hang. 120.0 is 4x the largest known-passing value (30.0) and still catches a genuine hang within two minutes. DO NOT disable the watchdog: this port has a documented NoC hang (PRESERVE_DECODE_TRACES) that this watchdog is the thing to catch."
348
+ },
349
+ "standalone_result_is_separate": "The 262144-token result in `measured` is STANDALONE (batch 1, direct generator calls, no vLLM) and remains correct as such. It is not a serving result and must not be quoted as one.",
350
+ "correction": "A previous version of this block published serving_validated_prefill_tokens 262080 with a 64-token gap, citing kv_pool_ceiling_sixwidth_production.json as evidence and describing a vLLM serving_config. That probe NEVER went through vLLM -- its own `what` field says 'measured through configure_paging rather than vLLM' and kv_pool_ceiling.py calls gen.prefill_forward() directly. A standalone single-user measurement was labelled serving-validated. The 4x timing discrepancy that should have exposed it was sitting in this same file unexamined: 192 s standalone at 262144 against 387 s SERVED at 168,901 tokens, i.e. serving prefill is ~4x slower per token. The real serving bound is 256000, not 262080.",
351
+ "why_our_own_runs_never_saw_this": "Every serving measurement in doc/batch_scaling/ was taken with TT_METAL_OPERATION_TIMEOUT_SECONDS UNSET, i.e. tt-metal's default 0.0 = watchdog DISABLED, because models/common/readiness_check does not set it. The TTI deployment sets 5.0. Our serving numbers were therefore measured under a different safety configuration from the one users actually deploy.",
352
+ "unresolved": {
353
+ "what": "the prefill cliff between 259,000 and 262,136 tokens",
354
+ "leading_hypothesis": "KV-pool edge. The pool is 263,168 tokens; a 262,144-token request leaves 1,024 tokens (32 blocks) while 259,000 leaves 4,168. The sharpness -- quadratic to within 0.3 % at 259,000, then >=3.5x the prediction 3.6 % later -- fits an allocation-constrained step.",
355
+ "second_hypothesis": "An L1/memory-config step function in the prefill path, analogous to the measured width-19 decode cliff where _decode_expert_memory_config flips L1->DRAM.",
356
+ "what_is_already_ruled_out": "Preemption/recompute THRASH. Engine telemetry shows ZERO loggers.py lines across the 59-minute attempt, and the scheduler runs (and logs) between steps, so a preempt/recompute loop would have emitted a line per interval. This does NOT rule out a SINGLE allocation-constrained step -- do not re-run this refutation believing it closed the question.",
357
+ "why_not_chased": "At 253,000 tokens TTFT is already 17.5 minutes, so no interactive deployment operates near the cliff; the usable range is bounded by latency far below it. Locating the mechanism would have cost 2-3 further 20-60 minute board runs for near-zero operational value."
358
+ }
359
+ }
360
+ }
code/models/demos/blackhole/qwen3_coder_30b_a3b/config/selected_precision_config.json ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation_dtype": "bfloat16",
3
+ "attention_fidelity": null,
4
+ "attention_qkv_dtype": "bfloat8_b",
5
+ "attention_wo_dtype": "bfloat8_b",
6
+ "ccl_dtype": null,
7
+ "embedding_dtype": "bfloat16",
8
+ "experts_down_dtype": "bfloat4_b",
9
+ "experts_down_in0_block_w": 24,
10
+ "experts_fidelity": "LoFi",
11
+ "experts_gate_up_dtype": "bfloat4_b",
12
+ "experts_gate_up_in0_block_w": 64,
13
+ "kv_cache_dtype": "bfloat16",
14
+ "lm_head_dtype": "bfloat8_b",
15
+ "lm_head_fidelity": "HiFi2",
16
+ "logits_dtype": "bfloat16",
17
+ "norm_fidelity": "HiFi4",
18
+ "norm_weight_dtype": "bfloat16",
19
+ "router_dtype": "bfloat16",
20
+ "router_window_fidelity": "HiFi4",
21
+ "sampling_dtype": "bfloat16"
22
+ }
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/__init__.py ADDED
File without changes
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/reference.py ADDED
@@ -0,0 +1,145 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Layer-only HuggingFace reference for Qwen3-Coder-30B-A3B-Instruct.
5
+
6
+ Loading the full 30.5B causal LM to test one decoder layer wastes ~57GB of host
7
+ RAM and several minutes, so this reads just the tensors for a single layer
8
+ straight out of the safetensors shards and populates one ``Qwen3MoeDecoderLayer``.
9
+
10
+ Checkpoint-vs-module weight layout
11
+ ----------------------------------
12
+ The checkpoint stores experts as 3 separate tensors per expert::
13
+
14
+ model.layers.L.mlp.experts.E.gate_proj.weight [moe_inter, hidden]
15
+ model.layers.L.mlp.experts.E.up_proj.weight [moe_inter, hidden]
16
+ model.layers.L.mlp.experts.E.down_proj.weight [hidden, moe_inter]
17
+
18
+ ``Qwen3MoeExperts`` instead holds them batched and gate/up fused::
19
+
20
+ gate_up_proj [num_experts, 2 * moe_inter, hidden] # [gate ; up] along dim 0
21
+ down_proj [num_experts, hidden, moe_inter]
22
+
23
+ The fusion order matters: ``Qwen3MoeExperts.forward`` does
24
+ ``linear(x, gate_up_proj[e]).chunk(2, dim=-1)``, so the FIRST half is gate and
25
+ the second is up. Concatenating in the other order silently swaps them and the
26
+ layer still runs -- it just produces wrong numbers.
27
+ """
28
+
29
+ from __future__ import annotations
30
+
31
+ import json
32
+ from collections import defaultdict
33
+
34
+ import torch
35
+ from huggingface_hub import hf_hub_download
36
+ from safetensors import safe_open
37
+ from transformers import AutoConfig
38
+ from transformers.models.qwen3_moe.modeling_qwen3_moe import Qwen3MoeDecoderLayer, Qwen3MoeRotaryEmbedding
39
+
40
+ HF_MODEL = "Qwen/Qwen3-Coder-30B-A3B-Instruct"
41
+
42
+
43
+ def load_config(hf_model: str = HF_MODEL):
44
+ return AutoConfig.from_pretrained(hf_model)
45
+
46
+
47
+ def layer_state_dict(layer_idx: int = 0, hf_model: str = HF_MODEL) -> dict[str, torch.Tensor]:
48
+ """Read only ``model.layers.<layer_idx>.*`` from the shards that hold them."""
49
+ index = json.load(open(hf_hub_download(hf_model, "model.safetensors.index.json")))["weight_map"]
50
+ prefix = f"model.layers.{layer_idx}."
51
+
52
+ per_shard: dict[str, list[str]] = defaultdict(list)
53
+ for name, shard in index.items():
54
+ if name.startswith(prefix):
55
+ per_shard[shard].append(name)
56
+ if not per_shard:
57
+ raise KeyError(f"no tensors found for layer {layer_idx}")
58
+
59
+ out: dict[str, torch.Tensor] = {}
60
+ for shard, names in per_shard.items():
61
+ path = hf_hub_download(hf_model, shard)
62
+ with safe_open(path, framework="pt") as f:
63
+ for name in names:
64
+ out[name[len(prefix) :]] = f.get_tensor(name)
65
+ return out
66
+
67
+
68
+ def build_reference_layer(layer_idx: int = 0, hf_model: str = HF_MODEL):
69
+ """Return ``(layer, config)`` with real checkpoint weights, in eval mode."""
70
+ config = load_config(hf_model)
71
+ sd = layer_state_dict(layer_idx, hf_model)
72
+
73
+ with torch.device("meta"):
74
+ layer = Qwen3MoeDecoderLayer(config, layer_idx)
75
+ layer.to_empty(device="cpu")
76
+
77
+ direct = [
78
+ "input_layernorm.weight",
79
+ "post_attention_layernorm.weight",
80
+ "self_attn.q_proj.weight",
81
+ "self_attn.k_proj.weight",
82
+ "self_attn.v_proj.weight",
83
+ "self_attn.o_proj.weight",
84
+ "self_attn.q_norm.weight",
85
+ "self_attn.k_norm.weight",
86
+ "mlp.gate.weight",
87
+ ]
88
+ params = dict(layer.named_parameters())
89
+ for key in direct:
90
+ params[key].data.copy_(sd[key])
91
+
92
+ # Fuse + stack the experts. gate first, then up -- see module docstring.
93
+ n_experts = config.num_experts
94
+ gate_up = torch.stack(
95
+ [
96
+ torch.cat([sd[f"mlp.experts.{e}.gate_proj.weight"], sd[f"mlp.experts.{e}.up_proj.weight"]], dim=0)
97
+ for e in range(n_experts)
98
+ ]
99
+ )
100
+ down = torch.stack([sd[f"mlp.experts.{e}.down_proj.weight"] for e in range(n_experts)])
101
+ params["mlp.experts.gate_up_proj"].data.copy_(gate_up)
102
+ params["mlp.experts.down_proj"].data.copy_(down)
103
+
104
+ return layer.eval(), config
105
+
106
+
107
+ def rotary_embeddings(config, seq_len: int, device="cpu"):
108
+ """Return the ``(cos, sin)`` pair the decoder layer expects."""
109
+ rope = Qwen3MoeRotaryEmbedding(config=config, device=device)
110
+ position_ids = torch.arange(seq_len, device=device).unsqueeze(0)
111
+ dummy = torch.zeros(1, seq_len, config.hidden_size, dtype=torch.float32, device=device)
112
+ return rope(dummy, position_ids)
113
+
114
+
115
+ def weight_stats(sd: dict[str, torch.Tensor]) -> dict[str, dict]:
116
+ """Per-tensor name/shape/dtype/mean/std, for deterministic synthetic weights."""
117
+ stats = {}
118
+ for name, t in sd.items():
119
+ f = t.float()
120
+ stats[name] = {
121
+ "shape": list(t.shape),
122
+ "dtype": str(t.dtype),
123
+ "mean": f.mean().item(),
124
+ "std": f.std().item(),
125
+ }
126
+ return stats
127
+
128
+
129
+ if __name__ == "__main__":
130
+ torch.manual_seed(0)
131
+
132
+ layer, config = build_reference_layer(0)
133
+ n_params = sum(p.numel() for p in layer.parameters())
134
+ print(f"layer 0 built: {n_params/1e9:.2f}B params, dtype={next(layer.parameters()).dtype}")
135
+
136
+ seq_len = 32
137
+ hidden = torch.randn(1, seq_len, config.hidden_size, dtype=torch.float32) * 0.02
138
+ cos, sin = rotary_embeddings(config, seq_len)
139
+
140
+ with torch.no_grad():
141
+ out = layer(hidden, position_embeddings=(cos, sin), attention_mask=None)
142
+ out = out[0] if isinstance(out, tuple) else out
143
+
144
+ print(f"forward OK: {tuple(hidden.shape)} -> {tuple(out.shape)}")
145
+ print(f" out mean={out.mean():.6f} std={out.std():.6f} finite={torch.isfinite(out).all().item()}")
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_attention.py ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """TTNN attention vs the HuggingFace reference layer's ``self_attn``.
5
+
6
+ The PCC test is the real check, but two cheap host-side tests run first because
7
+ they isolate failures that PCC alone reports as one undifferentiated low number:
8
+
9
+ * ``test_qk_norm_is_applied`` -- proves the per-head norm actually changed the
10
+ tensor. A QK-norm quietly skipped (wrong key, missing weight) still yields
11
+ a plausible ~0.9 PCC that looks like ordinary bf16 loss.
12
+ * ``test_causality`` -- perturbing a late token must not move an early one.
13
+ A non-causal run scores high PCC on short sequences, so ``is_causal`` being
14
+ dropped is otherwise invisible here and only surfaces as garbage generation
15
+ much later.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import pytest
21
+ import torch
22
+ from loguru import logger
23
+
24
+ import ttnn
25
+ from models.common.utility_functions import comp_allclose, comp_pcc
26
+
27
+ from ..tt.functional_decoder import AttentionConfig, attention_prefill, build_rope_cache, upload_attention_weights
28
+ from ..tt.weight_mapping import convert_attention_weights
29
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
30
+
31
+ LAYER_IDX = 0
32
+ PCC_REQUIRED = 0.99 # attention_bias=False, so no large-bias PCC degradation applies
33
+
34
+
35
+ @pytest.fixture(scope="module")
36
+ def reference():
37
+ return build_reference_layer(LAYER_IDX)
38
+
39
+
40
+ @pytest.fixture(scope="module")
41
+ def torch_weights():
42
+ return convert_attention_weights(
43
+ {k: v for k, v in layer_state_dict(LAYER_IDX).items()},
44
+ n_heads=32,
45
+ n_kv_heads=4,
46
+ head_dim=128,
47
+ )
48
+
49
+
50
+ def _hidden(config, seq_len, seed=0):
51
+ torch.manual_seed(seed)
52
+ return torch.randn(1, seq_len, config.hidden_size, dtype=torch.float32) * 0.02
53
+
54
+
55
+ def _causal_mask(seq_len):
56
+ return torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
57
+
58
+
59
+ def _reference_attention(layer, config, hidden):
60
+ seq_len = hidden.shape[1]
61
+ cos, sin = rotary_embeddings(config, seq_len)
62
+ with torch.no_grad():
63
+ out = layer.self_attn(
64
+ hidden_states=hidden,
65
+ position_embeddings=(cos, sin),
66
+ attention_mask=_causal_mask(seq_len),
67
+ )
68
+ return out[0] if isinstance(out, tuple) else out
69
+
70
+
71
+ def test_qk_norm_weights_are_non_trivial(torch_weights):
72
+ """A QK-norm weight of all ones would make the norm undetectable in PCC."""
73
+ for name in ("q_norm", "k_norm"):
74
+ w = torch_weights[name]
75
+ assert w.shape == (128,), f"{name} has shape {tuple(w.shape)}, expected (head_dim,)"
76
+ assert not torch.allclose(w, torch.ones_like(w)), f"{name} is all ones -- cannot detect a skipped norm"
77
+
78
+
79
+ def test_causality(reference):
80
+ """Changing the last token must leave earlier outputs untouched."""
81
+ layer, config = reference
82
+ hidden = _hidden(config, 32)
83
+ baseline = _reference_attention(layer, config, hidden)
84
+
85
+ perturbed = hidden.clone()
86
+ perturbed[:, -1, :] += 1.0
87
+ after = _reference_attention(layer, config, perturbed)
88
+
89
+ assert torch.allclose(baseline[:, :-1], after[:, :-1], atol=1e-5), "reference attention is not causal"
90
+
91
+
92
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
93
+ @pytest.mark.parametrize("seq_len", [32, 128, 512], ids=["s32", "s128", "s512"])
94
+ def test_attention_prefill_vs_reference(mesh_device, reference, torch_weights, seq_len):
95
+ layer, hf_config = reference
96
+ config = AttentionConfig.from_hf(hf_config)
97
+ hidden = _hidden(hf_config, seq_len)
98
+
99
+ ref_out = _reference_attention(layer, hf_config, hidden)
100
+
101
+ weights = upload_attention_weights(torch_weights, mesh_device)
102
+ cos_cache, sin_cache = build_rope_cache(hf_config, seq_len, mesh_device)
103
+
104
+ tt_in = ttnn.from_torch(
105
+ hidden.unsqueeze(0), # [1, 1, S, hidden]
106
+ dtype=ttnn.bfloat16,
107
+ layout=ttnn.TILE_LAYOUT,
108
+ device=mesh_device,
109
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
110
+ )
111
+ tt_out = attention_prefill(tt_in, weights, config, cos_cache, sin_cache)
112
+ tt_out_torch = ttnn.to_torch(tt_out).squeeze(0)
113
+
114
+ passing, pcc_message = comp_pcc(ref_out, tt_out_torch, PCC_REQUIRED)
115
+ logger.info(comp_allclose(ref_out, tt_out_torch))
116
+ logger.info(f"attention prefill seq={seq_len}: {pcc_message}")
117
+ assert passing, f"attention prefill (seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_attention_decode.py ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Single-token decode attention against the KV cache.
5
+
6
+ Decode is validated against the *prefill reference*, not against a separate
7
+ decode reference: prefill token S and decode token S must produce the same
8
+ output, because causal attention at position S sees exactly the same context
9
+ either way. That equivalence is the whole contract of a KV cache, so testing it
10
+ directly catches the errors that matter -- an off-by-one write position, RoPE
11
+ applied at the wrong index, or a cache that was never seeded by prefill.
12
+
13
+ This is also the first exercise of the Blackhole ``nlp_create_qkv_heads_decode``
14
+ DRAM bug (tt-metal #16667), which zeroes odd-indexed Q rows. The workaround
15
+ lives in ``attention_decode``; ``test_decode_q_rows_are_all_live`` is the
16
+ regression guard, because with half of Q zeroed the output is still finite,
17
+ still plausible, and still scores a deceptively high PCC.
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ import pytest
23
+ import torch
24
+ from loguru import logger
25
+
26
+ import ttnn
27
+ from models.common.utility_functions import comp_allclose, comp_pcc
28
+
29
+ from ..tt.functional_decoder import (
30
+ AttentionConfig,
31
+ attention_decode,
32
+ attention_prefill,
33
+ build_rope_cache,
34
+ create_kv_cache,
35
+ upload_attention_weights,
36
+ )
37
+ from ..tt.weight_mapping import convert_attention_weights
38
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
39
+
40
+ LAYER_IDX = 0
41
+ PCC_REQUIRED = 0.99
42
+ MAX_SEQ = 256
43
+
44
+
45
+ @pytest.fixture(scope="module")
46
+ def reference():
47
+ return build_reference_layer(LAYER_IDX)
48
+
49
+
50
+ @pytest.fixture(scope="module")
51
+ def torch_weights():
52
+ return convert_attention_weights(layer_state_dict(LAYER_IDX), n_heads=32, n_kv_heads=4, head_dim=128)
53
+
54
+
55
+ def _hidden(hf_config, seq_len, seed=0):
56
+ torch.manual_seed(seed)
57
+ return torch.randn(1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
58
+
59
+
60
+ def _reference_attention(layer, hf_config, hidden):
61
+ seq_len = hidden.shape[1]
62
+ cos, sin = rotary_embeddings(hf_config, seq_len)
63
+ mask = torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
64
+ with torch.no_grad():
65
+ out = layer.self_attn(hidden_states=hidden, position_embeddings=(cos, sin), attention_mask=mask)
66
+ return out[0] if isinstance(out, tuple) else out
67
+
68
+
69
+ def _to_device(t, mesh_device, dtype=ttnn.bfloat16):
70
+ return ttnn.from_torch(
71
+ t, dtype=dtype, layout=ttnn.TILE_LAYOUT, device=mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG
72
+ )
73
+
74
+
75
+ def _prefill_then_decode(mesh_device, hf_config, torch_weights, hidden_full, prompt_len, block_size=None):
76
+ """Prefill ``prompt_len`` tokens, then decode the token at ``prompt_len``."""
77
+ config = AttentionConfig.from_hf(hf_config)
78
+ weights = upload_attention_weights(torch_weights, mesh_device)
79
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
80
+ kv_cache = create_kv_cache(mesh_device, config, max_batch=1, max_seq_len=MAX_SEQ, block_size=block_size)
81
+
82
+ prompt = hidden_full[:, :prompt_len, :]
83
+ attention_prefill(_to_device(prompt.unsqueeze(0), mesh_device), weights, config, cos_cache, sin_cache, kv_cache)
84
+
85
+ # [1, 1, batch=1, hidden]
86
+ next_tok = hidden_full[:, prompt_len, :].reshape(1, 1, 1, hf_config.hidden_size)
87
+ current_pos = ttnn.from_torch(torch.tensor([prompt_len], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
88
+ out = attention_decode(
89
+ _to_device(next_tok, mesh_device),
90
+ weights,
91
+ config,
92
+ cos_cache,
93
+ sin_cache,
94
+ kv_cache,
95
+ current_pos,
96
+ token_index=prompt_len,
97
+ )
98
+ return ttnn.to_torch(out).reshape(1, hf_config.hidden_size)
99
+
100
+
101
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
102
+ @pytest.mark.parametrize("prompt_len", [32, 128], ids=["p32", "p128"])
103
+ @pytest.mark.parametrize("block_size", [None, 32, 64], ids=["contiguous", "paged32", "paged64"])
104
+ def test_decode_matches_prefill_at_same_position(mesh_device, reference, torch_weights, prompt_len, block_size):
105
+ layer, hf_config = reference
106
+ hidden_full = _hidden(hf_config, prompt_len + 1)
107
+
108
+ # Reference: run the whole prompt+1 through prefill and take the last row.
109
+ ref_out = _reference_attention(layer, hf_config, hidden_full)[:, prompt_len, :]
110
+
111
+ tt_out = _prefill_then_decode(mesh_device, hf_config, torch_weights, hidden_full, prompt_len, block_size)
112
+
113
+ passing, pcc_message = comp_pcc(ref_out, tt_out, PCC_REQUIRED)
114
+ logger.info(comp_allclose(ref_out, tt_out))
115
+ kind = "contiguous" if block_size is None else f"paged(block={block_size})"
116
+ logger.info(f"decode at pos {prompt_len} [{kind}]: {pcc_message}")
117
+ assert passing, f"decode at position {prompt_len} [{kind}] below {PCC_REQUIRED}: {pcc_message}"
118
+
119
+
120
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
121
+ def test_decode_q_rows_are_all_live(mesh_device, reference, torch_weights):
122
+ """Guard for tt-metal #16667: no systematic zeroing of alternating rows.
123
+
124
+ The Blackhole bug zeroes odd-indexed Q rows when the fused QKV is read from
125
+ DRAM. Rather than reach into the op, this checks the observable
126
+ consequence: the head-dim structure of the output must not show a
127
+ stripe of exactly-zero alternating entries.
128
+ """
129
+ layer, hf_config = reference
130
+ hidden_full = _hidden(hf_config, 33)
131
+ out = _prefill_then_decode(mesh_device, hf_config, torch_weights, hidden_full, 32).float()
132
+
133
+ assert torch.isfinite(out).all(), "decode produced non-finite values"
134
+ zero_fraction = (out == 0).float().mean().item()
135
+ logger.info(f"decode output zero fraction = {zero_fraction:.4f}")
136
+ assert zero_fraction < 0.1, (
137
+ f"{zero_fraction:.1%} of decode outputs are exactly zero -- looks like the "
138
+ "Blackhole nlp_create_qkv_heads_decode DRAM bug (#16667); the fused QKV "
139
+ "must be staged through L1 before the split"
140
+ )
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decode_compaction_fifo.py ADDED
@@ -0,0 +1,596 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """The async boundary the compaction un-permutation has to survive.
5
+
6
+ `Qwen3CoderForCausalLM` may decode a compacted batch in a graph narrower than
7
+ `max_num_seqs`, which means the sampled tokens come back in **graph-row** order
8
+ and have to be scattered back to vLLM's slots. The mapping that does that is
9
+ chosen when a forward is *issued*; the scatter happens when the output is
10
+ *read*. Under `--async-scheduling` those are different steps.
11
+
12
+ So the hazard is not the permutation arithmetic -- that is covered by
13
+ `doc/batch_scaling/probes/compaction_identity.py` on real weights -- it is the
14
+ **pairing**: if a later step installs a new mapping before an earlier step's
15
+ tokens are read, the earlier tokens get scattered with the wrong permutation and
16
+ every one of them lands on the wrong request. Silently, and as a correctness
17
+ bug rather than a slowdown.
18
+
19
+ The vLLM plugin happens to order this safely today (it drains pending async
20
+ decodes on a layout change, and only a layout change can move the mapping), but
21
+ that invariant lives in a repository this one must not modify and cannot pin.
22
+ `_pending_orders` removes the dependency by pairing each output with the mapping
23
+ its own forward used, and these tests are what hold that property in place.
24
+
25
+ Deliberately **device-free**: the whole hazard is adapter bookkeeping, so a fake
26
+ generator exercises it exactly and the tests run in milliseconds. Three earlier
27
+ pieces of evidence -- a real-weights identity probe, a 158-test suite and a
28
+ serving A/B -- all passed while this bug was present, because every one of them
29
+ either used the synchronous read path or never moved the mapping. That is the
30
+ gap these tests close.
31
+ """
32
+
33
+ from __future__ import annotations
34
+
35
+ from unittest.mock import patch
36
+
37
+ import pytest
38
+ import torch
39
+
40
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt import generator_vllm as gv
41
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator import Qwen3CoderGenerator
42
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator_vllm import Qwen3CoderForCausalLM
43
+
44
+ SLOTS = 32
45
+
46
+
47
+ class _Handle:
48
+ """Stands in for the device tensor a decode forward returns."""
49
+
50
+ def __init__(self, tag: int):
51
+ self.tag = tag
52
+
53
+
54
+ class _FakeGenerator:
55
+ """Only the surface `decode_forward` / `process_decode_output_host` touch.
56
+
57
+ `read_sampled_tokens` returns a vector that encodes **which forward** it came
58
+ from and **which graph row** each entry is, so a mis-paired scatter is
59
+ visible in the values rather than having to be inferred.
60
+ """
61
+
62
+ def __init__(self):
63
+ self.model = object()
64
+ self.mesh_device = object()
65
+ self.pages_per_user = 8
66
+ self.page_block_size = 32
67
+ self.num_blocks = SLOTS * 8
68
+ self.issued: list[_Handle] = []
69
+ self.trace_stats: dict = {}
70
+
71
+ # -- the calls the adapter makes on the decode path --------------------
72
+ def set_sampling_params(self, **kwargs):
73
+ return None
74
+
75
+ def set_penalty_params(self, **kwargs):
76
+ return False, False
77
+
78
+ def decode_device_state(self):
79
+ return None # first install every time: the adapter takes the host view
80
+
81
+ def prefill_forward(self, tokens, **kwargs):
82
+ # Host-sampled shape: the adapter reshapes this to [active, 1, vocab].
83
+ return torch.zeros((int(tokens.shape[0]), 8))
84
+
85
+ def decode_forward(self, *args, **kwargs):
86
+ if kwargs.get("sampling_mode") == "host":
87
+ # Host-sampled decode returns logits, not a device handle; the
88
+ # adapter reshapes them to [rows, 1, vocab].
89
+ return torch.zeros((SLOTS, 8))
90
+ handle = _Handle(len(self.issued))
91
+ self.issued.append(handle)
92
+ return handle
93
+
94
+ def read_sampled_tokens(self, tt_out, count):
95
+ # row r of forward `tag` -> 1000 * (tag + 1) + r
96
+ return torch.tensor([1000 * (tt_out.tag + 1) + r for r in range(count)], dtype=torch.long)
97
+
98
+
99
+ class _Sampling:
100
+ def __init__(self, rows: int):
101
+ self.temperature = [0.0] * rows
102
+ self.top_k = [1] * rows
103
+ self.top_p = [1.0] * rows
104
+ self.seed = [None] * rows
105
+ self.repetition_penalty = [1.0] * rows
106
+ self.presence_penalty = [0.0] * rows
107
+ self.frequency_penalty = [0.0] * rows
108
+
109
+
110
+ def _adapter(widths: str = "1,2,4,8,16,32") -> Qwen3CoderForCausalLM:
111
+ with patch.dict("os.environ", {"QWEN3_DECODE_WIDTHS": widths}):
112
+ adapter = Qwen3CoderForCausalLM(_FakeGenerator(), max_model_len=4096, max_num_seqs=SLOTS)
113
+ adapter.kv_cache = ["fake-cache"]
114
+ return adapter
115
+
116
+
117
+ def _batch(live_rows):
118
+ """vLLM's padded decode batch: position -1 on every unoccupied slot."""
119
+ positions = torch.full((SLOTS,), -1, dtype=torch.int64)
120
+ tokens = torch.zeros((SLOTS, 1), dtype=torch.int64)
121
+ for row in live_rows:
122
+ positions[row] = 128
123
+ tokens[row, 0] = 5000 + row
124
+ return tokens, positions
125
+
126
+
127
+ def _issue(adapter, live_rows, *, reset=True):
128
+ """One decode forward, output deliberately NOT read (the async path)."""
129
+ tokens, positions = _batch(live_rows)
130
+ return adapter.decode_forward(
131
+ tokens=tokens,
132
+ page_table=torch.zeros((SLOTS, 8), dtype=torch.int32),
133
+ kv_cache=adapter.kv_cache,
134
+ start_pos=positions,
135
+ sampling_params=_Sampling(SLOTS),
136
+ reset_batch=reset,
137
+ read_from_device=False,
138
+ )
139
+
140
+
141
+ def _read(adapter, handle):
142
+ with patch.object(gv.ttnn, "is_tensor_storage_on_device", lambda _t: False):
143
+ return adapter.process_decode_output_host(handle, is_tokens=True).reshape(-1)
144
+
145
+
146
+ def test_output_uses_the_mapping_its_own_forward_was_issued_with():
147
+ """Issue A, then issue B with a different mapping, then read A.
148
+
149
+ This is the exact interleaving `--async-scheduling` produces and the one
150
+ nothing else covers. With the mapping stored on the adapter rather than
151
+ queued, reading A after B scatters A's tokens through **B's** permutation.
152
+ """
153
+ adapter = _adapter()
154
+
155
+ handle_a = _issue(adapter, [3, 17, 29]) # width 4, order starts 3,17,29
156
+ order_a = adapter._compaction.clone()
157
+ handle_b = _issue(adapter, [0, 1, 2, 3, 4, 5]) # width 8, a different order
158
+ order_b = adapter._compaction.clone()
159
+
160
+ # The premise of the test: the adapter's live mapping is no longer A's.
161
+ assert not torch.equal(
162
+ order_a[: min(len(order_a), len(order_b))], order_b[: min(len(order_a), len(order_b))]
163
+ ), "the two steps must have different mappings or this test proves nothing"
164
+
165
+ tokens_a = _read(adapter, handle_a)
166
+ # Forward 0, graph rows 0,1,2 -> vLLM slots 3,17,29.
167
+ assert tokens_a[3] == 1000, tokens_a[[3, 17, 29]]
168
+ assert tokens_a[17] == 1001, tokens_a[[3, 17, 29]]
169
+ assert tokens_a[29] == 1002, tokens_a[[3, 17, 29]]
170
+
171
+ tokens_b = _read(adapter, handle_b)
172
+ # Forward 1, contiguous live rows -> identity over the first six slots.
173
+ for row in range(6):
174
+ assert tokens_b[row] == 2000 + row, tokens_b[:6]
175
+
176
+ assert adapter._audit["compaction_fifo_underflows"] == 0
177
+ assert len(adapter._pending_orders) == 0
178
+
179
+
180
+ def test_three_forwards_in_flight_are_read_in_issue_order():
181
+ """FIFO depth > 2, so the pairing cannot be a lucky one-slot swap."""
182
+ adapter = _adapter()
183
+ handles = [_issue(adapter, rows) for rows in ([3, 17, 29], [0, 1], [7])]
184
+ assert adapter._audit["compaction_fifo_max_depth"] == 3
185
+
186
+ expected = ({3: 1000, 17: 1001, 29: 1002}, {0: 2000, 1: 2001}, {7: 3000})
187
+ for handle, wanted in zip(handles, expected):
188
+ got = _read(adapter, handle)
189
+ for slot, value in wanted.items():
190
+ assert got[slot] == value, (slot, value, got[slot])
191
+ assert adapter._audit["compaction_fifo_underflows"] == 0
192
+
193
+
194
+ def test_full_width_steps_queue_a_null_mapping():
195
+ """At full occupancy there is no permutation, and that must still be paired.
196
+
197
+ A `None` entry is meaningful: it says "this step needs no un-permutation".
198
+ Skipping the push for full-width steps would misalign the queue for every
199
+ narrow step behind them.
200
+ """
201
+ adapter = _adapter()
202
+ handle_full = _issue(adapter, list(range(SLOTS)))
203
+ assert adapter._compaction is None
204
+ handle_narrow = _issue(adapter, [9])
205
+ assert adapter._compaction is not None
206
+
207
+ tokens_full = _read(adapter, handle_full)
208
+ for row in (0, 5, 31):
209
+ assert tokens_full[row] == 1000 + row, tokens_full[:3]
210
+
211
+ tokens_narrow = _read(adapter, handle_narrow)
212
+ assert tokens_narrow[9] == 2000, tokens_narrow[:3]
213
+ assert adapter._audit["compaction_fifo_underflows"] == 0
214
+
215
+
216
+ def _unset_widths_adapter() -> Qwen3CoderForCausalLM:
217
+ """An adapter built with `QWEN3_DECODE_WIDTHS` genuinely absent."""
218
+ import os
219
+
220
+ with patch.dict("os.environ", {}, clear=False):
221
+ os.environ.pop("QWEN3_DECODE_WIDTHS", None)
222
+ adapter = Qwen3CoderForCausalLM(_FakeGenerator(), max_model_len=4096, max_num_seqs=SLOTS)
223
+ adapter.kv_cache = ["fake-cache"]
224
+ return adapter
225
+
226
+
227
+ def test_the_ladder_is_on_when_the_variable_is_unset():
228
+ """Unset must mean the ladder, not the fixed-width graph.
229
+
230
+ The previous default was off, and it was *known wrong*: a `max_num_seqs=32`
231
+ server that simply does not set an environment variable decodes one user at
232
+ 4.3464 t/s/u instead of 49.3636. This test is what stops that default coming
233
+ back by accident.
234
+ """
235
+ adapter = _unset_widths_adapter()
236
+
237
+ assert adapter._decode_widths == [1, 2, 4, 8, 16, SLOTS]
238
+ assert adapter._compaction_enabled is True
239
+ handle = _issue(adapter, [3, 17, 29])
240
+ assert adapter._compaction is not None, "unset must still compact"
241
+ assert len(adapter._pending_orders) == 1
242
+ tokens = _read(adapter, handle)
243
+ # Un-permuted back to vLLM slots, exactly as with the ladder set explicitly.
244
+ for row in (3, 17, 29):
245
+ assert tokens[row] == 1000 + [3, 17, 29].index(row), tokens[[3, 17, 29]]
246
+
247
+
248
+ def test_a_single_width_restores_the_fixed_width_path_exactly():
249
+ """`QWEN3_DECODE_WIDTHS=32` is the escape hatch back to the old behaviour.
250
+
251
+ With one width there is nothing to compact, so the path takes none of the
252
+ bookkeeping and cannot raise any of the pairing errors below -- which is the
253
+ property the previous default provided and which must remain reachable.
254
+ """
255
+ adapter = _adapter(widths=str(SLOTS))
256
+
257
+ assert adapter._decode_widths == [SLOTS]
258
+ assert adapter._compaction_enabled is False
259
+ handle = _issue(adapter, [3, 17, 29])
260
+ assert adapter._compaction is None, "no compaction may happen with a single width"
261
+ assert len(adapter._pending_orders) == 0, "the single-width path must not touch the queue"
262
+ tokens = _read(adapter, handle)
263
+ # Straight through: graph row r is slot r.
264
+ for row in (0, 3, 17, 29):
265
+ assert tokens[row] == 1000 + row, tokens[[0, 3, 17, 29]]
266
+
267
+
268
+ def test_widths_above_max_num_seqs_are_dropped_from_the_default():
269
+ """A smaller server must not try to capture a graph wider than its slots."""
270
+ import os
271
+
272
+ with patch.dict("os.environ", {}, clear=False):
273
+ os.environ.pop("QWEN3_DECODE_WIDTHS", None)
274
+ adapter = Qwen3CoderForCausalLM(_FakeGenerator(), max_model_len=4096, max_num_seqs=8)
275
+
276
+ assert adapter._decode_widths == [1, 2, 4, 8]
277
+
278
+
279
+ @pytest.mark.parametrize("live_rows", [[0], [3, 17, 29], list(range(16)), list(range(SLOTS))])
280
+ def test_queue_drains_exactly_once_per_forward(live_rows):
281
+ adapter = _adapter()
282
+ handle = _issue(adapter, live_rows)
283
+ assert len(adapter._pending_orders) == 1
284
+ _read(adapter, handle)
285
+ assert len(adapter._pending_orders) == 0
286
+ assert adapter._audit["compaction_fifo_underflows"] == 0
287
+
288
+
289
+ # -- the guards -------------------------------------------------------------
290
+ #
291
+ # The pairing rests on decode forwards being finalized exactly once and in issue
292
+ # order. Those are invariants of `vllm-tt-plugin`, which this repository must not
293
+ # modify, so they are *checked* here rather than trusted. The route that makes
294
+ # this concrete: `async_decode.py::ensure_finalized` sets `_finalized = True`
295
+ # only after `_get_output_impl()` returns, and the pop happens inside that call
296
+ # -- so a raise anywhere after the pop leaves the step un-finalized and a later
297
+ # `wait_for_all_pending_async_steps` finalizes it a second time.
298
+
299
+
300
+ def test_underflow_raises_instead_of_guessing(expect_error):
301
+ """No queued mapping means the pairing is broken; a wrong guess is worse.
302
+
303
+ Falling back to the adapter's current mapping here would apply exactly the
304
+ permutation the queue exists to prevent, in the one state where it is known
305
+ not to belong to these tokens -- every token to the wrong request, silently.
306
+ """
307
+ adapter = _adapter()
308
+ handle = _issue(adapter, [3, 17, 29])
309
+ _read(adapter, handle)
310
+ with expect_error(RuntimeError, "no queued row mapping"):
311
+ _read(adapter, handle) # the second finalize of the same forward
312
+ assert adapter._audit["compaction_fifo_underflows"] == 1
313
+
314
+
315
+ def test_double_finalize_is_caught_by_the_tag_before_it_can_mis_scatter(expect_error):
316
+ """A re-finalize consumes the *next* step's mapping; the tag says so.
317
+
318
+ This is the overflow-direction desync: the queue drains faster than it
319
+ fills. Depth alone cannot see it -- the tags can.
320
+ """
321
+ adapter = _adapter()
322
+ a = _issue(adapter, [3, 17, 29])
323
+ _issue(adapter, [0, 1])
324
+ _read(adapter, a) # legitimate: pops tag 0
325
+ with expect_error(RuntimeError, "out of step"):
326
+ # A second finalize of forward A pops tag 1, which belongs to forward B.
327
+ # Without the tag it would silently scatter A's tokens through B's map.
328
+ adapter._pending_orders.appendleft((99, adapter._compaction))
329
+ _read(adapter, a)
330
+
331
+
332
+ def test_queue_cap_raises_rather_than_growing_without_bound(expect_error):
333
+ """Outputs never read is a leak; fail at the cap instead of mis-pairing later."""
334
+ adapter = _adapter()
335
+ adapter._pending_orders_cap = 4
336
+ with expect_error(RuntimeError, "row-mapping queue reached"):
337
+ for _ in range(adapter._pending_orders_cap + 2):
338
+ _issue(adapter, [3, 17, 29])
339
+
340
+
341
+ def test_reset_realigns_tags_so_later_pops_still_pair():
342
+ """A released trace makes queued outputs unreadable; the reset must not desync.
343
+
344
+ Clearing without realigning the tags would make the next legitimate pop look
345
+ like a skipped step and raise on a perfectly healthy server.
346
+ """
347
+ adapter = _adapter()
348
+ _issue(adapter, [3, 17, 29])
349
+ _issue(adapter, [0, 1])
350
+ assert len(adapter._pending_orders) == 2
351
+ adapter._reset_pending_orders()
352
+ assert len(adapter._pending_orders) == 0
353
+
354
+ handle = _issue(adapter, [7])
355
+ got = _read(adapter, handle)
356
+ assert got[7] == 3000, got[:8] # third forward issued -> tag 1000*(2+1)
357
+ assert adapter._audit["compaction_fifo_underflows"] == 0
358
+
359
+
360
+ def test_prefill_resets_the_queue():
361
+ """Prefill may release the decode traces, so queued outputs die with them."""
362
+ adapter = _adapter()
363
+ _issue(adapter, [3, 17, 29])
364
+ assert len(adapter._pending_orders) == 1
365
+ adapter.prefill_forward(
366
+ tokens=torch.zeros((1, 8), dtype=torch.int64),
367
+ page_table=torch.zeros((SLOTS, 8), dtype=torch.int32),
368
+ kv_cache=adapter.kv_cache,
369
+ prompt_lens=[8],
370
+ sampling_params=None,
371
+ )
372
+ assert len(adapter._pending_orders) == 0
373
+ handle = _issue(adapter, [5])
374
+ _read(adapter, handle)
375
+ assert adapter._audit["compaction_fifo_underflows"] == 0
376
+
377
+
378
+ # ---------------------------------------------------------------------------
379
+ # The penalty path under the ladder.
380
+ #
381
+ # `_apply_penalties` reorders the token history into graph-row order alongside
382
+ # the per-row penalty scalars. The scalars are genuine python lists, so a list
383
+ # comprehension is right for them; the histories are vLLM `[rows, L]` **torch
384
+ # tensors**, and rebuilding one as a list of 1-D tensors makes
385
+ # `Qwen3CoderGenerator._row_token_ids` raise `TypeError: only integer tensors of
386
+ # a single element can be converted to an index` on its `torch.as_tensor` call.
387
+ #
388
+ # That crashed a real penalised request in production. Nothing caught it: the
389
+ # ladder must be ON (with it off the history is passed through untouched) *and*
390
+ # the request must carry a non-neutral penalty, and `_FakeGenerator` above
391
+ # accepts `set_penalty_params(**kwargs)` without ever looking at the history.
392
+ # These tests run the **real** `_row_token_ids` over whatever the adapter
393
+ # actually passed, so a type that the generator cannot consume fails here.
394
+ # ---------------------------------------------------------------------------
395
+
396
+
397
+ class _PenaltyRecordingGenerator(_FakeGenerator):
398
+ """Captures the kwargs `_apply_penalties` hands to the generator."""
399
+
400
+ def __init__(self):
401
+ super().__init__()
402
+ self.penalty_calls: list[dict] = []
403
+
404
+ def set_penalty_params(self, **kwargs):
405
+ self.penalty_calls.append(kwargs)
406
+ return False, False
407
+
408
+
409
+ def _penalty_adapter(widths: str = "1,2,4,8,16,32") -> Qwen3CoderForCausalLM:
410
+ with patch.dict("os.environ", {"QWEN3_DECODE_WIDTHS": widths}):
411
+ adapter = Qwen3CoderForCausalLM(_PenaltyRecordingGenerator(), max_model_len=4096, max_num_seqs=SLOTS)
412
+ adapter.kv_cache = ["fake-cache"]
413
+ return adapter
414
+
415
+
416
+ def _history(rows: int, width: int = 6) -> torch.Tensor:
417
+ """A vLLM `[rows, L]` history: slot r holds tokens 100*r+1.., -1 padded.
418
+
419
+ The -1 padding and the batch padded to `max_num_reqs` are what
420
+ `_row_token_ids` documents, so this mirrors the real contract.
421
+ """
422
+ hist = torch.full((rows, width), -1, dtype=torch.int64)
423
+ for row in range(rows):
424
+ hist[row, :3] = torch.tensor([100 * row + 1, 100 * row + 2, 100 * row + 3])
425
+ return hist
426
+
427
+
428
+ def _penalised(rows: int) -> _Sampling:
429
+ sampling = _Sampling(rows)
430
+ sampling.repetition_penalty = [1.2] * rows # non-neutral: takes the staged path
431
+ return sampling
432
+
433
+
434
+ def test_penalised_decode_under_the_ladder_does_not_crash_on_a_tensor_history():
435
+ """The production crash, reduced.
436
+
437
+ A `[rows, L]` tensor history reordered into a python list of 1-D tensors
438
+ reaches `_row_token_ids` as something `torch.as_tensor` cannot index.
439
+ """
440
+ adapter = _penalty_adapter()
441
+ live = [3, 17, 29]
442
+ order = adapter._compaction_order(_batch(live)[1], 4, SLOTS)
443
+
444
+ adapter._apply_penalties(_penalised(SLOTS), SLOTS, _history(SLOTS), _history(SLOTS), order=order, graph_rows=4)
445
+
446
+ call = adapter.generator.penalty_calls[-1]
447
+ for name in ("prompt_tokens", "output_tokens"):
448
+ # The real consumer, on the real object the adapter passed.
449
+ ids = Qwen3CoderGenerator._row_token_ids(call[name], 0)
450
+ assert ids.numel() == 3, f"{name} row 0 unreadable by the generator"
451
+
452
+
453
+ def test_penalty_history_follows_its_own_slot_through_the_compaction():
454
+ """Graph row g must carry the history of the slot the mapping sent there.
455
+
456
+ A reorder that is merely type-correct but inverted would apply one user's
457
+ repetition penalty to another user's tokens -- wrong output, no crash.
458
+ """
459
+ adapter = _penalty_adapter()
460
+ live = [3, 17, 29]
461
+ order = adapter._compaction_order(_batch(live)[1], 4, SLOTS)
462
+
463
+ adapter._apply_penalties(_penalised(SLOTS), SLOTS, _history(SLOTS), _history(SLOTS), order=order, graph_rows=4)
464
+
465
+ call = adapter.generator.penalty_calls[-1]
466
+ for graph_row, slot in enumerate(int(v) for v in order.tolist()):
467
+ ids = Qwen3CoderGenerator._row_token_ids(call["prompt_tokens"], graph_row)
468
+ expected = torch.tensor([100 * slot + 1, 100 * slot + 2, 100 * slot + 3])
469
+ assert torch.equal(ids, expected), f"graph row {graph_row} carries slot {slot}'s history"
470
+
471
+
472
+ def test_penalty_scalars_and_history_are_reordered_the_same_way():
473
+ """The scalar and the history for one slot must not come apart.
474
+
475
+ They are reordered by separate statements; if only one of them tracked the
476
+ mapping, a user would get another user's penalty strength.
477
+ """
478
+ adapter = _penalty_adapter()
479
+ live = [3, 17, 29]
480
+ order = adapter._compaction_order(_batch(live)[1], 4, SLOTS)
481
+
482
+ sampling = _Sampling(SLOTS)
483
+ # A distinct penalty per slot, so a mis-pairing is visible in the value.
484
+ sampling.repetition_penalty = [1.0 + 0.01 * r for r in range(SLOTS)]
485
+
486
+ adapter._apply_penalties(sampling, SLOTS, _history(SLOTS), None, order=order, graph_rows=4)
487
+
488
+ call = adapter.generator.penalty_calls[-1]
489
+ for graph_row, slot in enumerate(int(v) for v in order.tolist()):
490
+ assert call["repetition"][graph_row] == pytest.approx(1.0 + 0.01 * slot)
491
+ ids = Qwen3CoderGenerator._row_token_ids(call["prompt_tokens"], graph_row)
492
+ assert int(ids[0]) == 100 * slot + 1, "history and scalar disagree about the slot"
493
+
494
+
495
+ @pytest.mark.parametrize("as_list", [False, True])
496
+ def test_history_reorder_preserves_the_type_it_was_given(as_list):
497
+ """Tensors stay tensors; genuine python sequences keep working.
498
+
499
+ `_row_token_ids` accepts both, but only via `torch.as_tensor`, which is what
500
+ the list-of-tensors form breaks.
501
+ """
502
+ adapter = _penalty_adapter()
503
+ live = [3, 17, 29]
504
+ order = adapter._compaction_order(_batch(live)[1], 4, SLOTS)
505
+
506
+ tensor_history = _history(SLOTS)
507
+ history = tensor_history.tolist() if as_list else tensor_history
508
+
509
+ adapter._apply_penalties(_penalised(SLOTS), SLOTS, history, None, order=order, graph_rows=4)
510
+
511
+ passed = adapter.generator.penalty_calls[-1]["prompt_tokens"]
512
+ if as_list:
513
+ assert isinstance(passed, list)
514
+ else:
515
+ assert isinstance(passed, torch.Tensor), "a tensor history must stay a tensor"
516
+ ids = Qwen3CoderGenerator._row_token_ids(passed, 0)
517
+ assert int(ids[0]) == 100 * int(order[0]) + 1
518
+
519
+
520
+ def test_ladder_off_leaves_the_history_exactly_as_vllm_sent_it():
521
+ """The shipped default path must not be touched by any of the above."""
522
+ adapter = _penalty_adapter(widths=str(SLOTS))
523
+ history = _history(SLOTS)
524
+
525
+ adapter._apply_penalties(_penalised(SLOTS), SLOTS, history, None, order=None, graph_rows=None)
526
+
527
+ passed = adapter.generator.penalty_calls[-1]["prompt_tokens"]
528
+ assert passed is history, "with the ladder off the history is passed through unchanged"
529
+
530
+
531
+ # ---------------------------------------------------------------------------
532
+ # The host-sampling demotion has to be audible.
533
+ #
534
+ # vLLM decides per request, in `check_perform_device_sampling`, whether a
535
+ # request may sample on device. On this 4-die mesh any request carrying
536
+ # `logprobs` -- including `logprobs: 0`, because the guard tests
537
+ # `max_num_logprobs is not None` before it ever looks at the value -- is routed
538
+ # to eager host sampling. That bypasses the captured trace and the width ladder
539
+ # and costs ~14x (measured: 3.595 t/s/u against 49.345), and the plugin emits no
540
+ # log line for it. The server-level `sample_on_device_mode: all` stays correct
541
+ # and stays silent.
542
+ #
543
+ # A 14x cliff whose only symptom is "the model got slow" is precisely the
544
+ # failure this port was first reported with. The adapter cannot prevent the
545
+ # demotion -- the guard lives in a repository this one must not modify -- so the
546
+ # least it must do is say so.
547
+ # ---------------------------------------------------------------------------
548
+
549
+
550
+ def _host_sampled_step(adapter):
551
+ """One decode step with `sampling_params=None`, i.e. vLLM's host-sampled route."""
552
+ tokens, positions = _batch([3, 17, 29])
553
+ return adapter.decode_forward(
554
+ tokens=tokens,
555
+ page_table=torch.zeros((SLOTS, 8), dtype=torch.int32),
556
+ kv_cache=adapter.kv_cache,
557
+ start_pos=positions,
558
+ sampling_params=None,
559
+ reset_batch=True,
560
+ read_from_device=False,
561
+ )
562
+
563
+
564
+ def test_host_sampled_decode_warns_once_with_the_cause_and_the_cost():
565
+ """The demotion must name what happened, why, and what it costs."""
566
+ adapter = _adapter()
567
+ with patch.object(gv.logger, "warning") as warn:
568
+ _host_sampled_step(adapter)
569
+
570
+ assert warn.call_count == 1, "the demotion must be reported"
571
+ message = warn.call_args[0][0]
572
+ for needle in ("HOST sampling", "logprobs", "3.595", "49.345", "14x"):
573
+ assert needle in message, f"the warning must mention {needle!r}"
574
+
575
+ assert adapter._audit["host_sampled_decode_steps"] == 1
576
+
577
+
578
+ def test_host_sampled_warning_does_not_repeat_every_step():
579
+ """One line, not one per token -- a per-step warning would be its own defect."""
580
+ adapter = _adapter()
581
+ with patch.object(gv.logger, "warning") as warn:
582
+ for _ in range(5):
583
+ _host_sampled_step(adapter)
584
+
585
+ assert warn.call_count == 1, "the warning must be once per server, not per step"
586
+ assert adapter._audit["host_sampled_decode_steps"] == 5
587
+
588
+
589
+ def test_device_sampled_steps_do_not_warn():
590
+ """The traced path must stay silent; a false alarm here trains people to ignore it."""
591
+ adapter = _adapter()
592
+ with patch.object(gv.logger, "warning") as warn:
593
+ _issue(adapter, [3, 17, 29])
594
+
595
+ assert warn.call_count == 0
596
+ assert adapter._audit["host_sampled_decode_steps"] == 0
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decoder_layer.py ADDED
@@ -0,0 +1,176 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """The composed decoder layer against the HuggingFace reference layer.
5
+
6
+ This is the stage-01 deliverable: norm -> attention -> residual -> norm -> MoE
7
+ -> residual, end to end, at PCC >= 0.995.
8
+
9
+ The submodules already pass on their own (attention 0.9994+, MoE 0.9981+), so
10
+ a shortfall here is a *composition* error -- a residual added in the wrong
11
+ place, the router fed the un-normed tensor, the two norms swapped -- rather
12
+ than accumulated precision. ``test_residual_path_is_present`` exists to
13
+ separate those two explanations: it checks the layer's output actually depends
14
+ on the residual stream, which is the composition mistake most likely to still
15
+ score respectably on PCC.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import pytest
21
+ import torch
22
+ from loguru import logger
23
+
24
+ import ttnn
25
+ from models.common.utility_functions import comp_allclose, comp_pcc
26
+
27
+ from ..tt.functional_decoder import (
28
+ DecoderLayerConfig,
29
+ build_expert_sparsity,
30
+ build_rope_cache,
31
+ decoder_layer_prefill,
32
+ upload_layer_weights,
33
+ )
34
+ from ..tt.weight_mapping import convert_layer_weights
35
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
36
+
37
+ LAYER_IDX = 0
38
+ PCC_REQUIRED = 0.995 # the stage-01 functional-decoder bar
39
+
40
+
41
+ @pytest.fixture(scope="module")
42
+ def reference():
43
+ return build_reference_layer(LAYER_IDX)
44
+
45
+
46
+ @pytest.fixture(scope="module")
47
+ def torch_weights(reference):
48
+ _, hf_config = reference
49
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
50
+
51
+
52
+ def _hidden(hf_config, seq_len, seed=0):
53
+ torch.manual_seed(seed)
54
+ return torch.randn(1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
55
+
56
+
57
+ def _causal_mask(seq_len):
58
+ return torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
59
+
60
+
61
+ def _reference_layer(layer, hf_config, hidden):
62
+ seq_len = hidden.shape[1]
63
+ cos, sin = rotary_embeddings(hf_config, seq_len)
64
+ with torch.no_grad():
65
+ out = layer(
66
+ hidden,
67
+ position_embeddings=(cos, sin),
68
+ attention_mask=_causal_mask(seq_len),
69
+ )
70
+ return out[0] if isinstance(out, tuple) else out
71
+
72
+
73
+ def _run_layer(mesh_device, hf_config, torch_weights, hidden):
74
+ config = DecoderLayerConfig.from_hf(hf_config)
75
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
76
+ cos_cache, sin_cache = build_rope_cache(hf_config, hidden.shape[1], mesh_device)
77
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
78
+
79
+ tt_in = ttnn.from_torch(
80
+ hidden.unsqueeze(0),
81
+ dtype=ttnn.bfloat16,
82
+ layout=ttnn.TILE_LAYOUT,
83
+ device=mesh_device,
84
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
85
+ )
86
+ tt_out = decoder_layer_prefill(tt_in, weights, config, cos_cache, sin_cache, sparsity)
87
+ return ttnn.to_torch(tt_out).squeeze(0)
88
+
89
+
90
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
91
+ @pytest.mark.parametrize(
92
+ "seq_len",
93
+ [32, 128, 512, 33, 100, 257],
94
+ ids=["s32", "s128", "s512", "s33", "s100", "s257"],
95
+ )
96
+ def test_decoder_layer_vs_reference(mesh_device, reference, torch_weights, seq_len):
97
+ """Tile-aligned and deliberately non-aligned sequence lengths.
98
+
99
+ 33, 100 and 257 exercise the zero-padding in ``moe_prefill``: one row past a
100
+ tile, a mid-tile length, and one past a large power of two. Real prompts are
101
+ almost never a multiple of 32, and padding bugs typically corrupt only the
102
+ tail tokens -- which a sequence-wide PCC can absorb, so these run as
103
+ separate cases rather than being folded into the aligned ones.
104
+ """
105
+ layer, hf_config = reference
106
+ hidden = _hidden(hf_config, seq_len)
107
+
108
+ ref_out = _reference_layer(layer, hf_config, hidden)
109
+ tt_out = _run_layer(mesh_device, hf_config, torch_weights, hidden)
110
+
111
+ passing, pcc_message = comp_pcc(ref_out, tt_out, PCC_REQUIRED)
112
+ logger.info(comp_allclose(ref_out, tt_out))
113
+ logger.info(f"decoder layer seq={seq_len}: {pcc_message}")
114
+ assert passing, f"decoder layer (seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
115
+
116
+
117
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
118
+ @pytest.mark.parametrize("seq_len", [33, 100], ids=["s33", "s100"])
119
+ def test_non_aligned_tail_tokens(mesh_device, reference, torch_weights, seq_len):
120
+ """Tokens near the pad boundary must be no worse than the rest of the sequence.
121
+
122
+ Stated *relatively*, on purpose. An absolute per-token bar cannot tell
123
+ "zero-padding corrupted the tail" apart from "this token was always noisy":
124
+ a couple of tokens sit near 0.9946 regardless of length because their router
125
+ top-8 contains a near-tie, and they stay low at seq_len=32 where no padding
126
+ exists at all. Comparing the tail against the sequence's own distribution
127
+ isolates the padding question, which is the only thing this test is for.
128
+ """
129
+ layer, hf_config = reference
130
+ hidden = _hidden(hf_config, seq_len)
131
+
132
+ ref_out = _reference_layer(layer, hf_config, hidden)
133
+ tt_out = _run_layer(mesh_device, hf_config, torch_weights, hidden)
134
+
135
+ def token_pcc(pos):
136
+ pair = torch.stack([ref_out[:, pos, :].flatten().float(), tt_out[:, pos, :].flatten().float()])
137
+ return float(torch.corrcoef(pair)[0, 1])
138
+
139
+ per_token = [token_pcc(p) for p in range(seq_len)]
140
+ tail = per_token[-3:]
141
+ body_worst = min(per_token[:-3])
142
+
143
+ logger.info(
144
+ f"seq={seq_len}: tail={[round(v, 5) for v in tail]} "
145
+ f"body_worst={body_worst:.5f} median={sorted(per_token)[len(per_token) // 2]:.5f}"
146
+ )
147
+
148
+ assert min(tail) >= 0.99, f"tail tokens of seq_len {seq_len} are outright wrong: {tail}"
149
+ # Padding, if broken, would make the tail distinctly worse than the body.
150
+ assert min(tail) >= body_worst - 1e-3, (
151
+ f"tail tokens ({min(tail):.5f}) are worse than the worst body token "
152
+ f"({body_worst:.5f}) at seq_len {seq_len} -- suspect the moe_prefill zero-padding"
153
+ )
154
+
155
+
156
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
157
+ def test_residual_path_is_present(mesh_device, reference, torch_weights):
158
+ """The output must track the input, and must not merely echo it.
159
+
160
+ A dropped residual still produces sane-looking activations; a layer that
161
+ returns its input unchanged does too. Both are composition bugs that PCC
162
+ against the reference would report as a single vague number.
163
+ """
164
+ layer, hf_config = reference
165
+ hidden = _hidden(hf_config, 32)
166
+ out = _run_layer(mesh_device, hf_config, torch_weights, hidden).float()
167
+ flat_in = hidden.squeeze(0).float()
168
+
169
+ assert torch.isfinite(out).all(), "layer produced non-finite values"
170
+ assert not torch.allclose(out, flat_in, atol=1e-3), "output equals input -- layer body is a no-op"
171
+
172
+ # With the residual intact the output stays correlated with the input;
173
+ # without it the sublayer outputs alone would decorrelate.
174
+ corr = torch.corrcoef(torch.stack([out.flatten(), flat_in.flatten()]))[0, 1]
175
+ logger.info(f"corr(output, input) = {corr:.4f}")
176
+ assert corr > 0.5, f"output barely tracks input (corr={corr:.4f}) -- residual likely dropped"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decoder_layer_decode.py ADDED
@@ -0,0 +1,155 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """The full decoder layer in decode mode: prefill a prompt, then step tokens.
5
+
6
+ Validated against the prefill reference at the same absolute position, which is
7
+ the KV cache's defining contract -- attending to a cached prompt must equal
8
+ attending to it inline.
9
+
10
+ ``test_multi_step_decode`` matters more than the single-step case. One step can
11
+ pass while the cache is subtly broken (a write that lands on the position being
12
+ read this turn still looks right); errors in the write position only diverge
13
+ once a later token has to read what an earlier step wrote. Three consecutive
14
+ steps against a three-token-longer reference catches that.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import pytest
20
+ import torch
21
+ from loguru import logger
22
+
23
+ import ttnn
24
+ from models.common.utility_functions import comp_allclose, comp_pcc
25
+
26
+ from ..tt.functional_decoder import (
27
+ DecoderLayerConfig,
28
+ build_expert_sparsity,
29
+ build_rope_cache,
30
+ create_kv_cache,
31
+ decoder_layer_decode,
32
+ decoder_layer_prefill,
33
+ upload_layer_weights,
34
+ )
35
+ from ..tt.weight_mapping import convert_layer_weights
36
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
37
+
38
+ LAYER_IDX = 0
39
+ PCC_REQUIRED = 0.99
40
+ MAX_SEQ = 256
41
+
42
+
43
+ @pytest.fixture(scope="module")
44
+ def reference():
45
+ return build_reference_layer(LAYER_IDX)
46
+
47
+
48
+ @pytest.fixture(scope="module")
49
+ def torch_weights(reference):
50
+ _, hf_config = reference
51
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
52
+
53
+
54
+ def _hidden(hf_config, seq_len, seed=0):
55
+ torch.manual_seed(seed)
56
+ return torch.randn(1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
57
+
58
+
59
+ def _reference_layer(layer, hf_config, hidden):
60
+ seq_len = hidden.shape[1]
61
+ cos, sin = rotary_embeddings(hf_config, seq_len)
62
+ mask = torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
63
+ with torch.no_grad():
64
+ out = layer(hidden, position_embeddings=(cos, sin), attention_mask=mask)
65
+ return out[0] if isinstance(out, tuple) else out
66
+
67
+
68
+ def _to_device(t, mesh_device):
69
+ return ttnn.from_torch(
70
+ t, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG
71
+ )
72
+
73
+
74
+ def _setup(mesh_device, hf_config, torch_weights, block_size=None):
75
+ config = DecoderLayerConfig.from_hf(hf_config)
76
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
77
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
78
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
79
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=block_size)
80
+ return config, weights, cos_cache, sin_cache, sparsity, kv_cache
81
+
82
+
83
+ def _decode_step(mesh_device, hf_config, ctx, token_hidden, position):
84
+ config, weights, cos_cache, sin_cache, _, kv_cache = ctx
85
+ current_pos = ttnn.from_torch(torch.tensor([position], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
86
+ tt_in = _to_device(token_hidden.reshape(1, 1, 1, hf_config.hidden_size), mesh_device)
87
+ out = decoder_layer_decode(
88
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, token_index=position
89
+ )
90
+ return ttnn.to_torch(out).reshape(1, hf_config.hidden_size)
91
+
92
+
93
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
94
+ @pytest.mark.parametrize("prompt_len", [32, 128], ids=["p32", "p128"])
95
+ @pytest.mark.parametrize("block_size", [None, 32], ids=["contiguous", "paged32"])
96
+ def test_decode_layer_matches_prefill(mesh_device, reference, torch_weights, prompt_len, block_size):
97
+ layer, hf_config = reference
98
+ hidden_full = _hidden(hf_config, prompt_len + 1)
99
+ ref_out = _reference_layer(layer, hf_config, hidden_full)[:, prompt_len, :]
100
+
101
+ ctx = _setup(mesh_device, hf_config, torch_weights, block_size)
102
+ config, weights, cos_cache, sin_cache, sparsity, kv_cache = ctx
103
+
104
+ # Prefill the prompt through the full layer so the cache is populated.
105
+ decoder_layer_prefill(
106
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
107
+ weights,
108
+ config,
109
+ cos_cache,
110
+ sin_cache,
111
+ sparsity,
112
+ kv_cache=kv_cache,
113
+ )
114
+
115
+ tt_out = _decode_step(mesh_device, hf_config, ctx, hidden_full[:, prompt_len, :], prompt_len)
116
+
117
+ passing, pcc_message = comp_pcc(ref_out, tt_out, PCC_REQUIRED)
118
+ logger.info(comp_allclose(ref_out, tt_out))
119
+ kind = "contiguous" if block_size is None else f"paged(block={block_size})"
120
+ logger.info(f"decode layer at pos {prompt_len} [{kind}]: {pcc_message}")
121
+ assert passing, f"decode layer at pos {prompt_len} [{kind}] below {PCC_REQUIRED}: {pcc_message}"
122
+
123
+
124
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
125
+ def test_multi_step_decode(mesh_device, reference, torch_weights):
126
+ """Three sequential decode steps, each checked against the prefill reference.
127
+
128
+ A cache written one position off still passes a single step; it only shows
129
+ up when a later token reads what an earlier step stored.
130
+ """
131
+ prompt_len, steps = 32, 3
132
+ layer, hf_config = reference
133
+ hidden_full = _hidden(hf_config, prompt_len + steps)
134
+ ref_out = _reference_layer(layer, hf_config, hidden_full)
135
+
136
+ # Paged: multi-step is where a block-table mapping error would surface.
137
+ ctx = _setup(mesh_device, hf_config, torch_weights, block_size=32)
138
+ config, weights, cos_cache, sin_cache, sparsity, kv_cache = ctx
139
+
140
+ decoder_layer_prefill(
141
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
142
+ weights,
143
+ config,
144
+ cos_cache,
145
+ sin_cache,
146
+ sparsity,
147
+ kv_cache=kv_cache,
148
+ )
149
+
150
+ for step in range(steps):
151
+ pos = prompt_len + step
152
+ tt_out = _decode_step(mesh_device, hf_config, ctx, hidden_full[:, pos, :], pos)
153
+ passing, pcc_message = comp_pcc(ref_out[:, pos, :], tt_out, PCC_REQUIRED)
154
+ logger.info(f"decode step {step} (pos {pos}): {pcc_message}")
155
+ assert passing, f"decode step {step} at position {pos} below {PCC_REQUIRED}: {pcc_message}"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_determinism.py ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Determinism and repeated-run stability of the TTNN decoder layer.
5
+
6
+ Bit-exactness is asserted, not PCC. Identical inputs through identical kernels
7
+ must give identical bits; anything less means the result depends on something
8
+ not in the inputs -- uninitialised memory, a race between cores, or a reduction
9
+ whose order varies run to run. Those defects are intermittent by nature, so a
10
+ tolerance-based check would hide exactly the cases worth catching.
11
+
12
+ ``test_repeated_decode_steps_are_stable`` runs a long decode rollout instead.
13
+ It is the counterpart test: nothing there is compared against a reference, it
14
+ just has to keep producing finite, non-degenerate activations for 64 steps.
15
+ Cache-indexing and accumulation faults tend to show up as slow drift rather
16
+ than a hard failure, and a handful of steps will not surface them.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import pytest
22
+ import torch
23
+ from loguru import logger
24
+
25
+ import ttnn
26
+
27
+ from ..tt.functional_decoder import (
28
+ DecoderLayerConfig,
29
+ build_expert_sparsity,
30
+ build_rope_cache,
31
+ create_kv_cache,
32
+ decoder_layer_decode,
33
+ decoder_layer_prefill,
34
+ upload_layer_weights,
35
+ )
36
+ from ..tt.weight_mapping import convert_layer_weights
37
+ from .reference import build_reference_layer, layer_state_dict
38
+
39
+ LAYER_IDX = 0
40
+ MAX_SEQ = 256
41
+ BLOCK_SIZE = 32
42
+
43
+
44
+ @pytest.fixture(scope="module")
45
+ def reference():
46
+ return build_reference_layer(LAYER_IDX)
47
+
48
+
49
+ @pytest.fixture(scope="module")
50
+ def torch_weights(reference):
51
+ _, hf_config = reference
52
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
53
+
54
+
55
+ def _hidden(hf_config, seq_len, seed=0):
56
+ torch.manual_seed(seed)
57
+ return torch.randn(1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
58
+
59
+
60
+ def _to_device(t, mesh_device):
61
+ return ttnn.from_torch(
62
+ t, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG
63
+ )
64
+
65
+
66
+ def _build(mesh_device, hf_config, torch_weights):
67
+ config = DecoderLayerConfig.from_hf(hf_config)
68
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
69
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
70
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
71
+ return config, weights, cos_cache, sin_cache, sparsity
72
+
73
+
74
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
75
+ def test_prefill_is_deterministic(mesh_device, reference, torch_weights):
76
+ _, hf_config = reference
77
+ config, weights, cos_cache, sin_cache, sparsity = _build(mesh_device, hf_config, torch_weights)
78
+ hidden = _hidden(hf_config, 128)
79
+
80
+ outs = []
81
+ for _ in range(3):
82
+ tt_in = _to_device(hidden.unsqueeze(0), mesh_device)
83
+ out = decoder_layer_prefill(tt_in, weights, config, cos_cache, sin_cache, sparsity)
84
+ outs.append(ttnn.to_torch(out).clone())
85
+
86
+ assert torch.equal(outs[0], outs[1]), "prefill run 1 != run 2 (bitwise)"
87
+ assert torch.equal(outs[0], outs[2]), "prefill run 1 != run 3 (bitwise)"
88
+ logger.info("prefill: 3 runs bit-identical")
89
+
90
+
91
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
92
+ def test_decode_is_deterministic(mesh_device, reference, torch_weights):
93
+ """Two independent prefill+decode sequences must agree bitwise.
94
+
95
+ Each repetition allocates a fresh paged cache, so this also checks that no
96
+ state leaks between runs through the cache or the page table.
97
+ """
98
+ _, hf_config = reference
99
+ config, weights, cos_cache, sin_cache, sparsity = _build(mesh_device, hf_config, torch_weights)
100
+ prompt_len = 32
101
+ hidden_full = _hidden(hf_config, prompt_len + 1)
102
+
103
+ outs = []
104
+ for _ in range(2):
105
+ kv_cache = create_kv_cache(
106
+ mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE
107
+ )
108
+ decoder_layer_prefill(
109
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
110
+ weights,
111
+ config,
112
+ cos_cache,
113
+ sin_cache,
114
+ sparsity,
115
+ kv_cache=kv_cache,
116
+ )
117
+ current_pos = ttnn.from_torch(
118
+ torch.tensor([prompt_len], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device
119
+ )
120
+ tt_in = _to_device(hidden_full[:, prompt_len, :].reshape(1, 1, 1, hf_config.hidden_size), mesh_device)
121
+ out = decoder_layer_decode(
122
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, token_index=prompt_len
123
+ )
124
+ outs.append(ttnn.to_torch(out).clone())
125
+
126
+ assert torch.equal(outs[0], outs[1]), "decode from two fresh caches differs bitwise"
127
+ logger.info("decode: 2 independent prefill+decode sequences bit-identical")
128
+
129
+
130
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
131
+ def test_repeated_decode_steps_are_stable(mesh_device, reference, torch_weights):
132
+ """A 64-step rollout must stay finite and non-degenerate.
133
+
134
+ Long rollouts are where cache-indexing and accumulation faults show up as
135
+ drift rather than an exception, so this watches the activation scale across
136
+ every step instead of only checking the last one.
137
+ """
138
+ _, hf_config = reference
139
+ config, weights, cos_cache, sin_cache, sparsity = _build(mesh_device, hf_config, torch_weights)
140
+ prompt_len, steps = 32, 64
141
+ hidden_full = _hidden(hf_config, prompt_len + steps)
142
+
143
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE)
144
+ decoder_layer_prefill(
145
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
146
+ weights,
147
+ config,
148
+ cos_cache,
149
+ sin_cache,
150
+ sparsity,
151
+ kv_cache=kv_cache,
152
+ )
153
+
154
+ stds = []
155
+ for step in range(steps):
156
+ pos = prompt_len + step
157
+ current_pos = ttnn.from_torch(torch.tensor([pos], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
158
+ tt_in = _to_device(hidden_full[:, pos, :].reshape(1, 1, 1, hf_config.hidden_size), mesh_device)
159
+ out = decoder_layer_decode(tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, token_index=pos)
160
+ t = ttnn.to_torch(out).float()
161
+ assert torch.isfinite(t).all(), f"decode step {step} (pos {pos}) produced non-finite values"
162
+ stds.append(float(t.std()))
163
+
164
+ lo, hi = min(stds), max(stds)
165
+ logger.info(f"64-step rollout: activation std min={lo:.5f} max={hi:.5f} ratio={hi / lo:.3f}")
166
+ assert lo > 1e-6, f"activations collapsed to zero during the rollout (min std {lo:.3e})"
167
+ assert hi / lo < 5.0, f"activation scale drifted {hi / lo:.1f}x across 64 steps -- suspect cache indexing"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_full_model.py ADDED
@@ -0,0 +1,976 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Full-model and generator gates for Qwen3-Coder-30B-A3B on the 4-die mesh.
5
+
6
+ Two tiers, selected by ``QWEN3_FULL_MODEL_LAYERS`` (default 2):
7
+
8
+ * the **reduced** tier, one real layer of each kind (there is only one kind
9
+ here) with every other shape, memory config, cache/page-table layout, terminal
10
+ norm/LM head and sampler call identical to the shipped path. Two layers load
11
+ in ~10 s, which is what makes these runnable as a normal test suite;
12
+ * the **all-layer** tier, ``QWEN3_FULL_MODEL_LAYERS=48``, which is the final
13
+ evidence and takes several minutes to load.
14
+
15
+ Accuracy against HuggingFace is *not* asserted here. A 48-layer torch reference
16
+ is a 61 GB CPU forward, so the accuracy gate is
17
+ ``models.common.readiness_check.run_prefill_check`` /
18
+ ``run_teacher_forcing`` against the AIME24 chat reference, reported in
19
+ ``doc/full_model/README.md``. What these tests own is everything that can be
20
+ wrong *without* moving PCC: the trace/feedback contract, position coherence,
21
+ page-table refresh policy, non-aligned prompt lengths, batch handling, cache
22
+ ownership, reset semantics and the runtime fallback audit.
23
+ """
24
+
25
+ from __future__ import annotations
26
+
27
+ import math
28
+ import os
29
+
30
+ import pytest
31
+ import torch
32
+
33
+ import ttnn
34
+ from models.common.utility_functions import comp_pcc
35
+
36
+ from ..tt import functional_decoder as FD
37
+ from ..tt import multichip_decoder as MC
38
+ from ..tt.generator import Qwen3CoderGenerator, _first_device_to_torch, build_generator, prefill_bucket_ladder
39
+ from ..tt.model import DEFAULT_TRACE_REGION_SIZE
40
+
41
+ MODEL_DIR = "models/demos/blackhole/qwen3_coder_30b_a3b"
42
+ CONTEXT = 8192
43
+
44
+
45
+ @pytest.fixture(scope="module")
46
+ def mesh_device():
47
+ """A **module-scoped** 4-die ring mesh, opened here rather than by conftest.
48
+
49
+ The repository's `mesh_device` fixture is function-scoped, so a module-scoped
50
+ generator cannot depend on it (`ScopeMismatch`). Reopening the mesh per test
51
+ would also mean reloading the model per test -- ten seconds at two layers and
52
+ over three minutes at forty-eight -- which would make the all-layer tier
53
+ unrunnable. `FABRIC_1D_RING` must be set before the open, exactly as the
54
+ stage-03/04 tests and every probe in `doc/full_model/probes/` do it.
55
+ """
56
+ ttnn.set_fabric_config(ttnn.FabricConfig.FABRIC_1D_RING)
57
+ mesh = ttnn.open_mesh_device(mesh_shape=ttnn.MeshShape(*MC.MESH_SHAPE), trace_region_size=DEFAULT_TRACE_REGION_SIZE)
58
+ yield mesh
59
+ ttnn.close_mesh_device(mesh)
60
+ ttnn.set_fabric_config(ttnn.FabricConfig.DISABLED)
61
+
62
+
63
+ #: Layer count for this run. An environment variable rather than a pytest
64
+ #: option so that the choice survives being run as part of the whole model
65
+ #: suite from the repository root, where a subdirectory ``conftest.py`` would
66
+ #: not have been loaded in time to register an option.
67
+ NUM_LAYERS = int(os.environ.get("QWEN3_FULL_MODEL_LAYERS", "2"))
68
+
69
+
70
+ @pytest.fixture(scope="module")
71
+ def num_layers():
72
+ return NUM_LAYERS
73
+
74
+
75
+ @pytest.fixture(scope="module")
76
+ def generator(mesh_device, num_layers):
77
+ gen = build_generator(
78
+ MODEL_DIR,
79
+ mesh_device,
80
+ override_num_layers=num_layers,
81
+ max_context_len=CONTEXT,
82
+ max_batch_size=1,
83
+ )
84
+ yield gen
85
+ gen.teardown()
86
+
87
+
88
+ @pytest.fixture(scope="module")
89
+ def batch_generator(mesh_device, num_layers):
90
+ gen = build_generator(
91
+ MODEL_DIR,
92
+ mesh_device,
93
+ override_num_layers=num_layers,
94
+ max_context_len=1024,
95
+ max_batch_size=4,
96
+ )
97
+ yield gen
98
+ gen.teardown()
99
+
100
+
101
+ @pytest.fixture(scope="module")
102
+ def small_rope_generator(mesh_device, num_layers):
103
+ """A generator whose cos/sin tables are far shorter than its context.
104
+
105
+ ``rope_cache_len`` defaults to 8192 against a 262144-token contract, so the
106
+ gap between "the table is sized" and "the context is advertised" is real on
107
+ the shipped configuration; 64 makes it reachable in a handful of tokens.
108
+ """
109
+ gen = build_generator(
110
+ MODEL_DIR,
111
+ mesh_device,
112
+ override_num_layers=num_layers,
113
+ max_context_len=CONTEXT,
114
+ max_batch_size=1,
115
+ rope_cache_len=64,
116
+ )
117
+ yield gen
118
+ gen.teardown()
119
+
120
+
121
+ def _prompt_ids(gen, text: str) -> list[int]:
122
+ rendered = gen.tokenizer.apply_chat_template(
123
+ [{"role": "user", "content": text}], add_generation_prompt=True, tokenize=False
124
+ )
125
+ return gen.tokenizer(rendered, add_special_tokens=False)["input_ids"]
126
+
127
+
128
+ # --- the generator contract ---------------------------------------------------
129
+
130
+
131
+ def test_generator_implements_the_contract(generator):
132
+ from models.common.readiness_check.contract import Generator
133
+
134
+ assert isinstance(generator, Generator)
135
+ assert isinstance(generator, Qwen3CoderGenerator)
136
+ assert generator.tokenizer is not None
137
+ import inspect
138
+
139
+ # The teacher-forcing runner requires an explicit keyword, not **kwargs.
140
+ assert "enable_trace" in inspect.signature(generator.generate).parameters
141
+
142
+
143
+ # --- split sampling, token feedback, position coherence -----------------------
144
+
145
+
146
+ def test_split_sampling_feeds_its_own_token_back_on_device(generator):
147
+ """Step N's sampled token *is* step N+1's token input, with no host copy."""
148
+ generator.reset()
149
+ prompt = _prompt_ids(generator, "List three prime numbers.")
150
+ kv_cache = generator._ensure_kv_cache()
151
+ page_table = generator.make_page_table([len(prompt) + 8])
152
+ sampled = generator.prefill_forward(
153
+ torch.tensor([prompt]),
154
+ page_table=page_table,
155
+ kv_cache=kv_cache,
156
+ prompt_lens=[len(prompt)],
157
+ sampling_mode="device",
158
+ )
159
+ first = int(generator._sampled_to_torch(sampled)[0].item())
160
+
161
+ def read(tensor):
162
+ return int(ttnn.to_torch(ttnn.get_device_tensors(tensor)[0]).reshape(-1)[0].item())
163
+
164
+ # Installing the trace also performs the first replay, so the prefill token
165
+ # has already been consumed by the time anything can be read back.
166
+ host_copies_before = generator.trace_stats["token_host_copies"]
167
+ generator.decode_forward(
168
+ None,
169
+ torch.tensor([len(prompt)]),
170
+ page_table=page_table,
171
+ kv_cache=kv_cache,
172
+ sampling_mode="device",
173
+ enable_trace=True,
174
+ active_batch=1,
175
+ )
176
+ token_in, current_pos, rotary_pos, _ = generator._trace_inputs
177
+
178
+ observed = [read(token_in)]
179
+ positions = [(read(current_pos), read(rotary_pos))]
180
+ # The sampler wrote through tt_out_tok, so the persistent decode token input
181
+ # already holds what the sampling trace produced -- same value, same tensor.
182
+ assert read(generator._trace_sampled) == observed[0]
183
+
184
+ for _ in range(3):
185
+ generator.decode_forward(
186
+ None, None, page_table=None, kv_cache=kv_cache, sampling_mode="device", enable_trace=True
187
+ )
188
+ observed.append(read(token_in))
189
+ positions.append((read(current_pos), read(rotary_pos)))
190
+ assert read(generator._trace_sampled) == observed[-1]
191
+
192
+ # Positions were advanced on device by the trace itself, one per replay,
193
+ # starting from the prompt length, and cache position and rotary position
194
+ # stayed in lockstep.
195
+ assert positions[0] == (len(prompt) + 1, len(prompt) + 1), positions
196
+ for step in range(len(positions) - 1):
197
+ assert positions[step + 1][0] == positions[step][0] + 1, positions
198
+ assert positions[step + 1][1] == positions[step][1] + 1, positions
199
+
200
+ # Nothing was written to the token input from the host across any of it.
201
+ assert generator.trace_stats["token_host_copies"] == host_copies_before
202
+
203
+ # And the tokens observed on device are exactly what the public generator
204
+ # returns for the same prompt: [prefill sample, then one per replay].
205
+ generator.reset()
206
+ produced = generator.generate(prompt, 1 + len(observed), enable_trace=True, sampling_mode="device")
207
+ assert produced == [first] + observed, (produced, first, observed)
208
+
209
+
210
+ def test_steady_state_decode_does_no_host_work(generator):
211
+ """Only ``replays`` may move between two steady-state tokens."""
212
+ generator.reset()
213
+ prompt = _prompt_ids(generator, "Say hello.")
214
+ generator.generate(prompt, 4, enable_trace=True, sampling_mode="device")
215
+ before = dict(generator.trace_stats)
216
+ generator.decode_forward(
217
+ None, None, page_table=None, kv_cache=generator._kv_cache, sampling_mode="device", enable_trace=True
218
+ )
219
+ after = dict(generator.trace_stats)
220
+ moved = {k: (before[k], after[k]) for k in before if before[k] != after[k]}
221
+ assert moved == {"replays": (before["replays"], before["replays"] + 1)}, moved
222
+
223
+
224
+ def test_unchanged_page_table_costs_no_host_copy(generator):
225
+ """A page table that has not changed must not be re-uploaded."""
226
+ generator.reset()
227
+ prompt = _prompt_ids(generator, "Say hello.")
228
+ generator.generate(prompt, 3, enable_trace=True, sampling_mode="device")
229
+ page_table = generator._trace_page_table_snapshot.clone()
230
+ before = generator.trace_stats["page_table_host_copies"]
231
+ generator._refresh_persistent_page_table(page_table, generator._trace_kv_cache, active_batch=1)
232
+ assert generator.trace_stats["page_table_host_copies"] == before
233
+
234
+ changed = page_table.clone()
235
+ changed[0, -1] = 0 if changed[0, -1] != 0 else 1
236
+ generator._refresh_persistent_page_table(changed, generator._trace_kv_cache, active_batch=1)
237
+ assert generator.trace_stats["page_table_host_copies"] == before + 1
238
+
239
+
240
+ def test_device_split_sampling_matches_host_argmax(generator):
241
+ """Greedy split sampling is semantically greedy, not merely close."""
242
+ generator.reset()
243
+ prompt = _prompt_ids(generator, "The capital of France is")
244
+ device_tokens = generator.generate(prompt, 6, enable_trace=True, sampling_mode="device")
245
+ generator.reset()
246
+ host_tokens = generator.generate(prompt, 6, sampling_mode="host")
247
+ assert device_tokens == host_tokens, (device_tokens, host_tokens)
248
+
249
+
250
+ def test_force_argmax_matches_split_sampling(generator):
251
+ """The rejected alternative gives the same token as the shipped one."""
252
+ generator.reset()
253
+ prompt = _prompt_ids(generator, "2 + 2 =")
254
+ kv_cache = generator._ensure_kv_cache()
255
+ page_table = generator.make_page_table([len(prompt) + 1])
256
+ split = generator.prefill_forward(
257
+ torch.tensor([prompt]),
258
+ page_table=page_table,
259
+ kv_cache=kv_cache,
260
+ prompt_lens=[len(prompt)],
261
+ sampling_mode="device",
262
+ )
263
+ split_token = int(generator._sampled_to_torch(split)[0].item())
264
+
265
+ generator.reset()
266
+ logits = generator.prefill_forward(
267
+ torch.tensor([prompt]),
268
+ page_table=page_table,
269
+ kv_cache=kv_cache,
270
+ prompt_lens=[len(prompt)],
271
+ sampling_mode="host",
272
+ )
273
+ assert split_token == int(logits[0, 0].argmax().item())
274
+
275
+
276
+ # --- rotary capacity on the low-level API -------------------------------------
277
+
278
+
279
+ def _low_level_decode(gen, prompt, steps, *, decode_horizon=None):
280
+ """Drive ``prefill_forward``/``decode_forward`` exactly as the docstring says.
281
+
282
+ This is the surface a serving adapter drives, and the only one that can
283
+ decode past the rotary table: ``generate`` sizes for its own horizon.
284
+ """
285
+ gen.reset()
286
+ kv_cache = gen._ensure_kv_cache()
287
+ page_table = gen.make_page_table([len(prompt) + steps + 1])
288
+ sampled = gen.prefill_forward(
289
+ torch.tensor([prompt]),
290
+ page_table=page_table,
291
+ kv_cache=kv_cache,
292
+ prompt_lens=[len(prompt)],
293
+ sampling_mode="device",
294
+ )
295
+ tokens = [int(gen._sampled_to_torch(sampled)[0].item())]
296
+ for step in range(steps):
297
+ initial = step == 0
298
+ sampled = gen.decode_forward(
299
+ None,
300
+ torch.tensor([len(prompt)]) if initial else None,
301
+ page_table=page_table if initial else None,
302
+ kv_cache=kv_cache,
303
+ sampling_mode="device",
304
+ enable_trace=True,
305
+ active_batch=1,
306
+ **({"decode_horizon": decode_horizon} if initial and decode_horizon is not None else {}),
307
+ )
308
+ tokens.append(int(gen._sampled_to_torch(sampled)[0].item()))
309
+ return tokens
310
+
311
+
312
+ def test_decode_past_the_rope_cache_length_through_the_low_level_api(small_rope_generator, expect_error):
313
+ """Walking off the cos/sin table must raise, and must be preventable.
314
+
315
+ The traced loop advances ``rotary_position`` with ``ttnn.plus_one`` and
316
+ nothing on device clamps it, so an out-of-range ``ttnn.embedding`` gather
317
+ would rotate at a wrong position and return a plausible-looking token. The
318
+ contract advertises 262144; the tables default to 8192.
319
+ """
320
+ gen = small_rope_generator
321
+ rope_len = gen.model.rope_cache_len
322
+ prompt = _prompt_ids(gen, "Count upwards.")
323
+ assert len(prompt) < rope_len, (len(prompt), rope_len)
324
+ steps = rope_len - len(prompt) + 6 # comfortably past the table
325
+
326
+ # 1. Undeclared horizon: the run must stop rather than silently gather out
327
+ # of range. The message has to name the fix.
328
+ with expect_error(RuntimeError, "rotary table"):
329
+ _low_level_decode(gen, prompt, steps)
330
+ assert gen.model.rope_cache_len == rope_len, "the failing path must not have grown the tables"
331
+
332
+ # 2. Declared horizon: the same run completes, and the tables grew.
333
+ horizon = len(prompt) + steps
334
+ declared = _low_level_decode(gen, prompt, steps, decode_horizon=horizon)
335
+ assert len(declared) == steps + 1
336
+ assert gen.model.rope_cache_len > rope_len
337
+ assert gen.model.rope_cache_len >= horizon
338
+
339
+ # 3. And it is the *right* answer: identical to what the high-level
340
+ # ``generate``, which sizes its own rotary horizon, produces.
341
+ gen.reset()
342
+ reference = gen.generate(prompt, steps + 1, enable_trace=True, sampling_mode="device")
343
+ assert declared == reference, (declared, reference)
344
+
345
+
346
+ def test_eager_decode_grows_the_rope_tables_for_its_position(small_rope_generator):
347
+ """The eager/host branch gathers cos/sin too, and holds no trace to protect."""
348
+ gen = small_rope_generator
349
+ gen.reset()
350
+ prompt_len = 8
351
+ position = gen.model.rope_cache_len + 5
352
+ kv_cache = gen._ensure_kv_cache()
353
+ page_table = gen.make_page_table([position + 1])
354
+ gen.prefill_forward(
355
+ torch.arange(1000, 1000 + prompt_len, dtype=torch.long).unsqueeze(0),
356
+ page_table=page_table,
357
+ kv_cache=kv_cache,
358
+ prompt_lens=[prompt_len],
359
+ sampling_mode="host",
360
+ )
361
+ logits = gen.decode_forward(
362
+ torch.tensor([[42]]),
363
+ torch.tensor([position]),
364
+ page_table=page_table,
365
+ kv_cache=kv_cache,
366
+ sampling_mode="host",
367
+ enable_trace=False,
368
+ )
369
+ assert torch.isfinite(logits).all()
370
+ assert gen.model.rope_cache_len > position
371
+
372
+
373
+ def test_decode_beyond_the_advertised_context_is_refused(generator, expect_error):
374
+ with expect_error(ValueError, "exceeds the supported context"):
375
+ generator.decode_forward(
376
+ None,
377
+ torch.tensor([generator.model.max_cache_len + 1]),
378
+ page_table=generator.make_page_table([8]),
379
+ kv_cache=generator._ensure_kv_cache(),
380
+ sampling_mode="device",
381
+ enable_trace=True,
382
+ active_batch=1,
383
+ )
384
+
385
+
386
+ # --- top-k / top-p sampling ---------------------------------------------------
387
+
388
+
389
+ def test_top_k_top_p_sampling_runs_through_a_traced_generate(generator):
390
+ """The stochastic route is exercised, not merely reachable.
391
+
392
+ ``README.md`` calls the top-k/top-p route "a live code path, not a promise";
393
+ this is the assertion behind that sentence. It drives ``sample_split``
394
+ through a captured trace and checks the generator really switched strategy.
395
+ """
396
+ gen = generator
397
+ gen.reset()
398
+ prompt = _prompt_ids(gen, "Name a colour.")
399
+ calls = {"split": 0, "argmax": 0}
400
+ real_split = gen.model.sample_split
401
+ real_argmax = gen.model.sample_greedy_argmax
402
+
403
+ def counting_split(*a, **k):
404
+ calls["split"] += 1
405
+ return real_split(*a, **k)
406
+
407
+ def counting_argmax(*a, **k):
408
+ calls["argmax"] += 1
409
+ return real_argmax(*a, **k)
410
+
411
+ gen.model.sample_split = counting_split
412
+ gen.model.sample_greedy_argmax = counting_argmax
413
+ try:
414
+ tokens = gen.generate(prompt, 5, enable_trace=True, sampling_mode="device", top_k=8, top_p=0.9, temperature=0.8)
415
+ finally:
416
+ gen.model.sample_split = real_split
417
+ gen.model.sample_greedy_argmax = real_argmax
418
+
419
+ assert len(tokens) == 5
420
+ assert all(0 <= t < gen.model.vocab_size for t in tokens), tokens
421
+ assert gen._sampling_stochastic is True
422
+ # Every sampler dispatch on this run went to the split path, and none to
423
+ # force-argmax -- warm-up, capture and prefill included.
424
+ assert calls["split"] > 0 and calls["argmax"] == 0, calls
425
+ assert gen._trace_model_id is not None and gen._trace_sampling_id is not None
426
+
427
+
428
+ def test_alternating_sampling_modes_recapture_the_traces(generator):
429
+ """The trace-id cache is keyed by sampling mode; prove the key is honoured.
430
+
431
+ A stale trace served across a greedy/stochastic flip would silently sample
432
+ with the wrong strategy, which no accuracy gate on this stage would catch.
433
+ """
434
+ gen = generator
435
+ gen.reset()
436
+ prompt = _prompt_ids(gen, "The capital of France is")
437
+
438
+ greedy_first = gen.generate(prompt, 4, enable_trace=True, sampling_mode="device", top_k=1)
439
+ assert gen._sampling_stochastic is False
440
+ releases_before = gen.trace_stats["releases"]
441
+ captures_before = gen.trace_stats["captures"]
442
+
443
+ gen.reset()
444
+ stochastic = gen.generate(
445
+ prompt, 4, enable_trace=True, sampling_mode="device", top_k=16, top_p=0.95, temperature=1.0
446
+ )
447
+ assert gen._sampling_stochastic is True
448
+ assert len(stochastic) == 4
449
+ assert gen.trace_stats["releases"] > releases_before, gen.trace_stats
450
+ assert gen.trace_stats["captures"] > captures_before, gen.trace_stats
451
+
452
+ gen.reset()
453
+ greedy_again = gen.generate(prompt, 4, enable_trace=True, sampling_mode="device", top_k=1)
454
+ assert gen._sampling_stochastic is False
455
+ # Flipping back must restore the greedy strategy exactly, not leave the
456
+ # stochastic trace installed.
457
+ assert greedy_again == greedy_first, (greedy_first, greedy_again)
458
+
459
+
460
+ def test_temperature_zero_is_spelled_as_greedy(generator):
461
+ """A serving stack spells greedy ``temperature=0``; it must not go stochastic."""
462
+ gen = generator
463
+ gen.reset()
464
+ prompt = _prompt_ids(gen, "The capital of France is")
465
+ greedy = gen.generate(prompt, 4, enable_trace=True, sampling_mode="device", top_k=1)
466
+ gen.reset()
467
+ as_temp_zero = gen.generate(prompt, 4, enable_trace=True, sampling_mode="device", top_k=0, temperature=0.0)
468
+ assert gen._sampling_stochastic is False
469
+ assert as_temp_zero == greedy, (greedy, as_temp_zero)
470
+
471
+
472
+ def test_set_sampling_params_releases_traces_only_on_a_mode_flip(generator):
473
+ """Changing k/p within one mode must not cost a recapture."""
474
+ gen = generator
475
+ gen.reset()
476
+ prompt = _prompt_ids(gen, "Say hello.")
477
+ gen.generate(prompt, 3, enable_trace=True, sampling_mode="device", top_k=8, top_p=0.9)
478
+ assert gen._sampling_stochastic is True
479
+ releases = gen.trace_stats["releases"]
480
+
481
+ gen.set_sampling_params(top_k=16, top_p=0.5, temperature=0.7, active_batch=1)
482
+ assert gen._sampling_stochastic is True
483
+ assert gen.trace_stats["releases"] == releases, "a k/p change inside one mode recaptured"
484
+ assert gen._trace_model_id is not None
485
+
486
+ gen.set_sampling_params(top_k=1, top_p=0.0, temperature=1.0, active_batch=1)
487
+ assert gen._sampling_stochastic is False
488
+ assert gen.trace_stats["releases"] == releases + 1, "the greedy flip did not recapture"
489
+
490
+
491
+ # --- prompt lengths -----------------------------------------------------------
492
+
493
+
494
+ @pytest.mark.parametrize("prompt_len", [1, 31, 33, 100, 127, 128, 129, 257, 1000])
495
+ def test_non_aligned_prompt_lengths(generator, prompt_len):
496
+ """Every length up to the context, aligned or not, through the public API."""
497
+ generator.reset()
498
+ tokens = torch.arange(1000, 1000 + prompt_len, dtype=torch.long).unsqueeze(0)
499
+ kv_cache = generator._ensure_kv_cache()
500
+ page_table = generator.make_page_table([prompt_len + 1])
501
+ logits = generator.prefill_forward(
502
+ tokens,
503
+ page_table=page_table,
504
+ kv_cache=kv_cache,
505
+ prompt_lens=[prompt_len],
506
+ sampling_mode="host",
507
+ )
508
+ assert tuple(logits.shape) == (1, 1, generator.model.vocab_size)
509
+ assert torch.isfinite(logits).all()
510
+
511
+
512
+ def test_return_all_logits_is_sliced_to_the_logical_length(generator):
513
+ generator.reset()
514
+ prompt_len = 37
515
+ tokens = torch.arange(500, 500 + prompt_len, dtype=torch.long).unsqueeze(0)
516
+ kv_cache = generator._ensure_kv_cache()
517
+ logits = generator.prefill_forward(
518
+ tokens,
519
+ page_table=generator.make_page_table([prompt_len]),
520
+ kv_cache=kv_cache,
521
+ prompt_lens=[prompt_len],
522
+ return_all_logits=True,
523
+ sampling_mode="host",
524
+ )
525
+ assert tuple(logits.shape) == (1, prompt_len, generator.model.vocab_size)
526
+ assert torch.isfinite(logits).all()
527
+
528
+
529
+ # --- batch, fixed slots, inactive rows ---------------------------------------
530
+
531
+
532
+ def test_mixed_length_batch_prefill_and_decode(batch_generator):
533
+ """Four users, four different prompt lengths, disjoint physical pages."""
534
+ gen = batch_generator
535
+ gen.reset()
536
+ lengths = [7, 33, 64, 129]
537
+ width = max(lengths)
538
+ tokens = torch.zeros(len(lengths), width, dtype=torch.long)
539
+ for user, length in enumerate(lengths):
540
+ tokens[user, :length] = torch.arange(100 + user * 50, 100 + user * 50 + length)
541
+ kv_cache = gen._ensure_kv_cache()
542
+ page_table = gen.make_page_table([length + 4 for length in lengths])
543
+ logits = gen.prefill_forward(
544
+ tokens, page_table=page_table, kv_cache=kv_cache, prompt_lens=lengths, sampling_mode="host"
545
+ )
546
+ assert tuple(logits.shape) == (len(lengths), 1, gen.model.vocab_size)
547
+ assert torch.isfinite(logits).all()
548
+
549
+ predicted = logits[:, 0].argmax(dim=-1)
550
+ decoded = gen.decode_forward(
551
+ predicted.reshape(-1, 1),
552
+ torch.tensor(lengths),
553
+ page_table=page_table,
554
+ kv_cache=kv_cache,
555
+ sampling_mode="host",
556
+ enable_trace=False,
557
+ )
558
+ assert tuple(decoded.shape) == (len(lengths), gen.model.vocab_size)
559
+ assert torch.isfinite(decoded).all()
560
+
561
+
562
+ def test_inactive_rows_are_expressible(batch_generator):
563
+ """A negative position marks an inactive slot and must not be validated."""
564
+ gen = batch_generator
565
+ gen.reset()
566
+ lengths = [16, 16]
567
+ tokens = torch.arange(2 * 16, dtype=torch.long).reshape(2, 16)
568
+ kv_cache = gen._ensure_kv_cache()
569
+ page_table = gen.make_page_table([20, 20])
570
+ gen.prefill_forward(tokens, page_table=page_table, kv_cache=kv_cache, prompt_lens=lengths, sampling_mode="host")
571
+ positions = torch.tensor([16, -1])
572
+ gen._validate_page_coverage(gen._normalise_page_table(page_table, 2), positions, 2)
573
+
574
+
575
+ def test_page_table_must_map_disjoint_pages(batch_generator, expect_error):
576
+ gen = batch_generator
577
+ table = gen.make_page_table([64, 64])
578
+ table[1, :] = table[0, :]
579
+ with expect_error(ValueError, "must map disjoint physical cache pages"):
580
+ gen._validate_page_coverage(gen._normalise_page_table(table, 2), torch.tensor([63, 63]), 2)
581
+
582
+
583
+ def test_sdpa_rounded_page_count_covers_the_read_window(generator):
584
+ """The allocation must match the kernel's rounded read, not the token count."""
585
+ for tokens, expected in ((1, 1), (32, 1), (33, 2), (96, 4), (256, 8), (257, 16), (320, 16), (513, 24)):
586
+ assert generator._sdpa_rounded_page_count(tokens) == expected, tokens
587
+ for tokens in (1, 33, 100, 257, 1000, 4095):
588
+ assert generator._sdpa_rounded_page_count(tokens) >= math.ceil(tokens / generator.page_block_size)
589
+
590
+
591
+ def test_distributed_argmax_is_exact_at_batch_above_one(batch_generator):
592
+ """The live-row slice must be right at every batch, not only at batch 1.
593
+
594
+ Stage 06 made ``_WatcherCleanSampling1D._sample_argmax`` **batch-dependent**:
595
+ it slices the 32-slot logit tile down to ``_dist_active_rows`` before the
596
+ per-die ``ttnn.argmax`` and pads the result back. Every other device-sampling
597
+ test in this file uses the ``max_batch_size=1`` fixture and the
598
+ ``max_batch_size=4`` fixture only ever samples on the host, so the branch
599
+ that the slice introduced was uncovered at batch > 1 -- which is exactly
600
+ where an off-by-one in the slice or the pad would live.
601
+
602
+ This drives the sampler directly with crafted logits, because the property
603
+ is about the reduction and not about what the model predicts:
604
+
605
+ * every **live** row returns the host argmax of the same bf16 logits;
606
+ * every **padding** row returns token 0, which is the value the shipped
607
+ 32-row reduction produces for a zero-logit row and the value the pad
608
+ writes back, so the 32-slot buffer is unchanged slot for slot;
609
+ * the caller's ``tt_out_tok`` object survives -- the traced decode loop
610
+ feeds that exact tensor back, so a new tensor would break feedback
611
+ silently.
612
+
613
+ The all-negative leg matters on its own: it is the case where "the padding
614
+ rows are zero" stops being harmless, because a zero padding row would beat
615
+ every live row if the slice were not there.
616
+ """
617
+ gen = batch_generator
618
+ model = gen.model
619
+ sampler = model.sampler
620
+ mesh = gen.mesh_device
621
+ dies = mesh.get_num_devices()
622
+ local_vocab = model.vocab_size // dies
623
+
624
+ assert sampler._dist_active_rows == model.max_batch_size > 1, (
625
+ f"this test is only meaningful when the sampler is batched: "
626
+ f"_dist_active_rows={sampler._dist_active_rows}, max_batch_size={model.max_batch_size}"
627
+ )
628
+ sampler.load_device_buffers()
629
+ assert getattr(sampler, "_dist_die_offset", None) is not None, "the distributed path is not active"
630
+ assert sampler._dist_local_vocab == local_vocab
631
+
632
+ slots = 32
633
+ active = sampler._dist_active_rows
634
+ torch.manual_seed(0)
635
+ legs = {
636
+ # random logits: the ordinary case, and the winner lands on a different
637
+ # die for different rows
638
+ "random": torch.randn(1, 1, slots, model.vocab_size),
639
+ # every live logit strictly negative: a padding row's exact 0.0 would win
640
+ # every row if the live-row slice were not doing its job
641
+ "all_negative": -1.0 - torch.rand(1, 1, slots, model.vocab_size),
642
+ }
643
+ for name, logits in legs.items():
644
+ # bf16 on the way in, so the host reference sees the same values the
645
+ # device compares -- otherwise near-ties round differently.
646
+ logits = logits.to(torch.bfloat16).to(torch.float32)
647
+ expected = logits[0, 0, :active].argmax(dim=-1).tolist()
648
+
649
+ device_logits = ttnn.from_torch(
650
+ logits,
651
+ dtype=ttnn.bfloat16,
652
+ layout=ttnn.TILE_LAYOUT,
653
+ device=mesh,
654
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
655
+ mesh_mapper=ttnn.ShardTensorToMesh(mesh, dim=-1),
656
+ )
657
+ out_tok = ttnn.from_torch(
658
+ torch.full((1, 1, 1, slots), 12345, dtype=torch.int32),
659
+ dtype=ttnn.uint32,
660
+ layout=ttnn.ROW_MAJOR_LAYOUT,
661
+ device=mesh,
662
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
663
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh),
664
+ )
665
+ returned, logprobs = sampler._sample_argmax(device_logits, out_tok)
666
+ assert logprobs is None
667
+ assert returned is out_tok, f"{name}: the sampler returned a new tensor, breaking token feedback"
668
+
669
+ tokens = [int(v) for v in ttnn.to_torch(ttnn.get_device_tensors(returned)[0]).reshape(-1)[:slots].tolist()]
670
+ assert tokens[:active] == expected, f"{name}: live rows {tokens[:active]} != host argmax {expected}"
671
+ assert tokens[active:] == [0] * (slots - active), f"{name}: padding rows are {tokens[active:]}, not 0"
672
+ ttnn.deallocate(device_logits)
673
+ ttnn.deallocate(out_tok)
674
+
675
+
676
+ # --- cache ownership, reset, determinism -------------------------------------
677
+
678
+
679
+ def test_caller_owned_cache_is_used_verbatim(generator):
680
+ """A caller-allocated cache must be honoured, not silently replaced."""
681
+ generator.reset()
682
+ caller_cache = generator.model.allocate_kv_cache(max_cache_len=1024, num_blocks=64)
683
+ prompt = _prompt_ids(generator, "Hi.")
684
+ generator.prefill_forward(
685
+ torch.tensor([prompt]),
686
+ page_table=generator.make_page_table([len(prompt)]),
687
+ kv_cache=caller_cache,
688
+ prompt_lens=[len(prompt)],
689
+ sampling_mode="host",
690
+ )
691
+ written = ttnn.to_torch(ttnn.get_device_tensors(caller_cache[0].k)[0])
692
+ assert written.abs().sum() > 0, "prefill did not write into the caller's cache"
693
+ for cache in caller_cache:
694
+ ttnn.deallocate(cache.k, True)
695
+ ttnn.deallocate(cache.v, True)
696
+
697
+
698
+ def test_reset_makes_generation_reproducible(generator):
699
+ prompt = _prompt_ids(generator, "Count to five.")
700
+ generator.reset()
701
+ first = generator.generate(prompt, 6, enable_trace=True, sampling_mode="device")
702
+ generator.reset()
703
+ second = generator.generate(prompt, 6, enable_trace=True, sampling_mode="device")
704
+ assert first == second, (first, second)
705
+
706
+
707
+ def test_reset_zeroes_the_cache(generator):
708
+ generator.reset()
709
+ prompt = _prompt_ids(generator, "Hello there.")
710
+ generator.generate(prompt, 2, enable_trace=True, sampling_mode="device")
711
+ generator.reset()
712
+ for cache in generator._kv_cache:
713
+ assert ttnn.to_torch(ttnn.get_device_tensors(cache.k)[0]).abs().sum() == 0
714
+ assert ttnn.to_torch(ttnn.get_device_tensors(cache.v)[0]).abs().sum() == 0
715
+
716
+
717
+ def test_prefill_logits_are_deterministic_across_runs(generator):
718
+ generator.reset()
719
+ prompt = _prompt_ids(generator, "Deterministic?")
720
+ args = dict(
721
+ page_table=generator.make_page_table([len(prompt)]),
722
+ kv_cache=generator._ensure_kv_cache(),
723
+ prompt_lens=[len(prompt)],
724
+ sampling_mode="host",
725
+ )
726
+ first = generator.prefill_forward(torch.tensor([prompt]), **args)
727
+ generator.reset()
728
+ second = generator.prefill_forward(torch.tensor([prompt]), **args)
729
+ assert torch.equal(first, second)
730
+
731
+
732
+ # --- the carried-forward decoder contract ------------------------------------
733
+
734
+
735
+ def test_runtime_fallback_audit_is_clean(generator):
736
+ audit = generator.model.runtime_fallback_audit()
737
+ assert audit["dram_sharded_taken"] is True
738
+ # Stage 07 retuned both to their full-K ceilings (+2.83% decode, no
739
+ # accuracy change); see doc/datatype_sweep/README.md.
740
+ assert audit["gate_up_in0_block_w"] == 64
741
+ assert audit["down_in0_block_w"] == 24
742
+ assert audit["expert_intermediate_buffer"] == "L1"
743
+ assert audit["local_heads"] == (8, 1)
744
+ assert audit["local_experts"] == 32
745
+ assert audit["norm_shard_feeds_qkv_directly"] is True
746
+ assert audit["decode_ccl_buffers_persistent"] is True
747
+ assert audit["host_logit_readback_on_token_out_path"] is False
748
+ assert audit["host_argmax_on_token_out_path"] is False
749
+ assert audit["vocab_padding"] == 0
750
+ assert audit["kv_cache_dtype"] == "bfloat16"
751
+ assert audit["collective_topology"] == "Topology.Ring"
752
+ assert (audit["prefill_num_links"], audit["decode_num_links"]) == (2, 1)
753
+
754
+
755
+ def test_inter_layer_residual_contract_is_preserved(generator):
756
+ """A layer's output must be indistinguishable from its input as a tensor.
757
+
758
+ This is the stage-04 contract restated at the full-model boundary: if it
759
+ holds, 48 layers stack with no conversion, which is what
760
+ ``decode_hidden``'s bare ``for`` loop assumes.
761
+ """
762
+ model = generator.model
763
+ generator.reset()
764
+ tokens = ttnn.from_torch(
765
+ torch.zeros((1, 1, 1, 32), dtype=torch.int32),
766
+ device=model.mesh_device,
767
+ dtype=ttnn.uint32,
768
+ layout=ttnn.ROW_MAJOR_LAYOUT,
769
+ mesh_mapper=ttnn.ReplicateTensorToMesh(model.mesh_device),
770
+ )
771
+ hidden = model.embed_decode(tokens)
772
+ assert tuple(hidden.shape) == (1, 1, model.max_batch_size, model.hidden_size)
773
+ assert hidden.dtype == ttnn.bfloat16
774
+ assert hidden.layout == ttnn.TILE_LAYOUT
775
+ assert hidden.memory_config() == ttnn.DRAM_MEMORY_CONFIG
776
+
777
+ current_pos = ttnn.from_torch(
778
+ torch.zeros(model.max_batch_size, dtype=torch.int32),
779
+ device=model.mesh_device,
780
+ dtype=ttnn.int32,
781
+ layout=ttnn.ROW_MAJOR_LAYOUT,
782
+ mesh_mapper=ttnn.ReplicateTensorToMesh(model.mesh_device),
783
+ )
784
+ rotary = ttnn.from_torch(
785
+ torch.zeros((1, model.max_batch_size), dtype=torch.int32),
786
+ device=model.mesh_device,
787
+ dtype=ttnn.uint32,
788
+ layout=ttnn.ROW_MAJOR_LAYOUT,
789
+ mesh_mapper=ttnn.ReplicateTensorToMesh(model.mesh_device),
790
+ )
791
+ caches = generator._ensure_kv_cache()
792
+ model.bind_page_table(caches, generator._prefill_page_table)
793
+ cos, sin = model.rope_decode_tables(rotary)
794
+ out = MC.decoder_layer_decode_multichip(
795
+ hidden,
796
+ model.layers[0],
797
+ model.config,
798
+ model.ctx,
799
+ cos,
800
+ sin,
801
+ caches[0],
802
+ current_pos,
803
+ 0,
804
+ rope=model._rope_decode,
805
+ )
806
+ assert tuple(out.shape) == tuple(hidden.shape)
807
+ assert out.dtype == hidden.dtype
808
+ assert out.layout == hidden.layout
809
+ assert out.memory_config() == hidden.memory_config()
810
+
811
+
812
+ # --- prefill bucketing -------------------------------------------------------
813
+ #
814
+ # A prefill program is compiled per sequence length, so serving at the exact
815
+ # logical length makes the shape space 1..max_cache_len and leaves every new
816
+ # prompt length paying a fresh compile. Bucketing collapses that space to a
817
+ # ladder that ``warmup_model_prefill`` can enumerate. These tests hold the line
818
+ # that it is a *pure* latency change: the padding must be invisible in the
819
+ # result, which is what makes it safe to leave on by default.
820
+
821
+
822
+ @pytest.mark.parametrize("name,expected_head", [("pow2", 1024), ("pow2_half", 1024), ("1k", 1024)])
823
+ def test_prefill_bucket_ladders_are_finite_ascending_and_total(name, expected_head):
824
+ """Every admissible length must land on a rung, and the ladder must be small."""
825
+ ladder = prefill_bucket_ladder(name, 256000)
826
+ assert ladder == tuple(sorted(set(ladder))), "rungs must be ascending and unique"
827
+ assert ladder[0] == 128 and ladder[-1] == 256000, "the ladder must span the context"
828
+ assert 128 in ladder and expected_head in ladder
829
+ # Finiteness is the whole point -- an unbounded ladder cannot be warmed.
830
+ assert len(ladder) < 300
831
+
832
+
833
+ def test_exact_ladder_is_empty_and_a_typo_is_rejected(expect_error):
834
+ """A misspelled ladder must fail loudly: silently ignoring it would cost the
835
+ very multi-second compiles the ladder exists to remove."""
836
+ assert prefill_bucket_ladder("exact", 256000) == ()
837
+ with expect_error(ValueError, "QWEN3_PREFILL_BUCKETS"):
838
+ prefill_bucket_ladder("pow_2", 256000)
839
+
840
+
841
+ @pytest.mark.parametrize("prompt_len", [1, 33, 100, 127, 128, 129, 257, 700, 1000, 1537])
842
+ def test_bucketed_prefill_gives_the_same_token_as_exact_prefill(generator, prompt_len):
843
+ """The padding rows are causally invisible, so the sampled token must not move.
844
+
845
+ Measured on 4 dies at two layers: token-identical at every length above,
846
+ aligned or not, and PCC >= 0.9996 against the exact-length logits -- the
847
+ residual is bf16 reduction order, the same order-of-magnitude difference two
848
+ different SDPA chunkings of the same prompt produce.
849
+ """
850
+ tokens = torch.arange(1000, 1000 + prompt_len, dtype=torch.long).unsqueeze(0)
851
+
852
+ def run(ladder):
853
+ generator._prefill_buckets = ladder
854
+ generator.reset()
855
+ return generator.prefill_forward(
856
+ tokens,
857
+ page_table=generator.make_page_table([prompt_len + 1]),
858
+ kv_cache=generator._ensure_kv_cache(),
859
+ prompt_lens=[prompt_len],
860
+ sampling_mode="host",
861
+ )
862
+
863
+ try:
864
+ exact = run(())
865
+ bucketed = run(prefill_bucket_ladder("pow2_half", generator.model.max_cache_len))
866
+ finally:
867
+ generator._prefill_buckets = None
868
+ assert int(bucketed.argmax(-1).item()) == int(exact.argmax(-1).item())
869
+ passing, message = comp_pcc(exact.float(), bucketed.float(), 0.999)
870
+ assert passing, message
871
+
872
+
873
+ def test_bucketed_split_prefill_matches_exact_at_every_chunk_size(generator):
874
+ """The cached-suffix branch buckets too, at all of its reachable chunk sizes.
875
+
876
+ ``sdpa_chunk_size(start) == min(256, start & -start)`` over block-aligned
877
+ starts is a closed set, and this walks it. The assertion on
878
+ ``PREFILL_ATTENTION_BRANCHES`` is not decoration: a PCC of ~1.0 on a split
879
+ prefill has two explanations, and "the chunked branch never ran" is one.
880
+ """
881
+ block = generator.page_block_size
882
+ ladder = prefill_bucket_ladder("pow2_half", generator.model.max_cache_len)
883
+ for start in (block, 2 * block, 3 * block, 4 * block, 6 * block):
884
+ prompt_len = start + 17 # a suffix aligned to nothing
885
+ tokens = torch.arange(2000, 2000 + prompt_len, dtype=torch.long).unsqueeze(0)
886
+
887
+ def run(chosen):
888
+ generator._prefill_buckets = chosen
889
+ generator.reset()
890
+ FD.PREFILL_ATTENTION_BRANCHES["chunked"] = 0
891
+ out = generator.prefill_forward(
892
+ tokens,
893
+ page_table=generator.make_page_table([prompt_len + 1]),
894
+ kv_cache=generator._ensure_kv_cache(),
895
+ prompt_lens=[prompt_len],
896
+ sampling_mode="host",
897
+ start_pos=[start],
898
+ )
899
+ assert FD.PREFILL_ATTENTION_BRANCHES["chunked"] > 0, "the chunked branch never ran"
900
+ return out
901
+
902
+ try:
903
+ exact = run(())
904
+ bucketed = run(ladder)
905
+ finally:
906
+ generator._prefill_buckets = None
907
+ assert int(bucketed.argmax(-1).item()) == int(exact.argmax(-1).item()), f"start={start}"
908
+ passing, message = comp_pcc(exact.float(), bucketed.float(), 0.999)
909
+ assert passing, message
910
+
911
+
912
+ def test_bucketing_keeps_the_cache_write_at_the_real_length(generator):
913
+ """Padding must not reach the KV cache -- it would run off the page table.
914
+
915
+ This is the one way bucketing could corrupt a *different* request rather
916
+ than merely slow this one down. vLLM allocates ``ceil(real/block)`` blocks
917
+ and not one more, so a write at the padded length would land in whatever
918
+ physical block comes next, which belongs to somebody else. ``fill_len`` is
919
+ what prevents that, and this asserts the block after the prompt's own is
920
+ still byte-for-byte what it was.
921
+ """
922
+ prompt_len = 129 # 5 blocks at a 32-token block; the 512 rung wants 16
923
+ block = generator.page_block_size
924
+ owned = math.ceil(prompt_len / block)
925
+ caches = generator._ensure_kv_cache()
926
+ assert generator.prefill_padded_len(prompt_len, start=0) > owned * block, "test needs real padding"
927
+
928
+ # Give the user exactly the blocks it is entitled to, and nothing after.
929
+ page_table = torch.full((generator.batch, generator.pages_per_user), -1, dtype=torch.int32)
930
+ page_table[0, :owned] = torch.arange(owned, dtype=torch.int32)
931
+ sentinel_block = owned
932
+
933
+ def read_sentinel():
934
+ host = _first_device_to_torch(caches[0].k)
935
+ return host[sentinel_block].clone()
936
+
937
+ generator.reset()
938
+ before = read_sentinel()
939
+ generator._prefill_buckets = prefill_bucket_ladder("pow2_half", generator.model.max_cache_len)
940
+ try:
941
+ generator.prefill_forward(
942
+ torch.arange(4000, 4000 + prompt_len, dtype=torch.long).unsqueeze(0),
943
+ page_table=page_table,
944
+ kv_cache=caches,
945
+ prompt_lens=[prompt_len],
946
+ sampling_mode="host",
947
+ )
948
+ finally:
949
+ generator._prefill_buckets = None
950
+ assert torch.equal(read_sentinel(), before), "bucket padding was written past the prompt's blocks"
951
+
952
+
953
+ @pytest.mark.parametrize("prompt_len", [128, 256, 1024])
954
+ def test_a_prompt_filling_the_whole_rope_table_does_not_free_it(generator, prompt_len):
955
+ """A power-of-two prompt must not leave the model's cos/sin tables freed.
956
+
957
+ ``ensure_rope_capacity`` rounds up to a power of two, so at these lengths
958
+ the window ``prefill_hidden`` wants is the *entire* table -- and
959
+ ``ttnn.slice`` returns a view, not a copy, when asked for the whole tensor.
960
+ Deallocating that view frees ``cos_table`` itself, and the failure lands on
961
+ the *next* prefill as "Input Tensor is not allocated", nowhere near the
962
+ cause. Two prefills, because one alone cannot see it.
963
+ """
964
+ tokens = torch.arange(1000, 1000 + prompt_len, dtype=torch.long).unsqueeze(0)
965
+ for _ in range(2):
966
+ generator.reset()
967
+ logits = generator.prefill_forward(
968
+ tokens,
969
+ page_table=generator.make_page_table([prompt_len + 1]),
970
+ kv_cache=generator._ensure_kv_cache(),
971
+ prompt_lens=[prompt_len],
972
+ sampling_mode="host",
973
+ )
974
+ assert torch.isfinite(logits).all()
975
+ assert generator.model.cos_table.is_allocated()
976
+ assert generator.model.sin_table.is_allocated()
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_moe.py ADDED
@@ -0,0 +1,232 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """TTNN MoE block vs the HuggingFace reference. Router first, experts after.
5
+
6
+ The router is tested on its own because it is the one place in this model where
7
+ a *discrete* disagreement is possible. Everything else degrades smoothly with
8
+ precision; top-k selection does not. HF softmaxes over all 128 experts in fp32
9
+ and we run bf16, so two experts with near-equal probability can swap places.
10
+ When that happens the token is routed to a genuinely different expert and the
11
+ MoE output for that token is unrelated to the reference -- a handful of such
12
+ tokens drags whole-block PCC down in a way that looks like a numerics problem
13
+ but is actually a selection problem.
14
+
15
+ ``test_router_selection_matches`` separates the two by comparing the chosen
16
+ expert *sets* directly, so a later PCC dip can be attributed rather than
17
+ guessed at.
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ import pytest
23
+ import torch
24
+ from loguru import logger
25
+
26
+ import ttnn
27
+ from models.common.utility_functions import comp_allclose, comp_pcc
28
+
29
+ from ..tt.functional_decoder import (
30
+ MoEConfig,
31
+ build_expert_sparsity,
32
+ moe_prefill,
33
+ router_forward,
34
+ upload_expert_weights,
35
+ upload_router_weight,
36
+ )
37
+ from ..tt.weight_mapping import convert_moe_weights
38
+ from .reference import build_reference_layer, layer_state_dict
39
+
40
+ LAYER_IDX = 0
41
+ PCC_REQUIRED = 0.99
42
+
43
+
44
+ @pytest.fixture(scope="module")
45
+ def reference():
46
+ return build_reference_layer(LAYER_IDX)
47
+
48
+
49
+ @pytest.fixture(scope="module")
50
+ def torch_weights():
51
+ return convert_moe_weights(layer_state_dict(LAYER_IDX), n_experts=128)
52
+
53
+
54
+ def _hidden(config, seq_len, seed=0):
55
+ torch.manual_seed(seed)
56
+ return torch.randn(1, seq_len, config.hidden_size, dtype=torch.float32) * 0.02
57
+
58
+
59
+ def _reference_router(layer, hidden):
60
+ """Return ``(dense [S, E], indices [S, k])`` from the reference router."""
61
+ flat = hidden.view(-1, hidden.shape[-1])
62
+ with torch.no_grad():
63
+ _, scores, indices = layer.mlp.gate(flat)
64
+ dense = torch.zeros(flat.shape[0], 128, dtype=scores.dtype)
65
+ dense.scatter_(-1, indices, scores)
66
+ return dense, indices
67
+
68
+
69
+ def _run_router(mesh_device, hf_config, torch_weights, hidden):
70
+ w = upload_router_weight(torch_weights["router"], mesh_device)
71
+ tt_in = ttnn.from_torch(
72
+ hidden.unsqueeze(0),
73
+ dtype=ttnn.bfloat16,
74
+ layout=ttnn.TILE_LAYOUT,
75
+ device=mesh_device,
76
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
77
+ )
78
+ out = router_forward(tt_in, w, MoEConfig.from_hf(hf_config))
79
+ return ttnn.to_torch(out).reshape(-1, 128).float()
80
+
81
+
82
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
83
+ @pytest.mark.parametrize("seq_len", [32, 128], ids=["s32", "s128"])
84
+ def test_router_dense_weights_vs_reference(mesh_device, reference, torch_weights, seq_len):
85
+ layer, hf_config = reference
86
+ hidden = _hidden(hf_config, seq_len)
87
+
88
+ ref_dense, _ = _reference_router(layer, hidden)
89
+ tt_dense = _run_router(mesh_device, hf_config, torch_weights, hidden)
90
+
91
+ passing, pcc_message = comp_pcc(ref_dense, tt_dense, PCC_REQUIRED)
92
+ logger.info(comp_allclose(ref_dense, tt_dense))
93
+ logger.info(f"router dense seq={seq_len}: {pcc_message}")
94
+ assert passing, f"router dense weights (seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
95
+
96
+
97
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
98
+ def test_router_selection_matches(mesh_device, reference, torch_weights):
99
+ """How many tokens pick a different set of 8 experts than the reference.
100
+
101
+ Reported explicitly rather than silently folded into PCC, because a
102
+ selection flip and a numerics error need completely different fixes.
103
+
104
+ The bound is measured, not guessed. Holding the softmax in fp32 and letting
105
+ only the projection run in bf16 costs 5/128 tokens on this checkpoint (a
106
+ host-side simulation of the same arithmetic agrees), so the floor is ~4%.
107
+ The threshold sits at 10% to leave room for a different activation sample
108
+ while still catching the regression that matters: dropping the softmax to
109
+ bf16 sends this straight to ~65%.
110
+ """
111
+ seq_len = 128
112
+ layer, hf_config = reference
113
+ hidden = _hidden(hf_config, seq_len)
114
+
115
+ ref_dense, ref_indices = _reference_router(layer, hidden)
116
+ tt_dense = _run_router(mesh_device, hf_config, torch_weights, hidden)
117
+ tt_indices = tt_dense.topk(8, dim=-1).indices
118
+
119
+ mismatched = 0
120
+ worst_missed_weight = 0.0
121
+ for token in range(seq_len):
122
+ ref_set = set(ref_indices[token].tolist())
123
+ tt_set = set(tt_indices[token].tolist())
124
+ if ref_set != tt_set:
125
+ mismatched += 1
126
+ # How much routing weight the reference put on experts we skipped.
127
+ # Near-ties sit at the bottom of the top-8, so this should be small;
128
+ # a large value means a genuinely wrong selection, not rounding.
129
+ for e in ref_set - tt_set:
130
+ worst_missed_weight = max(worst_missed_weight, float(ref_dense[token, e]))
131
+
132
+ logger.info(
133
+ f"router selection: {mismatched}/{seq_len} tokens differ from the fp32 reference; "
134
+ f"largest missed routing weight {worst_missed_weight:.4f}"
135
+ )
136
+ assert mismatched <= seq_len * 0.10, (
137
+ f"{mismatched}/{seq_len} tokens routed to a different expert set -- "
138
+ "far above the ~4% bf16-projection floor. Check that the router softmax "
139
+ "and topk are still running in fp32 before suspecting anything else."
140
+ )
141
+
142
+
143
+ def _upload_dense(dense: torch.Tensor, mesh_device):
144
+ return ttnn.from_torch(
145
+ dense.reshape(1, 1, *dense.shape).float(),
146
+ dtype=ttnn.bfloat16,
147
+ layout=ttnn.TILE_LAYOUT,
148
+ device=mesh_device,
149
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
150
+ )
151
+
152
+
153
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
154
+ @pytest.mark.parametrize("seq_len", [32, 128], ids=["s32", "s128"])
155
+ def test_experts_with_reference_routing(mesh_device, reference, torch_weights, seq_len):
156
+ """Expert math alone, with the router's own selection taken out of play.
157
+
158
+ The reference's exact fp32 routing weights are fed straight to our experts,
159
+ so any shortfall here is the sparse_matmul / SwiGLU / reduce path and not
160
+ the ~5/128 near-tie selection differences the router legitimately has.
161
+ Isolating the two is what makes the end-to-end number interpretable.
162
+ """
163
+ layer, hf_config = reference
164
+ config = MoEConfig.from_hf(hf_config)
165
+ hidden = _hidden(hf_config, seq_len)
166
+
167
+ with torch.no_grad():
168
+ ref_out = layer.mlp(hidden)
169
+
170
+ ref_dense, _ = _reference_router(layer, hidden)
171
+ weights = upload_expert_weights(torch_weights, mesh_device, config)
172
+ sparsity = build_expert_sparsity(mesh_device, config.num_experts)
173
+
174
+ tt_in = ttnn.from_torch(
175
+ hidden.unsqueeze(0),
176
+ dtype=ttnn.bfloat16,
177
+ layout=ttnn.TILE_LAYOUT,
178
+ device=mesh_device,
179
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
180
+ )
181
+ tt_out = moe_prefill(tt_in, _upload_dense(ref_dense, mesh_device), weights, config, sparsity)
182
+ tt_out_torch = ttnn.to_torch(tt_out).squeeze(0)
183
+
184
+ passing, pcc_message = comp_pcc(ref_out, tt_out_torch, PCC_REQUIRED)
185
+ logger.info(comp_allclose(ref_out, tt_out_torch))
186
+ logger.info(f"experts (reference routing) seq={seq_len}: {pcc_message}")
187
+ assert passing, f"experts (seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
188
+
189
+
190
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
191
+ @pytest.mark.parametrize("seq_len", [32, 128], ids=["s32", "s128"])
192
+ def test_moe_block_end_to_end(mesh_device, reference, torch_weights, seq_len):
193
+ """Our router driving our experts, against the reference MoE block."""
194
+ layer, hf_config = reference
195
+ config = MoEConfig.from_hf(hf_config)
196
+ hidden = _hidden(hf_config, seq_len)
197
+
198
+ with torch.no_grad():
199
+ ref_out = layer.mlp(hidden)
200
+
201
+ w_router = upload_router_weight(torch_weights["router"], mesh_device)
202
+ weights = upload_expert_weights(torch_weights, mesh_device, config)
203
+ sparsity = build_expert_sparsity(mesh_device, config.num_experts)
204
+
205
+ tt_in = ttnn.from_torch(
206
+ hidden.unsqueeze(0),
207
+ dtype=ttnn.bfloat16,
208
+ layout=ttnn.TILE_LAYOUT,
209
+ device=mesh_device,
210
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
211
+ )
212
+ routing = router_forward(tt_in, w_router, config)
213
+ tt_out = moe_prefill(tt_in, routing, weights, config, sparsity)
214
+ tt_out_torch = ttnn.to_torch(tt_out).squeeze(0)
215
+
216
+ passing, pcc_message = comp_pcc(ref_out, tt_out_torch, PCC_REQUIRED)
217
+ logger.info(comp_allclose(ref_out, tt_out_torch))
218
+ logger.info(f"MoE block end-to-end seq={seq_len}: {pcc_message}")
219
+ assert passing, f"MoE block (seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
220
+
221
+
222
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
223
+ def test_router_weights_sum_to_one(mesh_device, reference, torch_weights):
224
+ """norm_topk_prob=True, so each token's 8 weights must renormalise to 1."""
225
+ layer, hf_config = reference
226
+ hidden = _hidden(hf_config, 32)
227
+ tt_dense = _run_router(mesh_device, hf_config, torch_weights, hidden)
228
+
229
+ sums = tt_dense.sum(dim=-1)
230
+ nonzero = (tt_dense > 0).sum(dim=-1)
231
+ assert torch.allclose(sums, torch.ones_like(sums), atol=2e-2), f"router weights not normalised: {sums[:4]}"
232
+ assert (nonzero == 8).all(), f"expected exactly 8 active experts per token, got {nonzero.unique().tolist()}"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_multichip_decoder.py ADDED
@@ -0,0 +1,1057 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Correctness of the multichip decoder on the full 4-die P300_X2 mesh.
5
+
6
+ The reference these tests compare against is the **single-chip TTNN optimized
7
+ decoder**, not HuggingFace, and it is run *on the same mesh with every tensor
8
+ replicated*. That is the whole trick of this file: a mesh op is SPMD, so
9
+ uploading the unsharded stage-02 weights with ``ReplicateTensorToMesh`` makes
10
+ each of the four dies independently compute the exact single-chip answer, in the
11
+ same process, from the same host tensors, with the same program cache. The
12
+ comparison then isolates sharding and collective bugs from every source of
13
+ numerical difference that HF-vs-TTNN would drag in.
14
+
15
+ ``test_baseline_upload_is_actually_replicated`` is what stops that reference
16
+ quietly becoming meaningless -- if ``from_torch`` ever stopped replicating, the
17
+ baseline would still produce *a* number and every PCC below would still pass.
18
+
19
+ Two tests here are load-bearing in a way their size does not suggest:
20
+
21
+ * ``test_topk_is_identical_across_dies``. The expert-parallel scheme assumes the
22
+ four dies agree on the global top-8 from bit-identical replicated logits, so
23
+ that the four 32-expert windows partition it. If ``ttnn.topk`` ever broke a
24
+ tie differently on one die the layer would be **silently wrong** -- no shape
25
+ error, no assert, just PCC drift -- so the property is asserted directly
26
+ rather than argued from "same program, same input".
27
+
28
+ * ``test_expert_window_can_be_empty``. Under EP the locally-live expert count is
29
+ data-dependent in 0..8, which is why decode must pass ``nnz=None``. The zero
30
+ case is the one that never happens by accident in a random test and is exactly
31
+ where an uninitialised output buffer would leak a NaN into the all-reduce.
32
+
33
+ Every test opens the mesh with ``fabric_config=FABRIC_1D_RING``. Without it the
34
+ collectives have no fabric to run on; with ``FABRIC_1D`` they would run but on
35
+ the linear topology the CCL sweep measured 1.2-1.8x slower.
36
+ """
37
+
38
+ from __future__ import annotations
39
+
40
+ import pytest
41
+ import torch
42
+ from loguru import logger
43
+
44
+ import ttnn
45
+ from models.common.modules.tt_ccl import default_topology
46
+ from models.common.utility_functions import comp_pcc
47
+
48
+ from ..tt import functional_decoder as F
49
+ from ..tt import multichip_decoder as MC
50
+ from ..tt import optimized_decoder as O
51
+ from ..tt.weight_mapping import convert_layer_weights
52
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
53
+
54
+ LAYER_IDX = 0
55
+ # Against the replicated single-chip baseline only sharding and the collectives
56
+ # differ, so 0.99 -- the threshold used against HF, where dtype and kernel choice
57
+ # differ too -- was two orders of magnitude looser than the measured margin. The
58
+ # actuals span 0.99945 (two stacked layers, the only one below 0.9996) to
59
+ # 1.0; 0.999 sits below the worst of them and still catches anything that would
60
+ # make the sharding wrong rather than merely different.
61
+ PCC_VS_SINGLE_CHIP = 0.999
62
+ PCC_VS_HF = 0.995
63
+ MAX_SEQ = 1024
64
+ BLOCK_SIZE = 32
65
+ TRACE_REGION_SIZE = 90000000
66
+
67
+ # Ring fabric must be configured before the mesh is opened, which is what this
68
+ # indirect parametrisation does (conftest.set_fabric runs ahead of the open).
69
+ MESH_PARAMS = {"trace_region_size": TRACE_REGION_SIZE, "fabric_config": ttnn.FabricConfig.FABRIC_1D_RING}
70
+ mesh_4 = pytest.mark.parametrize("mesh_device", [MC.MESH_SHAPE], ids=["1x4"], indirect=True)
71
+ ring_fabric = pytest.mark.parametrize("device_params", [MESH_PARAMS], indirect=True)
72
+
73
+
74
+ @pytest.fixture(scope="module")
75
+ def reference():
76
+ return build_reference_layer(LAYER_IDX)
77
+
78
+
79
+ @pytest.fixture(scope="module")
80
+ def torch_weights(reference):
81
+ _, hf_config = reference
82
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
83
+
84
+
85
+ def _hidden(hf_config, seq_len, seed=0):
86
+ torch.manual_seed(seed)
87
+ return torch.randn(1, 1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
88
+
89
+
90
+ def _reference_layer(layer, hf_config, hidden):
91
+ """HF layer output for a ``[1, 1, S, H]`` input, returned as ``[1, S, H]``."""
92
+ hidden = hidden.reshape(1, -1, hf_config.hidden_size)
93
+ seq_len = hidden.shape[1]
94
+ cos, sin = rotary_embeddings(hf_config, seq_len)
95
+ mask = torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
96
+ with torch.no_grad():
97
+ out = layer(hidden, position_embeddings=(cos, sin), attention_mask=mask)
98
+ return out[0] if isinstance(out, tuple) else out
99
+
100
+
101
+ def _replicate(t, mesh_device, dtype=ttnn.bfloat16):
102
+ return ttnn.from_torch(
103
+ t,
104
+ dtype=dtype,
105
+ layout=ttnn.TILE_LAYOUT,
106
+ device=mesh_device,
107
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
108
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
109
+ )
110
+
111
+
112
+ def _per_die(t, mesh_device, dim: int = 0) -> torch.Tensor:
113
+ """All four dies' copies of a tensor, concatenated along ``dim``.
114
+
115
+ Every activation in this layer is ``[1, 1, ., .]``, so concatenating on dim 0
116
+ puts die *d* at index *d* and reads like a stack -- but it *is* a
117
+ concatenation, which matters for the KV cache, whose dim 0 is the block index
118
+ and which therefore has to be reassembled on the head axis instead.
119
+ """
120
+ return ttnn.to_torch(t, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=dim))
121
+
122
+
123
+ class Fixture:
124
+ """Both paths uploaded side by side onto the same mesh.
125
+
126
+ ``baseline`` is the stage-02 optimized decoder with every tensor replicated,
127
+ so each die computes the full single-chip layer; ``multichip`` is the sharded
128
+ stage-03 path. They share the router weight tensor, the RoPE caches and the
129
+ host weights, so the only difference between them is the parallelisation.
130
+ """
131
+
132
+ def __init__(self, mesh_device, hf_config, torch_weights):
133
+ self.mesh = mesh_device
134
+ self.hf = hf_config
135
+ self.config = MC.MeshDecoderConfig.from_hf(hf_config)
136
+ self.ctx = MC.mesh_context(mesh_device)
137
+ self.torch_router = torch_weights["router"]
138
+ self.multichip = MC.upload_multichip_weights(torch_weights, mesh_device, self.config)
139
+ self.baseline_experts = O.upload_optimized_weights(torch_weights, mesh_device, self.config.global_config.moe)
140
+ self.baseline = F.DecoderLayerWeights(
141
+ input_layernorm=self.multichip.input_layernorm,
142
+ post_attention_layernorm=self.multichip.post_attention_layernorm,
143
+ attention=None, # the optimized path reads OptimizedWeights.attention
144
+ router=self.multichip.router,
145
+ experts=None,
146
+ )
147
+ self.cos, self.sin = F.build_rope_cache(hf_config, MAX_SEQ, mesh_device)
148
+ self.baseline_sparsity = F.build_expert_sparsity(mesh_device, self.config.global_config.moe.num_experts)
149
+ self.sparsity = MC.build_local_sparsity(mesh_device, self.config.local_moe)
150
+
151
+ def rep(self, t):
152
+ return _replicate(t, self.mesh)
153
+
154
+ def dies(self, t, dim: int = 0):
155
+ return _per_die(t, self.mesh, dim)
156
+
157
+ def baseline_prefill(self, x, kv_cache=None, user_id=0):
158
+ return O.decoder_layer_prefill_optimized(
159
+ self.rep(x),
160
+ self.baseline,
161
+ self.config.global_config,
162
+ self.cos,
163
+ self.sin,
164
+ self.baseline_sparsity,
165
+ self.baseline_experts,
166
+ kv_cache=kv_cache,
167
+ user_id=user_id,
168
+ )
169
+
170
+ def multichip_prefill(self, x, kv_cache=None, user_id=0):
171
+ return MC.decoder_layer_prefill_multichip(
172
+ self.rep(x),
173
+ self.multichip,
174
+ self.config,
175
+ self.ctx,
176
+ self.cos,
177
+ self.sin,
178
+ self.sparsity,
179
+ kv_cache=kv_cache,
180
+ user_id=user_id,
181
+ )
182
+
183
+
184
+ @pytest.fixture
185
+ def fixture(mesh_device, reference, torch_weights):
186
+ _, hf_config = reference
187
+ return Fixture(mesh_device, hf_config, torch_weights)
188
+
189
+
190
+ # --- host-side weight transforms ---------------------------------------------
191
+
192
+
193
+ def test_wqkv_column_split_is_head_interleaved(reference, torch_weights):
194
+ """Die *d* must get Q heads 8d..8d+7 plus K head d and V head d.
195
+
196
+ A contiguous 4-way split of the checkpoint's ``[Wq | Wk | Wv]`` gives die 0
197
+ nothing but Q heads and die 3 nothing but K and V, and produces **no shape
198
+ error** -- ``nlp_create_qkv_heads_decode(num_heads=8, num_kv_heads=1)``
199
+ accepts 1280 columns whatever is in them. This is a host-only test because
200
+ that is where the bug would live and where it is cheapest to catch.
201
+ """
202
+ _, hf_config = reference
203
+ cfg = MC.MeshDecoderConfig.from_hf(hf_config)
204
+ a = cfg.global_config.attention
205
+ n, hd = cfg.num_devices, a.head_dim
206
+ full = torch_weights["wqkv"].reshape(a.hidden_size, -1)
207
+ permuted = MC.head_interleaved_wqkv(full, a, n)
208
+
209
+ q_end = a.num_attention_heads * hd
210
+ k_end = q_end + a.num_key_value_heads * hd
211
+ per_die = permuted.shape[-1] // n
212
+ q_per = a.num_attention_heads // n
213
+
214
+ for d in range(n):
215
+ shard = permuted[:, d * per_die : (d + 1) * per_die]
216
+ assert shard.shape[-1] == q_per * hd + 2 * hd, shard.shape
217
+ expect_q = full[:, d * q_per * hd : (d + 1) * q_per * hd]
218
+ expect_k = full[:, q_end + d * hd : q_end + (d + 1) * hd]
219
+ expect_v = full[:, k_end + d * hd : k_end + (d + 1) * hd]
220
+ assert torch.equal(shard[:, : q_per * hd], expect_q), f"die {d} Q heads"
221
+ assert torch.equal(shard[:, q_per * hd : q_per * hd + hd], expect_k), f"die {d} K head"
222
+ assert torch.equal(shard[:, q_per * hd + hd :], expect_v), f"die {d} V head"
223
+
224
+ # And it is a permutation, not a rewrite: every column survives exactly once.
225
+ # Compared as sorted multisets rather than row sums -- the reordering changes
226
+ # float addition order, so ``sum`` differs in the last bits even when nothing
227
+ # has been lost, and asserting on it fails for the wrong reason.
228
+ assert torch.equal(permuted.sort(dim=-1).values, full.sort(dim=-1).values)
229
+ logger.info(f"wqkv head-interleaved split verified for {n} dies, per-die N = {per_die}")
230
+
231
+
232
+ def test_repo_default_topology_is_wrong_for_this_mesh():
233
+ """Documents *why* Ring is passed explicitly, so it cannot be "simplified" away.
234
+
235
+ ``tt_ccl.default_topology()`` only returns ``Ring`` for 8-device T3K and
236
+ Galaxy; for this 4-device Blackhole mesh it returns ``Linear``, which the CCL
237
+ sweep measured at 1.21x slower at decode size and 1.79x at 2 MB. The module
238
+ constant must therefore disagree with the helper.
239
+ """
240
+ assert MC.TOPOLOGY is ttnn.Topology.Ring
241
+ assert MC.NUM_LINKS == 2
242
+ logger.info(f"multichip_decoder overrides default_topology (callable: {default_topology.__name__})")
243
+
244
+
245
+ # --- the replicated baseline itself ------------------------------------------
246
+
247
+
248
+ @ring_fabric
249
+ @mesh_4
250
+ def test_baseline_upload_is_actually_replicated(fixture):
251
+ """The single-chip reference is only a reference if all four dies agree.
252
+
253
+ Everything else in this file divides by this. If ``ReplicateTensorToMesh``
254
+ ever stopped replicating, or the mesh stopped being SPMD, the baseline would
255
+ still return numbers and every PCC below would still pass against them.
256
+ """
257
+ out = fixture.dies(fixture.baseline_prefill(_hidden(fixture.hf, 128)))
258
+ spread = (out - out[0:1]).abs().max().item()
259
+ logger.info(f"replicated single-chip baseline: max spread across 4 dies = {spread:.3e}")
260
+ assert spread == 0.0, f"replicated baseline differs across dies by {spread}; it is not a valid reference"
261
+
262
+
263
+ # --- the determinism assumption the whole scheme rests on --------------------
264
+
265
+
266
+ @ring_fabric
267
+ @mesh_4
268
+ @pytest.mark.parametrize("seed", [0, 1, 2, 3, 4, 5, 6, 7])
269
+ def test_topk_is_identical_across_dies(fixture, seed):
270
+ """All four dies must select the same 8 experts from the replicated logits.
271
+
272
+ Expert parallelism partitions the 128 experts into four 32-wide windows and
273
+ each die keeps only the winners inside its own. That is the global top-8 only
274
+ if the four dies agree; if they disagree the layer double-counts some experts
275
+ and drops others, with no error of any kind. Random inputs are checked, and
276
+ so is an all-zero input, where every logit is the router's bias-free
277
+ projection of zero and the top-8 is decided **entirely by tie-breaking** --
278
+ the degenerate case an ordinary test never reaches.
279
+ """
280
+ hidden = _hidden(fixture.hf, 128, seed=seed) if seed else torch.zeros(1, 1, 128, fixture.hf.hidden_size)
281
+ logits = ttnn.linear(
282
+ fixture.rep(hidden), fixture.multichip.router, dtype=ttnn.float32, memory_config=ttnn.DRAM_MEMORY_CONFIG
283
+ )
284
+ _, indices = ttnn.topk(logits, k=fixture.config.global_config.moe.num_experts_per_tok, dim=-1, sorted=True)
285
+ per_die = fixture.dies(indices)
286
+ for d in range(1, fixture.config.num_devices):
287
+ assert torch.equal(per_die[0], per_die[d]), (
288
+ f"seed={seed}: die {d} selected different experts than die 0 "
289
+ f"({int((per_die[0] != per_die[d]).sum())} of {per_die[0].numel()} slots differ). "
290
+ "The four expert windows are no longer a partition of the global top-8."
291
+ )
292
+ logger.info(f"topk seed={seed}: 4 dies bit-identical over {per_die[0].numel()} selections")
293
+
294
+
295
+ @ring_fabric
296
+ @mesh_4
297
+ @pytest.mark.parametrize("seq_len", [1, 33, 128], ids=["decode", "s33", "s128"])
298
+ def test_router_windows_partition_global_routing(fixture, seq_len):
299
+ """Concatenating the four local windows must reproduce the global dense routing.
300
+
301
+ This is the direct statement of the EP contract: the multichip router returns
302
+ ``[1, 1, S, 32]`` per die, and stitching them in device order must equal what
303
+ the single-chip router returns as one ``[1, 1, S, 128]`` row -- same experts,
304
+ same weights, normalised by the same global denominator.
305
+ """
306
+ x = fixture.rep(_hidden(fixture.hf, seq_len))
307
+ moe = fixture.config.global_config.moe
308
+ global_dense = fixture.dies(O.router_forward_optimized(x, fixture.multichip.router, moe))[0].float()
309
+ local = fixture.dies(
310
+ MC.router_forward_multichip(
311
+ x, fixture.multichip.router, fixture.multichip.expert_window, moe, fixture.config.local_moe
312
+ )
313
+ ).float()
314
+
315
+ n_local = fixture.config.local_moe.num_experts
316
+ stitched = torch.cat([local[d].reshape(-1, n_local) for d in range(fixture.config.num_devices)], dim=-1)
317
+ reference = global_dense.reshape(-1, moe.num_experts)
318
+ delta = (stitched - reference).abs().max().item()
319
+ logger.info(f"router windows seq={seq_len}: max |stitched - global| = {delta:.3e}")
320
+ assert torch.equal(stitched > 0, reference > 0), "the four windows do not select the global top-8"
321
+ assert delta == 0.0, f"routing weights differ by {delta}; the window matmul is not exact"
322
+ assert ((stitched > 0).sum(dim=-1) == moe.num_experts_per_tok).all()
323
+
324
+
325
+ # --- prefill ------------------------------------------------------------------
326
+
327
+
328
+ @ring_fabric
329
+ @mesh_4
330
+ @pytest.mark.parametrize("seq_len", [32, 128, 512, 33, 100, 257], ids=["s32", "s128", "s512", "s33", "s100", "s257"])
331
+ def test_multichip_prefill_vs_single_chip(fixture, seq_len):
332
+ """Prefill against the single-chip TTNN baseline, aligned and non-aligned.
333
+
334
+ The non-aligned lengths are the point of the parametrisation: nothing in the
335
+ multichip path may turn a decoder that accepted any prompt length into one
336
+ that only accepts multiples of a chunk, tile, page or collective block. The
337
+ collectives scatter on dim 3 (hidden, 2048), which is independent of S.
338
+ """
339
+ hidden = _hidden(fixture.hf, seq_len)
340
+ base = fixture.dies(fixture.baseline_prefill(hidden))[0:1].float()
341
+ multi = fixture.dies(fixture.multichip_prefill(hidden))
342
+
343
+ spread = (multi - multi[0:1]).abs().max().item()
344
+ assert spread == 0.0, f"S={seq_len}: layer output differs across dies by {spread}; the all-reduce is not complete"
345
+ assert tuple(multi.shape) == (fixture.config.num_devices, 1, seq_len, fixture.hf.hidden_size), (
346
+ f"S={seq_len}: four dies of [1,1,S,H] concatenated on dim 0 came back {tuple(multi.shape)}; "
347
+ "the replicated layer contract is not intact"
348
+ )
349
+
350
+ passing, message = comp_pcc(base, multi[0:1].float(), PCC_VS_SINGLE_CHIP)
351
+ logger.info(f"multichip prefill S={seq_len} vs single-chip TTNN: {message}")
352
+ assert passing, f"multichip prefill S={seq_len} below {PCC_VS_SINGLE_CHIP} vs single-chip: {message}"
353
+
354
+
355
+ @ring_fabric
356
+ @mesh_4
357
+ @pytest.mark.parametrize("seq_len", [128, 33], ids=["s128", "s33"])
358
+ def test_multichip_prefill_vs_hf(fixture, reference, seq_len):
359
+ """The end-to-end bar: the same 0.995 PCC against HF the single chip clears."""
360
+ layer, hf_config = reference
361
+ hidden = _hidden(hf_config, seq_len)
362
+ ref = _reference_layer(layer, hf_config, hidden)
363
+ multi = fixture.dies(fixture.multichip_prefill(hidden))[0].reshape(1, seq_len, hf_config.hidden_size)
364
+ passing, message = comp_pcc(ref, multi.float(), PCC_VS_HF)
365
+ logger.info(f"multichip prefill S={seq_len} vs HF: {message}")
366
+ assert passing, f"multichip prefill S={seq_len} below {PCC_VS_HF} vs HF: {message}"
367
+
368
+
369
+ @ring_fabric
370
+ @mesh_4
371
+ def test_multichip_prefill_is_deterministic(fixture):
372
+ """Bitwise repeatability, including through two collectives per layer."""
373
+ hidden = _hidden(fixture.hf, 128)
374
+ outs = [fixture.dies(fixture.multichip_prefill(hidden)).clone() for _ in range(3)]
375
+ assert torch.equal(outs[0], outs[1]), "multichip prefill run 1 != run 2 (bitwise)"
376
+ assert torch.equal(outs[0], outs[2]), "multichip prefill run 1 != run 3 (bitwise)"
377
+ logger.info("multichip prefill: 3 runs bit-identical on all 4 dies")
378
+
379
+
380
+ # --- KV cache and decode ------------------------------------------------------
381
+
382
+
383
+ @ring_fabric
384
+ @mesh_4
385
+ def test_local_kv_cache_layout(fixture):
386
+ """Each die owns exactly one KV head, and the four together hold the whole cache.
387
+
388
+ This is the memory half of the TP decision: 512 B per token per layer per die
389
+ instead of 2048. The test does not merely check the *shape* -- it prefills
390
+ both paths from the same prompt and asserts that stacking the four dies' K
391
+ caches on the head axis reproduces the single-chip cache, which is what
392
+ proves the head *assignment* matches the wqkv column split rather than just
393
+ the head count.
394
+ """
395
+ cfg = fixture.config
396
+ base_kv = F.create_kv_cache(fixture.mesh, cfg.global_config.attention, 1, 128, block_size=BLOCK_SIZE)
397
+ mc_kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, 1, 128, block_size=BLOCK_SIZE)
398
+
399
+ assert mc_kv.k.shape[1] == 1, f"per-die KV cache has {mc_kv.k.shape[1]} heads, expected 1"
400
+ assert base_kv.k.shape[1] == cfg.global_config.attention.num_key_value_heads
401
+ assert mc_kv.is_paged and base_kv.is_paged
402
+
403
+ hidden = _hidden(fixture.hf, 64)
404
+ fixture.baseline_prefill(hidden, kv_cache=base_kv)
405
+ fixture.multichip_prefill(hidden, kv_cache=mc_kv)
406
+
407
+ # The cache's dim 0 is the physical block index, so the four dies are
408
+ # reassembled on the *head* axis -- which is also exactly what makes this a
409
+ # test of head ownership rather than of head count.
410
+ n_kv = cfg.global_config.attention.num_key_value_heads
411
+ # The baseline cache is replicated, so any one die's copy is the reference;
412
+ # take die 0's four heads out of the 16 the concat produces.
413
+ base_k = fixture.dies(base_kv.k, dim=1)[:, :n_kv].float() # [blocks, 4, block, head_dim]
414
+ stitched = fixture.dies(mc_kv.k, dim=1).float() # 4 x [blocks, 1, block, head_dim]
415
+ assert stitched.shape == base_k.shape, (stitched.shape, base_k.shape)
416
+ passing, message = comp_pcc(base_k, stitched, 0.999)
417
+ logger.info(f"local KV head layout: 4x[.,1,.,.] stitched vs single-chip [.,4,.,.]: {message}")
418
+ assert passing, f"per-die KV heads are not the single-chip heads in device order: {message}"
419
+
420
+ bytes_per_token = mc_kv.k.shape[1] * cfg.local_attention.head_dim * 2 * 2
421
+ logger.info(f"per-die KV: {bytes_per_token} B/token/layer (single chip: {bytes_per_token * 4})")
422
+ assert bytes_per_token == 512
423
+
424
+
425
+ @ring_fabric
426
+ @mesh_4
427
+ @pytest.mark.parametrize("block_size", [None, 32], ids=["contiguous", "paged32"])
428
+ def test_multichip_decode_vs_single_chip(fixture, block_size):
429
+ """One decode step against the single-chip baseline, both cache modes.
430
+
431
+ Both paths are prefilled with the same prompt into their own caches, so this
432
+ covers the paged write path (``paged_fill_cache``), the paged update
433
+ (``paged_update_cache`` at 1 KV head, which the design phase flagged as
434
+ unexercised), the page table, and ``cur_pos_tensor``, not just the matmuls.
435
+ """
436
+ cfg = fixture.config
437
+ prompt = 32
438
+ full = _hidden(fixture.hf, prompt + 1)
439
+ base_kv = F.create_kv_cache(fixture.mesh, cfg.global_config.attention, 1, MAX_SEQ, block_size=block_size)
440
+ mc_kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, 1, MAX_SEQ, block_size=block_size)
441
+ fixture.baseline_prefill(full[:, :, :prompt, :], kv_cache=base_kv)
442
+ fixture.multichip_prefill(full[:, :, :prompt, :], kv_cache=mc_kv)
443
+
444
+ pos = ttnn.from_torch(
445
+ torch.tensor([prompt], dtype=torch.int32),
446
+ dtype=ttnn.int32,
447
+ device=fixture.mesh,
448
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
449
+ )
450
+ token = full[:, :, prompt : prompt + 1, :]
451
+
452
+ base = fixture.dies(
453
+ O.decoder_layer_decode_optimized(
454
+ fixture.rep(token),
455
+ fixture.baseline,
456
+ cfg.global_config,
457
+ fixture.cos,
458
+ fixture.sin,
459
+ base_kv,
460
+ pos,
461
+ prompt,
462
+ packed_experts=fixture.baseline_experts,
463
+ )
464
+ )[0:1].float()
465
+ multi = fixture.dies(
466
+ MC.decoder_layer_decode_multichip(
467
+ fixture.rep(token), fixture.multichip, cfg, fixture.ctx, fixture.cos, fixture.sin, mc_kv, pos, prompt
468
+ )
469
+ )
470
+ spread = (multi - multi[0:1]).abs().max().item()
471
+ assert spread == 0.0, f"decode output differs across dies by {spread}"
472
+ kind = "contiguous" if block_size is None else f"paged({block_size})"
473
+ passing, message = comp_pcc(base, multi[0:1].float(), PCC_VS_SINGLE_CHIP)
474
+ logger.info(f"multichip decode [{kind}] vs single-chip TTNN: {message}")
475
+ assert passing, f"multichip decode [{kind}] below {PCC_VS_SINGLE_CHIP}: {message}"
476
+
477
+
478
+ @ring_fabric
479
+ @mesh_4
480
+ def test_multichip_decode_contiguous_batch8(fixture):
481
+ """The contiguous-cache SDPA workaround at batch > 1, against the single chip.
482
+
483
+ ``_sdpa_program_config`` is the layer's one hand-written program config and
484
+ the only place the multichip path departs from stage 02's tuning. It exists
485
+ because at TP=4 the contiguous cache has 1 KV head per die, so at batch 1
486
+ SDPA-decode asks for all 110 worker cores on that head and
487
+ ``sdpa_decode_program_factory.cpp:245`` refuses anything over 64. But
488
+ ``num_cores_per_head`` divides by the batch, so at batch 8 the op would have
489
+ asked for 13 and been legal *without* the config -- and the config is
490
+ supplied unconditionally on the contiguous path.
491
+
492
+ That makes batch > 1 the case that actually tests the workaround rather than
493
+ the failure it works around: here the cap is not rescuing anything, it is
494
+ only constraining, and so are the ``q_chunk_size``/``k_chunk_size`` of 32
495
+ that come with it (the default path picked its own). Every other contiguous
496
+ test is batch 1 and every batch > 1 test is paged, so without this one both
497
+ the cap and the chunk sizes are exercised in exactly one configuration.
498
+ """
499
+ cfg = fixture.config
500
+ batch, prompt = 8, 32
501
+ per_user = [_hidden(fixture.hf, prompt + 1, seed=u) for u in range(batch)]
502
+
503
+ base_kv = F.create_kv_cache(fixture.mesh, cfg.global_config.attention, batch, 128, block_size=None)
504
+ mc_kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, batch, 128, block_size=None)
505
+ assert not mc_kv.is_paged
506
+ for user, hidden in enumerate(per_user):
507
+ fixture.baseline_prefill(hidden[:, :, :prompt, :], kv_cache=base_kv, user_id=user)
508
+ fixture.multichip_prefill(hidden[:, :, :prompt, :], kv_cache=mc_kv, user_id=user)
509
+
510
+ tokens = torch.cat([h[:, :, prompt, :] for h in per_user], dim=1).reshape(1, 1, batch, fixture.hf.hidden_size)
511
+ pos = ttnn.from_torch(
512
+ torch.full((batch,), prompt, dtype=torch.int32),
513
+ dtype=ttnn.int32,
514
+ device=fixture.mesh,
515
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
516
+ )
517
+
518
+ base = (
519
+ fixture.dies(
520
+ O.decoder_layer_decode_optimized(
521
+ fixture.rep(tokens),
522
+ fixture.baseline,
523
+ cfg.global_config,
524
+ fixture.cos,
525
+ fixture.sin,
526
+ base_kv,
527
+ pos,
528
+ prompt,
529
+ packed_experts=fixture.baseline_experts,
530
+ )
531
+ )[0]
532
+ .reshape(-1, fixture.hf.hidden_size)[:batch]
533
+ .float()
534
+ )
535
+ multi_all = fixture.dies(
536
+ MC.decoder_layer_decode_multichip(
537
+ fixture.rep(tokens), fixture.multichip, cfg, fixture.ctx, fixture.cos, fixture.sin, mc_kv, pos, prompt
538
+ )
539
+ )
540
+ spread = (multi_all - multi_all[0:1]).abs().max().item()
541
+ assert spread == 0.0, f"decode output differs across dies by {spread}"
542
+ multi = multi_all[0].reshape(-1, fixture.hf.hidden_size)[:batch].float()
543
+
544
+ assert (multi - multi[0:1]).abs().max().item() > 1e-3, "all users identical (broadcast bug)"
545
+ for user in range(batch):
546
+ passing, message = comp_pcc(base[user : user + 1], multi[user : user + 1], PCC_VS_SINGLE_CHIP)
547
+ logger.info(f"multichip decode [contiguous, batch 8] user {user} vs single-chip TTNN: {message}")
548
+ assert passing, f"contiguous batch-8 user {user} below {PCC_VS_SINGLE_CHIP}: {message}"
549
+
550
+
551
+ @ring_fabric
552
+ @mesh_4
553
+ def test_multichip_multi_step_decode_vs_hf(fixture, reference):
554
+ """Four consecutive decode steps against HF, each at its own position."""
555
+ layer, hf_config = reference
556
+ cfg = fixture.config
557
+ prompt, steps = 32, 4
558
+ full = _hidden(hf_config, prompt + steps)
559
+ ref = _reference_layer(layer, hf_config, full)
560
+
561
+ kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, 1, MAX_SEQ, block_size=BLOCK_SIZE)
562
+ fixture.multichip_prefill(full[:, :, :prompt, :], kv_cache=kv)
563
+
564
+ for step in range(steps):
565
+ p = prompt + step
566
+ pos = ttnn.from_torch(
567
+ torch.tensor([p], dtype=torch.int32),
568
+ dtype=ttnn.int32,
569
+ device=fixture.mesh,
570
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
571
+ )
572
+ out = MC.decoder_layer_decode_multichip(
573
+ fixture.rep(full[:, :, p : p + 1, :]),
574
+ fixture.multichip,
575
+ cfg,
576
+ fixture.ctx,
577
+ fixture.cos,
578
+ fixture.sin,
579
+ kv,
580
+ pos,
581
+ p,
582
+ )
583
+ got = fixture.dies(out)[0].reshape(1, -1).float()
584
+ passing, message = comp_pcc(ref[:, p, :], got, 0.99)
585
+ logger.info(f"multichip decode step {step} (pos {p}) vs HF: {message}")
586
+ assert passing, f"multichip decode step {step} below 0.99: {message}"
587
+
588
+
589
+ @ring_fabric
590
+ @mesh_4
591
+ @pytest.mark.parametrize("batch", [1, 2, 8, 32], ids=["b1", "b2", "b8", "b32"])
592
+ def test_multichip_decode_batch(fixture, reference, batch):
593
+ """Multi-user decode, each user against its own HF reference.
594
+
595
+ 32 is the ceiling and it is a TTNN op limit that TP does not move:
596
+ ``nlp_create_qkv_heads_decode_device_operation.cpp:51`` asserts
597
+ ``num_users <= 32``, and that op is on the per-die path too. Per-user
598
+ references are what prove routing is per-user rather than broadcast, which
599
+ matters more under EP than on one die -- a die whose window is empty for one
600
+ user and full for another exercises the dynamic ``nnz`` path in both
601
+ directions inside a single program.
602
+ """
603
+ layer, hf_config = reference
604
+ cfg = fixture.config
605
+ prompt = 32
606
+ kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, batch, 128, block_size=BLOCK_SIZE)
607
+ per_user = [_hidden(hf_config, prompt + 1, seed=u) for u in range(batch)]
608
+ for user, hidden in enumerate(per_user):
609
+ fixture.multichip_prefill(hidden[:, :, :prompt, :], kv_cache=kv, user_id=user)
610
+
611
+ tokens = torch.cat([h[:, :, prompt, :] for h in per_user], dim=1).reshape(1, 1, batch, hf_config.hidden_size)
612
+ pos = ttnn.from_torch(
613
+ torch.full((batch,), prompt, dtype=torch.int32),
614
+ dtype=ttnn.int32,
615
+ device=fixture.mesh,
616
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
617
+ )
618
+ out = MC.decoder_layer_decode_multichip(
619
+ fixture.rep(tokens), fixture.multichip, cfg, fixture.ctx, fixture.cos, fixture.sin, kv, pos, prompt
620
+ )
621
+ got = fixture.dies(out)[0].reshape(-1, hf_config.hidden_size)[:batch].float()
622
+ assert torch.isfinite(got).all(), f"batch={batch} produced non-finite values"
623
+
624
+ for user, hidden in enumerate(per_user):
625
+ ref_user = _reference_layer(layer, hf_config, hidden)[:, prompt, :]
626
+ passing, message = comp_pcc(ref_user, got[user : user + 1], 0.99)
627
+ logger.info(f"multichip decode batch={batch} user {user} vs HF: {message}")
628
+ assert passing, f"batch={batch} user {user} below 0.99: {message}"
629
+
630
+ if batch > 1:
631
+ assert (got - got[0]).abs().max().item() > 1e-3, f"batch={batch}: all users identical (broadcast bug)"
632
+
633
+
634
+ # --- stage 04: the layer's own shape and layout contract ----------------------
635
+
636
+
637
+ @ring_fabric
638
+ @mesh_4
639
+ @pytest.mark.parametrize("batch", [1, 8], ids=["b1", "b8"])
640
+ def test_decode_output_layout_matches_input(fixture, batch):
641
+ """The decode layer must return exactly the tensor contract it takes.
642
+
643
+ Replicated ``[1, 1, B, 2048]``, bfloat16, TILE, DRAM-interleaved, logical
644
+ shape included -- that is what lets 48 layers stack with no boundary
645
+ conversion, and it is the *inter-layer residual layout contract* that
646
+ ``doc/optimized_multichip_decoder/README.md`` writes down for full-model
647
+ bringup.
648
+
649
+ It is asserted rather than assumed because stage 04's persistent collective
650
+ buffers can break it silently. A persistent output buffer imposes its own
651
+ logical shape on the op's result, and the layer's two all-reduces have the
652
+ same *padded* shape but different *logical* ones -- the attention partial is
653
+ 32 rows out of ``wo``, the expert partial is ``batch``. Keyed on the padded
654
+ shape alone they collide and the layer returns a 32-row tensor; every test
655
+ that compares a path against itself still passes.
656
+ """
657
+ cfg = fixture.config
658
+ prompt = 32
659
+ kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, batch, 128, block_size=BLOCK_SIZE)
660
+ hidden = _hidden(fixture.hf, prompt + 1)
661
+ for user in range(batch):
662
+ fixture.multichip_prefill(hidden[:, :, :prompt, :], kv_cache=kv, user_id=user)
663
+
664
+ token = fixture.rep(hidden[:, :, prompt, :].reshape(1, 1, 1, -1).repeat(1, 1, batch, 1))
665
+ pos = ttnn.from_torch(
666
+ torch.full((batch,), prompt, dtype=torch.int32),
667
+ dtype=ttnn.int32,
668
+ device=fixture.mesh,
669
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
670
+ )
671
+ out = MC.decoder_layer_decode_multichip(
672
+ token, fixture.multichip, cfg, fixture.ctx, fixture.cos, fixture.sin, kv, pos, prompt
673
+ )
674
+ assert list(out.shape) == list(token.shape), (
675
+ f"decode layer changed the logical shape: in {list(token.shape)}, out {list(out.shape)}. "
676
+ "48 of these stack, so the output contract must equal the input contract."
677
+ )
678
+ assert out.dtype == token.dtype, f"dtype changed: {token.dtype} -> {out.dtype}"
679
+ assert out.layout == token.layout, f"layout changed: {token.layout} -> {out.layout}"
680
+ assert out.memory_config() == token.memory_config(), (
681
+ f"memory config changed: {token.memory_config()} -> {out.memory_config()}; "
682
+ "there must be no inter-layer reshard"
683
+ )
684
+ spread = (fixture.dies(out) - fixture.dies(out)[0:1]).abs().max().item()
685
+ assert spread == 0.0, f"decode output differs across dies by {spread}"
686
+
687
+
688
+ # --- the dynamic-nnz hazard ---------------------------------------------------
689
+
690
+
691
+ @ring_fabric
692
+ @mesh_4
693
+ def test_expert_window_can_be_empty(fixture):
694
+ """A die holding none of the global top-8 must contribute an exact zero.
695
+
696
+ Constructed rather than hoped for: the router weight keeps amplified real
697
+ rows for experts 0..31 and zeroed rows for 32..127, so the routing collapses
698
+ onto the low end of the expert range and at least two dies end up holding
699
+ none of the global top-8. Measured, the split is **[6, 2, 0, 0]** -- die 0
700
+ full, die 1 partial, dies 2 and 3 empty -- which is a better test than a
701
+ clean [8,0,0,0] would have been, because it exercises both hazards at once:
702
+
703
+ * ``E_local = 0``, where a host-computed ``nnz`` would be 8 against zero live
704
+ sparsity entries and where an uninitialised ``sparse_matmul`` output would
705
+ put a NaN into the all-reduce and poison the layer on *every* die;
706
+ * ``0 < E_local < top_k``, the ordinary EP case, which is data-dependent and
707
+ is exactly why no single ``nnz`` can be computed on the host for a program
708
+ that runs on four dies at once.
709
+
710
+ So the assertions are on the *properties*, not on the exact split: the live
711
+ counts must sum to top-8 (the windows are a partition), at least one die must
712
+ be empty (the hazard is reached), every empty die must contribute an exact
713
+ zero, and the whole layer must still match the single-chip baseline.
714
+ """
715
+ cfg = fixture.config
716
+ hf = fixture.hf
717
+ n_local = cfg.local_moe.num_experts
718
+
719
+ # Amplified real router rows for experts 0..31, zeroed rows for 32..127. The
720
+ # zero rows give those experts a logit of exactly 0, so they only win a slot
721
+ # when fewer than 8 of the first 32 come out positive -- which is what
722
+ # produces the [6, 2, 0, 0] split rather than [8, 0, 0, 0].
723
+ #
724
+ # **The gain is 4x, and that number is a hazard, not a taste.** A first
725
+ # version of this test used 10x *synthetic* rows, which -- against the
726
+ # rms-normed activation, not the 0.02-scaled raw hidden -- produces logits
727
+ # with a standard deviation near 450 and a top-8 spread past 1000.
728
+ # ``exp(-1000)`` is exactly zero in bf16, so two of the eight routing weights
729
+ # underflowed, ``count_nonzero(sparsity)`` fell below the ``nnz = top_k *
730
+ # batch`` that the *single-chip* baseline passes, and the board deadlocked
731
+ # exactly as ``sparse_matmul_device_operation.cpp:205-211`` says it will --
732
+ # it had to be killed and reset. The multichip leg, on ``nnz=None``, was
733
+ # unaffected. That is this stage's own reproduction of the hazard the design
734
+ # phase only read about, and it is recorded in ``work_log.md``. At 4x on the
735
+ # real rows the top-8 spread is a few units and every weight stays normal.
736
+ forced = torch.zeros(cfg.global_config.moe.num_experts, hf.hidden_size)
737
+ forced[:n_local] = 4.0 * fixture.torch_router[:n_local]
738
+ forced_router = _replicate(
739
+ forced.T.contiguous().reshape(1, 1, hf.hidden_size, cfg.global_config.moe.num_experts), fixture.mesh
740
+ )
741
+
742
+ x = fixture.rep(_hidden(hf, 1))
743
+ routing = MC.router_forward_multichip(
744
+ x, forced_router, fixture.multichip.expert_window, cfg.global_config.moe, cfg.local_moe
745
+ )
746
+ live = [int((fixture.dies(routing)[d].reshape(-1) > 0).sum()) for d in range(cfg.num_devices)]
747
+ logger.info(f"forced routing: live experts per die = {live} (sum {sum(live)})")
748
+ assert sum(live) == cfg.global_config.moe.num_experts_per_tok, (
749
+ f"the four windows hold {sum(live)} experts, not the global top-"
750
+ f"{cfg.global_config.moe.num_experts_per_tok}: {live}"
751
+ )
752
+ empty = [d for d in range(cfg.num_devices) if live[d] == 0]
753
+ assert empty, f"the forced routing did not empty any die: {live}"
754
+
755
+ partial = fixture.dies(MC.moe_decode_multichip(x, routing, fixture.multichip.experts, cfg.local_moe)).float()
756
+ assert torch.isfinite(partial).all(), "an empty expert window produced non-finite output"
757
+ empty_max = partial[empty].abs().max().item()
758
+ logger.info(f"empty-window partials: max |value| on dies {empty} = {empty_max}")
759
+ assert empty_max == 0.0, f"a die with no live experts contributed {empty_max}, not an exact zero"
760
+
761
+ forced_baseline = F.DecoderLayerWeights(
762
+ input_layernorm=fixture.multichip.input_layernorm,
763
+ post_attention_layernorm=fixture.multichip.post_attention_layernorm,
764
+ attention=None,
765
+ router=forced_router,
766
+ experts=None,
767
+ )
768
+ forced_multichip = MC.MultichipWeights(
769
+ input_layernorm=fixture.multichip.input_layernorm,
770
+ post_attention_layernorm=fixture.multichip.post_attention_layernorm,
771
+ router=forced_router,
772
+ expert_window=fixture.multichip.expert_window,
773
+ experts=fixture.multichip.experts,
774
+ )
775
+ hidden = _hidden(hf, 128)
776
+ base = fixture.dies(
777
+ O.decoder_layer_prefill_optimized(
778
+ fixture.rep(hidden),
779
+ forced_baseline,
780
+ cfg.global_config,
781
+ fixture.cos,
782
+ fixture.sin,
783
+ fixture.baseline_sparsity,
784
+ fixture.baseline_experts,
785
+ )
786
+ )[0:1].float()
787
+ multi = fixture.dies(
788
+ MC.decoder_layer_prefill_multichip(
789
+ fixture.rep(hidden),
790
+ forced_multichip,
791
+ cfg,
792
+ fixture.ctx,
793
+ fixture.cos,
794
+ fixture.sin,
795
+ fixture.sparsity,
796
+ )
797
+ )[0:1].float()
798
+ passing, message = comp_pcc(base, multi, PCC_VS_SINGLE_CHIP)
799
+ logger.info(f"layer under maximally unbalanced routing vs single-chip: {message}")
800
+ assert passing, message
801
+
802
+
803
+ # --- stacking and trace -------------------------------------------------------
804
+
805
+
806
+ @ring_fabric
807
+ @mesh_4
808
+ def test_stacked_layer_io_contract(fixture):
809
+ """The layer's output must be usable as its own input, unmodified.
810
+
811
+ Stage 04 stacks 48 of these. The contract is a replicated
812
+ ``[1, 1, B, 2048]`` in DRAM in both directions, so feeding the output
813
+ straight back in must work with no gather, reshard, layout change or dtype
814
+ cast in between -- and must still match the single-chip baseline stacked the
815
+ same way, which is what rules out a per-layer boundary conversion hiding
816
+ inside the comparison.
817
+ """
818
+ hidden = _hidden(fixture.hf, 128)
819
+
820
+ base_out = fixture.baseline_prefill(hidden)
821
+ base_out2 = O.decoder_layer_prefill_optimized(
822
+ base_out,
823
+ fixture.baseline,
824
+ fixture.config.global_config,
825
+ fixture.cos,
826
+ fixture.sin,
827
+ fixture.baseline_sparsity,
828
+ fixture.baseline_experts,
829
+ )
830
+
831
+ multi_out = fixture.multichip_prefill(hidden)
832
+ assert multi_out.memory_config() == ttnn.DRAM_MEMORY_CONFIG
833
+ assert multi_out.layout == ttnn.TILE_LAYOUT and multi_out.dtype == ttnn.bfloat16
834
+ multi_out2 = MC.decoder_layer_prefill_multichip(
835
+ multi_out,
836
+ fixture.multichip,
837
+ fixture.config,
838
+ fixture.ctx,
839
+ fixture.cos,
840
+ fixture.sin,
841
+ fixture.sparsity,
842
+ )
843
+ assert multi_out2.shape == multi_out.shape
844
+ assert multi_out2.memory_config() == multi_out.memory_config()
845
+ assert multi_out2.dtype == multi_out.dtype and multi_out2.layout == multi_out.layout
846
+
847
+ passing, message = comp_pcc(
848
+ fixture.dies(base_out2)[0:1].float(), fixture.dies(multi_out2)[0:1].float(), PCC_VS_SINGLE_CHIP
849
+ )
850
+ logger.info(f"two stacked multichip layers vs two stacked single-chip layers: {message}")
851
+ assert passing, f"stacked layers diverge: {message}"
852
+
853
+
854
+ @ring_fabric
855
+ @mesh_4
856
+ def test_multichip_decode_is_traceable(fixture):
857
+ """Warmed trace capture and replay on the mesh, with a live input buffer.
858
+
859
+ Trace capture and CCL interact: global semaphores and any persistent CCL
860
+ buffer must exist before ``begin_trace_capture`` and nothing may allocate
861
+ inside it. ``MeshContext`` allocates its semaphores at construction and the
862
+ ``_ones_column`` constant is populated by the eager warm-up call below, which
863
+ is what makes the capture legal.
864
+ """
865
+ cfg = fixture.config
866
+ prompt = 32
867
+ full = _hidden(fixture.hf, prompt + 1)
868
+ kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, 1, MAX_SEQ, block_size=BLOCK_SIZE)
869
+ fixture.multichip_prefill(full[:, :, :prompt, :], kv_cache=kv)
870
+
871
+ tt_in = fixture.rep(full[:, :, prompt : prompt + 1, :])
872
+ pos = ttnn.from_torch(
873
+ torch.tensor([prompt], dtype=torch.int32),
874
+ dtype=ttnn.int32,
875
+ device=fixture.mesh,
876
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
877
+ )
878
+
879
+ def step():
880
+ return MC.decoder_layer_decode_multichip(
881
+ tt_in, fixture.multichip, cfg, fixture.ctx, fixture.cos, fixture.sin, kv, pos, prompt
882
+ )
883
+
884
+ eager = fixture.dies(step()).clone()
885
+ ttnn.synchronize_device(fixture.mesh)
886
+
887
+ trace_id = ttnn.begin_trace_capture(fixture.mesh, cq_id=0)
888
+ traced_out = step()
889
+ ttnn.end_trace_capture(fixture.mesh, trace_id, cq_id=0)
890
+
891
+ ttnn.execute_trace(fixture.mesh, trace_id, cq_id=0, blocking=True)
892
+ replayed = fixture.dies(traced_out).clone()
893
+ passing, message = comp_pcc(eager.float(), replayed.float(), 0.999)
894
+ logger.info(f"multichip traced decode vs eager: {message}")
895
+ assert passing, f"traced replay disagrees with eager: {message}"
896
+
897
+ other = torch.randn(1, 1, 1, fixture.hf.hidden_size) * 0.02
898
+ ttnn.copy_host_to_device_tensor(
899
+ ttnn.from_torch(
900
+ other,
901
+ dtype=ttnn.bfloat16,
902
+ layout=ttnn.TILE_LAYOUT,
903
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
904
+ ),
905
+ tt_in,
906
+ )
907
+ ttnn.execute_trace(fixture.mesh, trace_id, cq_id=0, blocking=True)
908
+ changed = fixture.dies(traced_out).clone()
909
+ delta = (replayed.float() - changed.float()).abs().max().item()
910
+ logger.info(f"multichip traced replay delta after input swap = {delta:.6f}")
911
+ assert delta > 1e-3, "the mesh trace is not reading the live input buffer"
912
+
913
+ ttnn.release_trace(fixture.mesh, trace_id)
914
+
915
+
916
+ @ring_fabric
917
+ @mesh_4
918
+ def test_multichip_decode_stress_is_deterministic(fixture):
919
+ """Repeated decode at the same position must be bit-identical, run to run.
920
+
921
+ The collectives are the new source of non-determinism here: two async CCLs
922
+ per layer with cycling semaphores, on a ring where four dies race to the same
923
+ reduction. Reduction order is fixed by the topology, so the result must be
924
+ exactly repeatable; a drift would mean the semaphore cycling is letting two
925
+ collectives overlap.
926
+ """
927
+ cfg = fixture.config
928
+ prompt = 32
929
+ full = _hidden(fixture.hf, prompt + 1)
930
+ kv = MC.create_mesh_kv_cache(fixture.mesh, cfg, 1, MAX_SEQ, block_size=BLOCK_SIZE)
931
+ fixture.multichip_prefill(full[:, :, :prompt, :], kv_cache=kv)
932
+ pos = ttnn.from_torch(
933
+ torch.tensor([prompt], dtype=torch.int32),
934
+ dtype=ttnn.int32,
935
+ device=fixture.mesh,
936
+ mesh_mapper=ttnn.ReplicateTensorToMesh(fixture.mesh),
937
+ )
938
+ token = fixture.rep(full[:, :, prompt : prompt + 1, :])
939
+
940
+ outs = []
941
+ for _ in range(20):
942
+ outs.append(
943
+ fixture.dies(
944
+ MC.decoder_layer_decode_multichip(
945
+ token, fixture.multichip, cfg, fixture.ctx, fixture.cos, fixture.sin, kv, pos, prompt
946
+ )
947
+ ).clone()
948
+ )
949
+ for i, o in enumerate(outs[1:], start=1):
950
+ assert torch.equal(outs[0], o), f"decode iteration {i} differs from iteration 0 (bitwise)"
951
+ logger.info("multichip decode: 20 consecutive steps bit-identical across all 4 dies")
952
+
953
+
954
+ # --- runtime fallback audit ---------------------------------------------------
955
+
956
+
957
+ @ring_fabric
958
+ @mesh_4
959
+ @pytest.mark.parametrize("batch", [1, 32], ids=["b1", "b32"])
960
+ def test_no_runtime_fallbacks(fixture, batch):
961
+ """None of the imported single-chip helpers may quietly take a slower path.
962
+
963
+ All three of them see different inputs under TP/EP than they were tuned
964
+ against, and all three fall back *silently* -- a PCC test cannot tell the
965
+ difference. In particular ``_dram_sharded_ok`` needs both weight dims
966
+ divisible by ``8 banks x 32 = 256``, and per-die wqkv N is 1280 = 5x256, one
967
+ factor of two from failing; if it ever did, stage 02's 1.11x DRAM-sharded
968
+ decode attention would disappear with no error at all.
969
+ """
970
+ audit = MC.fallback_audit(fixture.multichip, fixture.config, batch)
971
+ logger.info(f"multichip fallback audit at batch {batch}: {audit}")
972
+ assert audit["dram_sharded_taken"], "decode attention fell back to the interleaved path"
973
+ assert audit["dram_sharded_qkv"] == (2048, 1280), audit
974
+ assert audit["dram_sharded_wo"] == (1024, 2048), audit
975
+ # **Literals, not the module constants.** ``EXPERT_IN0_BLOCK_W_*`` are now
976
+ # derived from ``DEFAULT_PRECISION``, which is the same value the audit
977
+ # resolves from -- so comparing them was an identity that could not fail and
978
+ # could no longer catch a width regression. These are the widths stage 07
979
+ # selected and measured; changing the default must fail here and be
980
+ # re-measured, not silently ratified. ``test_precision_config.py`` pins the
981
+ # same two literals through the full construction path.
982
+ assert audit["gate_up_in0_block_w"] == 64, "gate/up block width moved off the stage-07 selection"
983
+ assert audit["down_in0_block_w"] == 24, "down block width moved off the stage-07 selection"
984
+ assert (O.EXPERT_IN0_BLOCK_W_GATE_UP, O.EXPERT_IN0_BLOCK_W_DOWN) == (
985
+ 64,
986
+ 24,
987
+ ), "the module constants no longer agree with the selected widths"
988
+ assert audit["local_heads"] == (8, 1) and audit["local_experts"] == 32, audit
989
+ # Batch 1 is the latency target and must keep the intermediates in L1 -- the
990
+ # traced A/B says L1 is 7.6% faster there. Batch 32 must not: the allocator
991
+ # refuses 234.88 MB outright (bank_manager.cpp:462), so a budget that let it
992
+ # through would be a crash, not a slow path. The inherited 40 MB constant
993
+ # would have separated these two correctly by luck while silently changing
994
+ # the answer for batches 2 to 16; the swept 128 MB is chosen against the
995
+ # measured L1-vs-DRAM crossover. See probes/l1_budget_probe.py.
996
+ assert audit["expert_intermediate_buffer"] == ("L1" if batch == 1 else "DRAM"), audit
997
+ # Stage 04. The sharded residual norm writes exactly the L1 shard the
998
+ # DRAM-sharded qkv projection reads, which is what lets the first norm's
999
+ # output cross into attention with no conversion at all. If that equality
1000
+ # ever breaks, TTNN inserts a reshard between them and the layer gets slower
1001
+ # with no error -- the same failure mode as the three above.
1002
+ assert audit["norm_shard_feeds_qkv_directly"], (
1003
+ "the sharded norm's output shard no longer matches attention's qkv input shard; "
1004
+ "a silent reshard has been reintroduced between them"
1005
+ )
1006
+ assert audit["norm_shard_cores"] == 8, audit
1007
+
1008
+
1009
+ def test_meta_rope_weights_match_hf():
1010
+ """The Meta channel permutation is a *pair*: Q/K rows and the QK-norm vectors.
1011
+
1012
+ Stage 01 chose HF-style RoPE precisely so neither permutation was needed.
1013
+ Stage 04 adopts ``rotary_embedding_llama`` on the decode path, which brings
1014
+ both back, and crossing them "runs fine and silently produces garbage"
1015
+ (``weight_mapping.py``). This asserts the whole convention on the host, with
1016
+ no device, so a mismatch fails here rather than as a PCC that is merely
1017
+ lower.
1018
+ """
1019
+ import torch
1020
+
1021
+ from ..tt.weight_mapping import hf_to_meta_channels, permute_head_vector_to_meta, permute_wqkv_to_meta
1022
+
1023
+ hd, nh, nkv, hidden = 128, 8, 1, 2048
1024
+ perm = hf_to_meta_channels(hd)
1025
+ inv = torch.argsort(perm)
1026
+
1027
+ # 1. The permutation is a permutation, and it is the interleave it claims.
1028
+ assert sorted(perm.tolist()) == list(range(hd))
1029
+ assert perm[0] == 0 and perm[1] == hd // 2 and perm[2] == 1
1030
+
1031
+ # 2. HF rope on HF-ordered channels == Meta rope on Meta-ordered channels.
1032
+ torch.manual_seed(0)
1033
+ x = torch.randn(4, hd)
1034
+ c, s = torch.randn(hd // 2).abs(), torch.randn(hd // 2)
1035
+ cos_hf, sin_hf = torch.cat([c, c]), torch.cat([s, s])
1036
+ hf = x * cos_hf + torch.cat([-x[:, hd // 2 :], x[:, : hd // 2]], dim=-1) * sin_hf
1037
+ xm = x[:, perm]
1038
+ cos_m, sin_m = cos_hf[perm], sin_hf[perm]
1039
+ rot = torch.stack([-xm[:, 1::2], xm[:, 0::2]], dim=-1).reshape(xm.shape)
1040
+ meta = xm * cos_m + rot * sin_m
1041
+ assert torch.allclose(meta[:, inv], hf, atol=1e-6), (meta[:, inv] - hf).abs().max()
1042
+
1043
+ # 3. permute_wqkv_to_meta touches Q and K and leaves V alone.
1044
+ wqkv = torch.randn(1, 1, hidden, (nh + 2 * nkv) * hd)
1045
+ out = permute_wqkv_to_meta(wqkv, n_heads=nh, n_kv_heads=nkv, head_dim=hd)
1046
+ assert out.shape == wqkv.shape
1047
+ v0 = (nh + nkv) * hd
1048
+ assert torch.equal(out[..., v0:], wqkv[..., v0:]), "V was permuted"
1049
+ for h in range(nh + nkv):
1050
+ lo = h * hd
1051
+ assert torch.equal(out[..., lo : lo + hd], wqkv[..., lo : lo + hd][..., perm])
1052
+ assert not torch.equal(out[..., :hd], wqkv[..., :hd]), "Q was not permuted"
1053
+
1054
+ # 4. Applying it twice is not identity -- i.e. forgetting it is detectable.
1055
+ vec = torch.randn(hd)
1056
+ assert not torch.equal(permute_head_vector_to_meta(vec, head_dim=hd), vec)
1057
+ assert torch.equal(permute_head_vector_to_meta(vec, head_dim=hd)[inv], vec)
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_optimized_decoder.py ADDED
@@ -0,0 +1,503 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Correctness of the optimized decoder — every stage-01 guarantee, re-checked.
5
+
6
+ Optimization here changed program configs and packed two weights together; it
7
+ did not change the maths. So these tests deliberately re-run the *same*
8
+ contracts stage 01 established rather than a reduced subset, because the way an
9
+ optimization usually breaks a model is by quietly narrowing what still works:
10
+ a program config legal only at tile-aligned lengths, a packed weight whose
11
+ halves are swapped, a trace that captured a stale buffer.
12
+
13
+ ``test_optimized_vs_functional_precision_delta`` is the sharpest of these.
14
+ Packing gate/up and widening ``in0_block_w`` are value-preserving, but the
15
+ optimized path also holds expert weights in bfloat4_b and attention projections
16
+ in bfloat8_b, so bit-identity is not available -- and asserting it would be
17
+ asserting the optimization away. Instead it bounds the gap between the two
18
+ implementations at 0.999 PCC. That is far tighter than either one's distance to
19
+ HF, so a real defect (swapped packed halves, a mis-sliced block) still cannot
20
+ hide inside it, while the quantisation that was measured and accepted can.
21
+ """
22
+
23
+ from __future__ import annotations
24
+
25
+ import pytest
26
+ import torch
27
+ from loguru import logger
28
+
29
+ import ttnn
30
+ from models.common.utility_functions import comp_allclose, comp_pcc
31
+
32
+ from ..tt import functional_decoder as F
33
+ from ..tt import optimized_decoder as O
34
+ from ..tt.weight_mapping import convert_layer_weights
35
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
36
+
37
+ LAYER_IDX = 0
38
+ PCC_REQUIRED = 0.995 # same bar as the functional decoder
39
+ MAX_SEQ = 1024
40
+ BLOCK_SIZE = 32
41
+ TRACE_REGION_SIZE = 50331648
42
+
43
+
44
+ @pytest.fixture(scope="module")
45
+ def reference():
46
+ return build_reference_layer(LAYER_IDX)
47
+
48
+
49
+ @pytest.fixture(scope="module")
50
+ def torch_weights(reference):
51
+ _, hf_config = reference
52
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
53
+
54
+
55
+ def _hidden(hf_config, seq_len, seed=0):
56
+ torch.manual_seed(seed)
57
+ return torch.randn(1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
58
+
59
+
60
+ def _reference_layer(layer, hf_config, hidden):
61
+ seq_len = hidden.shape[1]
62
+ cos, sin = rotary_embeddings(hf_config, seq_len)
63
+ mask = torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
64
+ with torch.no_grad():
65
+ out = layer(hidden, position_embeddings=(cos, sin), attention_mask=mask)
66
+ return out[0] if isinstance(out, tuple) else out
67
+
68
+
69
+ def _to_device(t, mesh_device):
70
+ return ttnn.from_torch(
71
+ t, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG
72
+ )
73
+
74
+
75
+ def _norm(t, mesh_device):
76
+ return ttnn.from_torch(
77
+ t.reshape(1, 1, 1, -1).float(),
78
+ dtype=ttnn.bfloat16,
79
+ layout=ttnn.TILE_LAYOUT,
80
+ device=mesh_device,
81
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
82
+ )
83
+
84
+
85
+ def _build(mesh_device, hf_config, torch_weights, *, functional_weights=False):
86
+ """Upload only what the path under test actually reads.
87
+
88
+ ``decoder_layer_prefill_optimized`` / ``decoder_layer_decode_optimized``
89
+ read exactly three tensors off ``DecoderLayerWeights`` -- the two norms and
90
+ the router -- and take every projection and expert weight from
91
+ ``OptimizedWeights``. Calling ``F.upload_layer_weights`` here as well
92
+ uploaded ~1.2 GB of bf16 experts plus a third copy of wqkv/wo that nothing
93
+ on the optimized path ever touches. That, not a model or op limit, is what
94
+ used to exhaust DRAM at batch 8, so the batch coverage below was capped by
95
+ the harness measuring its own waste.
96
+
97
+ ``functional_weights=True`` restores the full set, needed only by the test
98
+ that runs the *functional* layer side by side for a precision delta.
99
+ """
100
+ config = F.DecoderLayerConfig.from_hf(hf_config)
101
+ if functional_weights:
102
+ weights = F.upload_layer_weights(torch_weights, mesh_device, config)
103
+ else:
104
+ weights = F.DecoderLayerWeights(
105
+ input_layernorm=_norm(torch_weights["input_layernorm"], mesh_device),
106
+ post_attention_layernorm=_norm(torch_weights["post_attention_layernorm"], mesh_device),
107
+ attention=None, # optimized path uses OptimizedWeights.attention
108
+ router=F.upload_router_weight(torch_weights["router"], mesh_device),
109
+ experts=None, # optimized path uses OptimizedWeights.gate_up_proj/down_proj
110
+ )
111
+ packed = O.upload_packed_expert_weights(torch_weights, mesh_device, config.moe)
112
+ cos_cache, sin_cache = F.build_rope_cache(hf_config, MAX_SEQ, mesh_device)
113
+ sparsity = F.build_expert_sparsity(mesh_device, config.moe.num_experts)
114
+ return config, weights, packed, cos_cache, sin_cache, sparsity
115
+
116
+
117
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
118
+ @pytest.mark.parametrize("seq_len", [32, 128, 512, 33, 100, 257], ids=["s32", "s128", "s512", "s33", "s100", "s257"])
119
+ def test_optimized_prefill_vs_reference(mesh_device, reference, torch_weights, seq_len):
120
+ """Aligned and non-aligned lengths must both survive the new program configs."""
121
+ layer, hf_config = reference
122
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights)
123
+ hidden = _hidden(hf_config, seq_len)
124
+
125
+ ref_out = _reference_layer(layer, hf_config, hidden)
126
+ tt_out = ttnn.to_torch(
127
+ O.decoder_layer_prefill_optimized(
128
+ _to_device(hidden.unsqueeze(0), mesh_device), weights, config, cos, sin, sparsity, packed
129
+ )
130
+ ).squeeze(0)
131
+
132
+ passing, pcc_message = comp_pcc(ref_out, tt_out, PCC_REQUIRED)
133
+ logger.info(f"optimized prefill seq={seq_len}: {pcc_message}")
134
+ assert passing, f"optimized prefill (seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
135
+
136
+
137
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
138
+ @pytest.mark.parametrize("seq_len", [128, 33], ids=["s128", "s33"])
139
+ def test_optimized_vs_functional_precision_delta(mesh_device, reference, torch_weights, seq_len):
140
+ """Bound how far the optimized path may drift from the functional one.
141
+
142
+ The two are deliberately *not* bit-identical: the optimized path holds
143
+ expert weights in bfloat4_b and attention projections in bfloat8_b, which is
144
+ what makes it fast. So this asserts a tight bound on the gap rather than
145
+ equality -- close enough that a real defect (swapped packed halves, a
146
+ mis-sliced block) still cannot hide, but loose enough to admit the
147
+ quantisation that was measured and accepted.
148
+
149
+ The bound is empirical and is budgeted, not guessed. Against HF the
150
+ functional layer scores ~0.9995 and the optimized one ~0.9990, so the two
151
+ can differ by ~0.0013 purely from quantisation; measured, they differ by
152
+ 0.00102 at S=128 and 0.00133 at S=33. 0.998 sits just outside that and well
153
+ inside anything a structural bug would produce -- swapping the packed
154
+ gate/up halves, for instance, drops this to below 0.5.
155
+ """
156
+ _, hf_config = reference
157
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights, functional_weights=True)
158
+ hidden = _hidden(hf_config, seq_len)
159
+
160
+ base = ttnn.to_torch(
161
+ F.decoder_layer_prefill(_to_device(hidden.unsqueeze(0), mesh_device), weights, config, cos, sin, sparsity)
162
+ )
163
+ opt = ttnn.to_torch(
164
+ O.decoder_layer_prefill_optimized(
165
+ _to_device(hidden.unsqueeze(0), mesh_device), weights, config, cos, sin, sparsity, packed
166
+ )
167
+ )
168
+
169
+ passing, message = comp_pcc(base, opt, 0.998)
170
+ logger.info(comp_allclose(base, opt))
171
+ logger.info(f"seq={seq_len} optimized vs functional (bf16 vs {O.EXPERT_WEIGHT_DTYPE}): {message}")
172
+ assert passing, f"optimized diverges from functional beyond expert quantisation: {message}"
173
+
174
+
175
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
176
+ @pytest.mark.parametrize("seq_len", [1, 128, 33], ids=["decode", "s128", "s33"])
177
+ def test_optimized_router_matches_functional(mesh_device, reference, torch_weights, seq_len):
178
+ """The optimized router must route *identically*, not merely closely.
179
+
180
+ ``router_forward_optimized`` removes both keepdim reductions -- the max
181
+ becomes column 0 of the sorted top-k, and the sum moves after the scatter
182
+ and becomes a matmul. Neither is supposed to change the answer, so the test
183
+ asserts the strong property: the same 8 experts, in the same slots, with
184
+ weights equal to the functional router's within bf16 representation error.
185
+ A weaker PCC bound would let a genuinely different routing pass, which is
186
+ the failure mode that matters -- misrouting a token replaces its experts
187
+ outright rather than perturbing its output.
188
+ """
189
+ _, hf_config = reference
190
+ config = F.DecoderLayerConfig.from_hf(hf_config)
191
+ w_router = F.upload_router_weight(torch_weights["router"], mesh_device)
192
+ x = _to_device(_hidden(hf_config, seq_len).unsqueeze(0), mesh_device)
193
+
194
+ base = ttnn.to_torch(F.router_forward(x, w_router, config.moe)).float().reshape(seq_len, -1)
195
+ opt = ttnn.to_torch(O.router_forward_optimized(x, w_router, config.moe)).float().reshape(seq_len, -1)
196
+
197
+ assert torch.equal(base > 0, opt > 0), (
198
+ f"seq={seq_len}: optimized router selected different experts "
199
+ f"({(( base > 0) != (opt > 0)).sum().item()} slots differ)"
200
+ )
201
+ assert ((base > 0).sum(dim=-1) == config.moe.num_experts_per_tok).all()
202
+ delta = (base - opt).abs().max().item()
203
+ sums = opt.sum(dim=-1)
204
+ logger.info(
205
+ f"router seq={seq_len}: max |functional - optimized| = {delta:.3e}, weight sums in "
206
+ f"[{sums.min():.5f}, {sums.max():.5f}]"
207
+ )
208
+ assert delta < 5e-3, f"seq={seq_len}: routing weights differ by {delta}"
209
+ assert torch.allclose(sums, torch.ones_like(sums), atol=2e-2), f"seq={seq_len}: weights not normalised"
210
+
211
+
212
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
213
+ @pytest.mark.parametrize("seq_len", [1, 33, 100, 128], ids=["decode", "s33", "s100", "s128"])
214
+ def test_optimized_router_padding_is_zero(mesh_device, reference, torch_weights, seq_len):
215
+ """The routing tensor's tile row-padding must be exact zero, not +inf.
216
+
217
+ ``router_forward_optimized`` moved the sum after the scatter, so the divide
218
+ runs over whole tiles. Rows ``seq_len``..``ceil(seq_len/32)*32`` have a zero
219
+ numerator *and* a zero denominator, and unguarded ``ttnn.div`` returns
220
+ **+inf** there -- which the functional router, dividing before the scatter,
221
+ never did. No consumer was found that observes it (``to_torch`` returns the
222
+ logical shape, the sparsity path drops the padding when it converts to
223
+ ROW_MAJOR, and the scale multiply, ``rms_norm`` and ``fast_reduce_nc``
224
+ reduce along axes that are tile-aligned or not the padded one), so this is
225
+ a latent hazard rather than a live bug -- which is exactly the kind that a
226
+ later consumer turns into a silent NaN. The divisor is clamped in
227
+ ``router_forward_optimized``; this test is what stops the clamp being
228
+ optimized back out.
229
+
230
+ ``to_torch_with_padded_shape`` is the point of the test: ``ttnn.to_torch``
231
+ slices to the logical shape and would pass no matter what the padding held.
232
+ """
233
+ _, hf_config = reference
234
+ config = F.DecoderLayerConfig.from_hf(hf_config)
235
+ w_router = F.upload_router_weight(torch_weights["router"], mesh_device)
236
+ x = _to_device(_hidden(hf_config, seq_len).unsqueeze(0), mesh_device)
237
+
238
+ out = O.router_forward_optimized(x, w_router, config.moe)
239
+ padded = out.cpu().to_torch_with_padded_shape().float()
240
+ assert torch.isfinite(padded).all(), (
241
+ f"seq={seq_len}: routing tensor has {int((~torch.isfinite(padded)).sum())} non-finite "
242
+ f"entries in its padded shape {tuple(padded.shape)}"
243
+ )
244
+ pad = padded[..., seq_len:, :]
245
+ logger.info(
246
+ f"router seq={seq_len}: padded {tuple(padded.shape)}, {pad.numel()} padding entries, "
247
+ f"max |pad| = {pad.abs().max().item() if pad.numel() else 0.0}"
248
+ )
249
+ assert (pad == 0).all(), (
250
+ f"seq={seq_len}: {int((pad != 0).sum())} of {pad.numel()} padding entries are non-zero "
251
+ f"(max {pad.abs().max().item()})"
252
+ )
253
+
254
+
255
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
256
+ @pytest.mark.parametrize("block_size", [None, 32], ids=["contiguous", "paged32"])
257
+ def test_optimized_decode_matches_prefill(mesh_device, reference, torch_weights, block_size):
258
+ """Paged and contiguous KV caches both still work on the optimized path."""
259
+ layer, hf_config = reference
260
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights)
261
+ prompt_len = 32
262
+ hidden_full = _hidden(hf_config, prompt_len + 1)
263
+ ref_out = _reference_layer(layer, hf_config, hidden_full)[:, prompt_len, :]
264
+
265
+ kv_cache = F.create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=block_size)
266
+ O.decoder_layer_prefill_optimized(
267
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
268
+ weights,
269
+ config,
270
+ cos,
271
+ sin,
272
+ sparsity,
273
+ packed,
274
+ kv_cache=kv_cache,
275
+ )
276
+ current_pos = ttnn.from_torch(torch.tensor([prompt_len], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
277
+ out = O.decoder_layer_decode_optimized(
278
+ _to_device(hidden_full[:, prompt_len, :].reshape(1, 1, 1, hf_config.hidden_size), mesh_device),
279
+ weights,
280
+ config,
281
+ cos,
282
+ sin,
283
+ kv_cache,
284
+ current_pos,
285
+ prompt_len,
286
+ packed_experts=packed,
287
+ )
288
+ tt_out = ttnn.to_torch(out).reshape(1, hf_config.hidden_size)
289
+
290
+ passing, pcc_message = comp_pcc(ref_out, tt_out, 0.99)
291
+ kind = "contiguous" if block_size is None else f"paged({block_size})"
292
+ logger.info(f"optimized decode [{kind}]: {pcc_message}")
293
+ assert passing, f"optimized decode [{kind}] below 0.99: {pcc_message}"
294
+
295
+
296
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
297
+ def test_optimized_multi_step_decode(mesh_device, reference, torch_weights):
298
+ """Several steps against a paged cache, each checked at its own position."""
299
+ layer, hf_config = reference
300
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights)
301
+ prompt_len, steps = 32, 4
302
+ hidden_full = _hidden(hf_config, prompt_len + steps)
303
+ ref_out = _reference_layer(layer, hf_config, hidden_full)
304
+
305
+ kv_cache = F.create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE)
306
+ O.decoder_layer_prefill_optimized(
307
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
308
+ weights,
309
+ config,
310
+ cos,
311
+ sin,
312
+ sparsity,
313
+ packed,
314
+ kv_cache=kv_cache,
315
+ )
316
+
317
+ for step in range(steps):
318
+ pos = prompt_len + step
319
+ current_pos = ttnn.from_torch(torch.tensor([pos], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
320
+ out = O.decoder_layer_decode_optimized(
321
+ _to_device(hidden_full[:, pos, :].reshape(1, 1, 1, hf_config.hidden_size), mesh_device),
322
+ weights,
323
+ config,
324
+ cos,
325
+ sin,
326
+ kv_cache,
327
+ current_pos,
328
+ pos,
329
+ packed_experts=packed,
330
+ )
331
+ passing, pcc_message = comp_pcc(ref_out[:, pos, :], ttnn.to_torch(out).reshape(1, -1), 0.99)
332
+ logger.info(f"optimized decode step {step} (pos {pos}): {pcc_message}")
333
+ assert passing, f"optimized decode step {step} below 0.99: {pcc_message}"
334
+
335
+
336
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
337
+ def test_optimized_prefill_is_deterministic(mesh_device, reference, torch_weights):
338
+ """Bitwise repeatability, as required of the functional path."""
339
+ _, hf_config = reference
340
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights)
341
+ hidden = _hidden(hf_config, 128)
342
+
343
+ outs = [
344
+ ttnn.to_torch(
345
+ O.decoder_layer_prefill_optimized(
346
+ _to_device(hidden.unsqueeze(0), mesh_device), weights, config, cos, sin, sparsity, packed
347
+ )
348
+ ).clone()
349
+ for _ in range(3)
350
+ ]
351
+ assert torch.equal(outs[0], outs[1]), "optimized prefill run 1 != run 2 (bitwise)"
352
+ assert torch.equal(outs[0], outs[2]), "optimized prefill run 1 != run 3 (bitwise)"
353
+ logger.info("optimized prefill: 3 runs bit-identical")
354
+
355
+
356
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
357
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
358
+ def test_optimized_decode_is_traceable(mesh_device, reference, torch_weights):
359
+ """Trace capture, bit-exact replay, and a live input buffer."""
360
+ _, hf_config = reference
361
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights)
362
+ prompt_len = 32
363
+ hidden_full = _hidden(hf_config, prompt_len + 1)
364
+
365
+ kv_cache = F.create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE)
366
+ O.decoder_layer_prefill_optimized(
367
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
368
+ weights,
369
+ config,
370
+ cos,
371
+ sin,
372
+ sparsity,
373
+ packed,
374
+ kv_cache=kv_cache,
375
+ )
376
+
377
+ tt_in = _to_device(hidden_full[:, prompt_len, :].reshape(1, 1, 1, hf_config.hidden_size), mesh_device)
378
+ current_pos = ttnn.from_torch(torch.tensor([prompt_len], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
379
+
380
+ def step():
381
+ return O.decoder_layer_decode_optimized(
382
+ tt_in, weights, config, cos, sin, kv_cache, current_pos, prompt_len, packed_experts=packed
383
+ )
384
+
385
+ eager = ttnn.to_torch(step()).clone()
386
+ trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0)
387
+ traced_out = step()
388
+ ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0)
389
+
390
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
391
+ replayed = ttnn.to_torch(traced_out).clone()
392
+ passing, message = comp_pcc(eager, replayed, 0.999)
393
+ logger.info(f"optimized traced vs eager: {message}")
394
+ assert passing, f"optimized traced replay disagrees with eager: {message}"
395
+
396
+ # The trace must read the live buffer, not a value captured at record time.
397
+ other = (_hidden(hf_config, 1, seed=99)).reshape(1, 1, 1, hf_config.hidden_size)
398
+ ttnn.copy_host_to_device_tensor(ttnn.from_torch(other, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT), tt_in)
399
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
400
+ changed = ttnn.to_torch(traced_out).clone()
401
+ delta = (replayed.float() - changed.float()).abs().max().item()
402
+ logger.info(f"optimized traced replay delta after input swap = {delta:.6f}")
403
+ assert delta > 1e-3, "optimized trace is not reading the live input buffer"
404
+
405
+ ttnn.release_trace(mesh_device, trace_id)
406
+
407
+
408
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
409
+ @pytest.mark.parametrize("batch", [1, 2, 8, 32], ids=["b1", "b2", "b8", "b32"])
410
+ def test_optimized_decode_batch(mesh_device, reference, torch_weights, batch):
411
+ """Multi-user decode. Batch 1 is the latency target, but capability must hold.
412
+
413
+ This did not work in stage 01 at any batch above 1: ``sparse_matmul``
414
+ resolves ``batch_length`` differently depending on which operand is flagged
415
+ sparse, and the down projection landed on a branch that ignores the batch
416
+ dimension entirely. Each user is given a different prompt so a broadcast bug
417
+ -- every row returning user 0's answer -- cannot pass.
418
+
419
+ Writing this test surfaced a second defect beyond the sparsity one:
420
+ ``attention_prefill`` hardcoded ``user_id=0``, so every user's prompt
421
+ overwrote slot 0 and the other slots stayed empty. Per-user PCC failed at
422
+ 0.92-0.93 until ``user_id`` was threaded through prefill. Both halves were
423
+ needed for multi-user decode to work end to end.
424
+
425
+ Coverage used to stop at 2, blamed on a harness limit: every
426
+ parametrisation re-uploaded **1.63 GB** of weights (the bf16 functional set
427
+ *and* the optimized set) without reclaiming the previous ones, so b8
428
+ exhausted DRAM in ``bank_manager.cpp`` before the layer ran. That was
429
+ self-inflicted -- the optimized path never reads the functional experts or
430
+ the functional attention copy. ``_build`` now uploads only what this path
431
+ touches, **0.38 GB**, and b8 and b32 run.
432
+
433
+ Both figures are derived from the shipped uploads rather than recalled:
434
+ functional experts 3 x 128 x 768 x 2048 elem at bf16 = 1.208 GB, plus
435
+ wqkv+wo bf16 = 37.75 MB; optimized experts (128 x 1536 x 2048 packed gate/up
436
+ plus 128 x 768 x 2048 down) at bfloat4_b 0.5625 B/elem = 339.7 MB, plus two
437
+ copies of wqkv+wo at bfloat8_b 1.0625 B/elem = 40.11 MB. Earlier revisions
438
+ of this docstring said ~2.4 GB and ~0.63 GB, and neither was derived.
439
+
440
+ 32 is the real ceiling, and it is a **TTNN op limit**, not this layer's
441
+ shape choice: ``nlp_create_qkv_heads_decode_device_operation.cpp:51``
442
+ asserts ``num_users <= num_users_supported`` with ``num_users_supported =
443
+ 32`` hardcoded at line 45 of that file, and that op is on the interleaved
444
+ attention path as well as the DRAM-sharded one. ``_dram_sharded_usable``
445
+ does refuse the sharded projections past B=32 -- ``_width_sharded_l1``
446
+ hardcodes a 32-row shard and ``_dram_sharded_program_config`` sets
447
+ ``per_core_M=1`` -- but the interleaved fallback it selects then fails in
448
+ ``nlp_create_qkv_heads_decode`` too. The guard buys a comprehensible
449
+ failure, not a working larger batch.
450
+ """
451
+ layer, hf_config = reference
452
+ config, weights, packed, cos, sin, sparsity = _build(mesh_device, hf_config, torch_weights)
453
+ prompt_len = 32
454
+
455
+ # Small cache on purpose: this test is about multi-user expert routing, and
456
+ # a full-length cache per user would exhaust DRAM before reaching the point.
457
+ kv_cache = F.create_kv_cache(mesh_device, config.attention, max_batch=batch, max_seq_len=128, block_size=BLOCK_SIZE)
458
+ per_user = [_hidden(hf_config, prompt_len + 1, seed=u) for u in range(batch)]
459
+ # Each user's prompt must actually land in the cache, or decode attends
460
+ # zeros and the test proves nothing about multi-user routing.
461
+ for user, hidden_full in enumerate(per_user):
462
+ O.decoder_layer_prefill_optimized(
463
+ _to_device(hidden_full[:, :prompt_len, :].unsqueeze(0), mesh_device),
464
+ weights,
465
+ config,
466
+ cos,
467
+ sin,
468
+ sparsity,
469
+ packed,
470
+ kv_cache=kv_cache,
471
+ user_id=user,
472
+ )
473
+
474
+ tokens = torch.cat([h[:, prompt_len, :] for h in per_user], dim=0) # [batch, hidden]
475
+ current_pos = ttnn.from_torch(
476
+ torch.full((batch,), prompt_len, dtype=torch.int32), dtype=ttnn.int32, device=mesh_device
477
+ )
478
+ out = O.decoder_layer_decode_optimized(
479
+ _to_device(tokens.reshape(1, 1, batch, hf_config.hidden_size), mesh_device),
480
+ weights,
481
+ config,
482
+ cos,
483
+ sin,
484
+ kv_cache,
485
+ current_pos,
486
+ prompt_len,
487
+ packed_experts=packed,
488
+ )
489
+ tt_out = ttnn.to_torch(out).reshape(-1, hf_config.hidden_size)[:batch].float()
490
+
491
+ assert torch.isfinite(tt_out).all(), f"batch={batch} produced non-finite values"
492
+
493
+ # Every user is checked against its own HF reference. This is what proves
494
+ # the routing is per-user rather than broadcast.
495
+ for user, hidden_full in enumerate(per_user):
496
+ ref_user = _reference_layer(layer, hf_config, hidden_full)[:, prompt_len, :]
497
+ passing, message = comp_pcc(ref_user, tt_out[user : user + 1], 0.99)
498
+ logger.info(f"optimized decode batch={batch} user {user} vs HF: {message}")
499
+ assert passing, f"batch={batch} user {user} decode below 0.99: {message}"
500
+
501
+ if batch > 1:
502
+ spread = (tt_out - tt_out[0]).abs().max().item()
503
+ assert spread > 1e-3, f"batch={batch}: all users returned identical output (broadcast bug)"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_perf.py ADDED
@@ -0,0 +1,611 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Warmed prefill and traced warmed decode latency for the functional decoder.
5
+
6
+ These numbers are the stage-01 baseline that the optimized decoder has to beat,
7
+ so the measurement conditions matter as much as the values:
8
+
9
+ * **Warmed.** The first call to any shape compiles kernels and populates the
10
+ program cache. Timing that measures the compiler. Every configuration runs
11
+ warmup iterations that are discarded.
12
+ * **Traced decode.** Decode is short enough that host dispatch overhead is a
13
+ large share of wall time, so an eager measurement mostly reports Python.
14
+ Replaying a captured trace is what a serving stack actually does.
15
+ * **Device-synchronised.** Dispatch is asynchronous; without an explicit
16
+ synchronise the host would time enqueue calls rather than execution.
17
+ * **Median, not mean.** One descheduled iteration should not move the number.
18
+
19
+ Results are written to ``doc/functional_decoder/`` as CSV so later stages can
20
+ diff against them rather than re-deriving a baseline.
21
+
22
+ Two harness rules the CSVs depend on:
23
+
24
+ * **Every test here is marked ``models_performance_bare_metal``.** These tests
25
+ overwrite the published CSVs, and ``TT_METAL_WATCHER=10`` inflates device
26
+ timings roughly 8x, so a watcher run over the whole suite silently replaced
27
+ the prefill baseline with 4358 us/token. Watcher runs must deselect them:
28
+ ``TT_METAL_WATCHER=10 pytest ... -m "not models_performance_bare_metal"``.
29
+ * **The decode CSVs are rewritten whole, not appended to.** The decode test is
30
+ parametrised over context length, so each parametrisation contributes one
31
+ row; rows accumulate in ``_DECODE_ROWS`` for the life of the process and the
32
+ file is rewritten from scratch each time. Appending stacked eight interleaved
33
+ runs into ``doc/functional_decoder/perf_decode.csv`` before this was fixed.
34
+ """
35
+
36
+ from __future__ import annotations
37
+
38
+ import csv
39
+ import statistics
40
+ import time
41
+ from pathlib import Path
42
+
43
+ import pytest
44
+ import torch
45
+ from loguru import logger
46
+
47
+ import ttnn
48
+
49
+ from ..tt.functional_decoder import (
50
+ DecoderLayerConfig,
51
+ build_expert_sparsity,
52
+ build_rope_cache,
53
+ create_kv_cache,
54
+ decoder_layer_decode,
55
+ decoder_layer_prefill,
56
+ upload_layer_weights,
57
+ )
58
+ from ..tt.weight_mapping import convert_layer_weights
59
+ from .reference import build_reference_layer, layer_state_dict
60
+
61
+ LAYER_IDX = 0
62
+ MAX_SEQ = 4096
63
+ BLOCK_SIZE = 32
64
+ TRACE_REGION_SIZE = 50331648
65
+
66
+ PREFILL_LENGTHS = [128, 512, 1024, 2048]
67
+ PREFILL_WARMUP, PREFILL_ITERS = 1, 5
68
+ DECODE_WARMUP, DECODE_ITERS = 10, 100
69
+
70
+ DOC_DIR = Path(__file__).resolve().parents[1] / "doc" / "functional_decoder"
71
+
72
+ DECODE_FIELDS = ["context_len", "median_ms", "min_ms", "max_ms", "tok_per_s_per_layer", "iters"]
73
+
74
+ # {csv path: [row, ...]} for this process only, so a rerun truncates.
75
+ _DECODE_ROWS: dict[Path, list[dict]] = {}
76
+
77
+
78
+ def _write_decode_row(path: Path, row: dict) -> None:
79
+ """Add a row and rewrite the whole file (see module docstring)."""
80
+ path.parent.mkdir(parents=True, exist_ok=True)
81
+ rows = _DECODE_ROWS.setdefault(path, [])
82
+ rows[:] = [r for r in rows if r["context_len"] != row["context_len"]] + [row]
83
+ rows.sort(key=lambda r: r["context_len"])
84
+ with path.open("w", newline="") as fh:
85
+ writer = csv.DictWriter(fh, fieldnames=DECODE_FIELDS)
86
+ writer.writeheader()
87
+ writer.writerows(rows)
88
+
89
+
90
+ @pytest.fixture(scope="module")
91
+ def reference():
92
+ return build_reference_layer(LAYER_IDX)
93
+
94
+
95
+ @pytest.fixture(scope="module")
96
+ def torch_weights(reference):
97
+ _, hf_config = reference
98
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
99
+
100
+
101
+ def _to_device(t, mesh_device):
102
+ return ttnn.from_torch(
103
+ t, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG
104
+ )
105
+
106
+
107
+ def _write_csv(name: str, fieldnames: list[str], rows: list[dict]) -> Path:
108
+ DOC_DIR.mkdir(parents=True, exist_ok=True)
109
+ path = DOC_DIR / name
110
+ with path.open("w", newline="") as fh:
111
+ writer = csv.DictWriter(fh, fieldnames=fieldnames)
112
+ writer.writeheader()
113
+ writer.writerows(rows)
114
+ logger.info(f"wrote {path}")
115
+ return path
116
+
117
+
118
+ @pytest.mark.models_performance_bare_metal
119
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
120
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
121
+ def test_prefill_latency(mesh_device, reference, torch_weights):
122
+ _, hf_config = reference
123
+ config = DecoderLayerConfig.from_hf(hf_config)
124
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
125
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
126
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
127
+
128
+ rows = []
129
+ for seq_len in PREFILL_LENGTHS:
130
+ torch.manual_seed(0)
131
+ hidden = torch.randn(1, 1, seq_len, hf_config.hidden_size) * 0.02
132
+ tt_in = _to_device(hidden, mesh_device)
133
+
134
+ def once():
135
+ out = decoder_layer_prefill(tt_in, weights, config, cos_cache, sin_cache, sparsity)
136
+ ttnn.synchronize_device(mesh_device)
137
+ ttnn.deallocate(out)
138
+
139
+ for _ in range(PREFILL_WARMUP):
140
+ once()
141
+
142
+ samples = []
143
+ for _ in range(PREFILL_ITERS):
144
+ t0 = time.perf_counter()
145
+ once()
146
+ samples.append((time.perf_counter() - t0) * 1e3)
147
+
148
+ median = statistics.median(samples)
149
+ per_tok = median / seq_len * 1e3 # us/token
150
+ logger.info(
151
+ f"prefill S={seq_len:>5}: median {median:8.2f} ms "
152
+ f"min {min(samples):8.2f} max {max(samples):8.2f} ({per_tok:6.1f} us/token)"
153
+ )
154
+ rows.append(
155
+ {
156
+ "seq_len": seq_len,
157
+ "median_ms": round(median, 3),
158
+ "min_ms": round(min(samples), 3),
159
+ "max_ms": round(max(samples), 3),
160
+ "us_per_token": round(per_tok, 2),
161
+ "iters": PREFILL_ITERS,
162
+ }
163
+ )
164
+ ttnn.deallocate(tt_in)
165
+
166
+ _write_csv(
167
+ "perf_prefill.csv",
168
+ ["seq_len", "median_ms", "min_ms", "max_ms", "us_per_token", "iters"],
169
+ rows,
170
+ )
171
+ assert all(r["median_ms"] > 0 for r in rows)
172
+
173
+
174
+ @pytest.mark.models_performance_bare_metal
175
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
176
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
177
+ @pytest.mark.parametrize("context_len", [128, 1024, 4096], ids=["ctx128", "ctx1k", "ctx4k"])
178
+ def test_decode_latency_traced(mesh_device, reference, torch_weights, context_len):
179
+ """Traced single-token decode latency at several cache depths.
180
+
181
+ Swept over context because decode cost is dominated by the SDPA read over
182
+ the cache, so a single depth would not say whether latency is flat or grows
183
+ with the conversation.
184
+ """
185
+ _, hf_config = reference
186
+ config = DecoderLayerConfig.from_hf(hf_config)
187
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
188
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
189
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE)
190
+
191
+ pos = context_len - 1
192
+ torch.manual_seed(0)
193
+ tt_in = _to_device(torch.randn(1, 1, 1, hf_config.hidden_size) * 0.02, mesh_device)
194
+ current_pos = ttnn.from_torch(torch.tensor([pos], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
195
+
196
+ def step():
197
+ return decoder_layer_decode(
198
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, token_index=pos
199
+ )
200
+
201
+ step() # compile outside the capture
202
+ ttnn.synchronize_device(mesh_device)
203
+
204
+ trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0)
205
+ step()
206
+ ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0)
207
+
208
+ for _ in range(DECODE_WARMUP):
209
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
210
+
211
+ samples = []
212
+ for _ in range(DECODE_ITERS):
213
+ t0 = time.perf_counter()
214
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
215
+ samples.append((time.perf_counter() - t0) * 1e3)
216
+
217
+ median = statistics.median(samples)
218
+ logger.info(
219
+ f"traced decode ctx={context_len:>5}: median {median:7.3f} ms "
220
+ f"min {min(samples):7.3f} max {max(samples):7.3f} "
221
+ f"({1e3 / median:7.1f} tok/s/layer)"
222
+ )
223
+
224
+ _write_decode_row(
225
+ DOC_DIR / "perf_decode.csv",
226
+ {
227
+ "context_len": context_len,
228
+ "median_ms": round(median, 4),
229
+ "min_ms": round(min(samples), 4),
230
+ "max_ms": round(max(samples), 4),
231
+ "tok_per_s_per_layer": round(1e3 / median, 1),
232
+ "iters": DECODE_ITERS,
233
+ },
234
+ )
235
+
236
+ ttnn.release_trace(mesh_device, trace_id)
237
+ assert median > 0
238
+
239
+
240
+ # --- optimized path -----------------------------------------------------------
241
+ # Same harness, same conditions, writing to doc/optimized_decoder/ so the two
242
+ # stages' CSVs are directly diffable rather than needing re-derivation.
243
+
244
+ OPT_DOC_DIR = Path(__file__).resolve().parents[1] / "doc" / "optimized_decoder"
245
+
246
+
247
+ @pytest.mark.models_performance_bare_metal
248
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
249
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
250
+ def test_optimized_prefill_latency(mesh_device, reference, torch_weights):
251
+ from ..tt import optimized_decoder as O
252
+
253
+ _, hf_config = reference
254
+ config = DecoderLayerConfig.from_hf(hf_config)
255
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
256
+ packed = O.upload_packed_expert_weights(torch_weights, mesh_device, config.moe)
257
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
258
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
259
+
260
+ rows = []
261
+ for seq_len in PREFILL_LENGTHS:
262
+ torch.manual_seed(0)
263
+ tt_in = _to_device(torch.randn(1, 1, seq_len, hf_config.hidden_size) * 0.02, mesh_device)
264
+
265
+ def once():
266
+ out = O.decoder_layer_prefill_optimized(tt_in, weights, config, cos_cache, sin_cache, sparsity, packed)
267
+ ttnn.synchronize_device(mesh_device)
268
+ ttnn.deallocate(out)
269
+
270
+ for _ in range(PREFILL_WARMUP):
271
+ once()
272
+ samples = []
273
+ for _ in range(PREFILL_ITERS):
274
+ t0 = time.perf_counter()
275
+ once()
276
+ samples.append((time.perf_counter() - t0) * 1e3)
277
+
278
+ median = statistics.median(samples)
279
+ logger.info(
280
+ f"OPTIMIZED prefill S={seq_len:>5}: median {median:8.2f} ms ({median / seq_len * 1e3:6.1f} us/token)"
281
+ )
282
+ rows.append(
283
+ {
284
+ "seq_len": seq_len,
285
+ "median_ms": round(median, 3),
286
+ "min_ms": round(min(samples), 3),
287
+ "max_ms": round(max(samples), 3),
288
+ "us_per_token": round(median / seq_len * 1e3, 2),
289
+ "iters": PREFILL_ITERS,
290
+ }
291
+ )
292
+ ttnn.deallocate(tt_in)
293
+
294
+ OPT_DOC_DIR.mkdir(parents=True, exist_ok=True)
295
+ with (OPT_DOC_DIR / "perf_prefill.csv").open("w", newline="") as fh:
296
+ wr = csv.DictWriter(fh, fieldnames=["seq_len", "median_ms", "min_ms", "max_ms", "us_per_token", "iters"])
297
+ wr.writeheader()
298
+ wr.writerows(rows)
299
+ assert all(r["median_ms"] > 0 for r in rows)
300
+
301
+
302
+ @pytest.mark.models_performance_bare_metal
303
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
304
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
305
+ @pytest.mark.parametrize("context_len", [128, 1024, 4096], ids=["ctx128", "ctx1k", "ctx4k"])
306
+ def test_optimized_decode_latency_traced(mesh_device, reference, torch_weights, context_len):
307
+ from ..tt import optimized_decoder as O
308
+
309
+ _, hf_config = reference
310
+ config = DecoderLayerConfig.from_hf(hf_config)
311
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
312
+ packed = O.upload_packed_expert_weights(torch_weights, mesh_device, config.moe)
313
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
314
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE)
315
+
316
+ pos = context_len - 1
317
+ torch.manual_seed(0)
318
+ tt_in = _to_device(torch.randn(1, 1, 1, hf_config.hidden_size) * 0.02, mesh_device)
319
+ current_pos = ttnn.from_torch(torch.tensor([pos], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
320
+
321
+ def step():
322
+ return O.decoder_layer_decode_optimized(
323
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, pos, packed_experts=packed
324
+ )
325
+
326
+ step()
327
+ ttnn.synchronize_device(mesh_device)
328
+ trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0)
329
+ step()
330
+ ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0)
331
+ for _ in range(DECODE_WARMUP):
332
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
333
+ samples = []
334
+ for _ in range(DECODE_ITERS):
335
+ t0 = time.perf_counter()
336
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
337
+ samples.append((time.perf_counter() - t0) * 1e3)
338
+
339
+ median = statistics.median(samples)
340
+ logger.info(
341
+ f"OPTIMIZED traced decode ctx={context_len:>5}: median {median:7.3f} ms ({1e3 / median:7.1f} tok/s/layer)"
342
+ )
343
+
344
+ _write_decode_row(
345
+ OPT_DOC_DIR / "perf_decode.csv",
346
+ {
347
+ "context_len": context_len,
348
+ "median_ms": round(median, 4),
349
+ "min_ms": round(min(samples), 4),
350
+ "max_ms": round(max(samples), 4),
351
+ "tok_per_s_per_layer": round(1e3 / median, 1),
352
+ "iters": DECODE_ITERS,
353
+ },
354
+ )
355
+ ttnn.release_trace(mesh_device, trace_id)
356
+ assert median > 0
357
+
358
+
359
+ # --- multichip path -----------------------------------------------------------
360
+ # Stage 03, on the full 4-die P300_X2 mesh. Three tests, and the first of them is
361
+ # the one that makes the other two mean anything:
362
+ #
363
+ # test_multichip_baseline_1x1_* re-measures the *single-chip* optimized layer,
364
+ # with this harness, in this tree, on one die, and writes it into
365
+ # doc/multichip_decoder/. Stage 02's CSVs already hold a warmed single-chip
366
+ # baseline, but they are also the artifact its own README quotes cell by
367
+ # cell, and re-running them here would move their third significant figure
368
+ # and silently invalidate that document. A stage-owned copy of the baseline
369
+ # is cheaper than a cross-stage prose/artifact mismatch.
370
+ # test_multichip_prefill_latency / test_multichip_decode_latency_traced
371
+ # measure the same lengths on the mesh, so speedup is one CSV cell divided
372
+ # by another rather than a number quoted from anywhere.
373
+ #
374
+ # All three are marked models_performance_bare_metal: they rewrite published
375
+ # CSVs and TT_METAL_WATCHER inflates device timings, so a watcher run must
376
+ # deselect them.
377
+
378
+ MC_DOC_DIR = Path(__file__).resolve().parents[1] / "doc" / "optimized_multichip_decoder"
379
+ # Stage 04 note. ``tt/multichip_decoder.py`` is optimized **in place**, so the
380
+ # four tests below now measure the stage-04 path. They therefore write into
381
+ # ``doc/optimized_multichip_decoder/``; ``doc/multichip_decoder/perf_*.csv`` are
382
+ # stage 03's frozen *before* numbers and are deliberately never regenerated --
383
+ # re-pointing this constant back would overwrite the baseline half of every
384
+ # before/after table in both READMEs. The stage-04 decode before/after is also
385
+ # measured in one process by
386
+ # ``doc/optimized_multichip_decoder/probes/layer_levers.py``, whose "stage 03"
387
+ # leg is a verbatim copy of the committed stage-03 layer body.
388
+
389
+ # Ring fabric must be set before the mesh opens; the conftest device_params hook
390
+ # does that, which is why it is spelled here rather than with set_fabric_config.
391
+ MC_DEVICE_PARAMS = {"trace_region_size": TRACE_REGION_SIZE, "fabric_config": ttnn.FabricConfig.FABRIC_1D_RING}
392
+ MC_MESH = (1, 4)
393
+
394
+ PREFILL_FIELDS = ["seq_len", "median_ms", "min_ms", "max_ms", "us_per_token", "iters"]
395
+
396
+
397
+ def _prefill_sweep(once_factory, lengths=PREFILL_LENGTHS, label="") -> list[dict]:
398
+ rows = []
399
+ for seq_len in lengths:
400
+ once = once_factory(seq_len)
401
+ for _ in range(PREFILL_WARMUP):
402
+ once()
403
+ samples = []
404
+ for _ in range(PREFILL_ITERS):
405
+ t0 = time.perf_counter()
406
+ once()
407
+ samples.append((time.perf_counter() - t0) * 1e3)
408
+ median = statistics.median(samples)
409
+ logger.info(f"{label} prefill S={seq_len:>5}: median {median:8.2f} ms ({median / seq_len * 1e3:6.1f} us/token)")
410
+ rows.append(
411
+ {
412
+ "seq_len": seq_len,
413
+ "median_ms": round(median, 3),
414
+ "min_ms": round(min(samples), 3),
415
+ "max_ms": round(max(samples), 3),
416
+ "us_per_token": round(median / seq_len * 1e3, 2),
417
+ "iters": PREFILL_ITERS,
418
+ }
419
+ )
420
+ return rows
421
+
422
+
423
+ def _write_rows(path: Path, fieldnames: list[str], rows: list[dict]) -> None:
424
+ path.parent.mkdir(parents=True, exist_ok=True)
425
+ with path.open("w", newline="") as fh:
426
+ wr = csv.DictWriter(fh, fieldnames=fieldnames)
427
+ wr.writeheader()
428
+ wr.writerows(rows)
429
+ logger.info(f"wrote {path}")
430
+
431
+
432
+ def _traced_decode_median(mesh_device, step) -> tuple[float, float, float]:
433
+ step()
434
+ ttnn.synchronize_device(mesh_device)
435
+ trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0)
436
+ step()
437
+ ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0)
438
+ for _ in range(DECODE_WARMUP):
439
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
440
+ samples = []
441
+ for _ in range(DECODE_ITERS):
442
+ t0 = time.perf_counter()
443
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
444
+ samples.append((time.perf_counter() - t0) * 1e3)
445
+ ttnn.release_trace(mesh_device, trace_id)
446
+ return statistics.median(samples), min(samples), max(samples)
447
+
448
+
449
+ @pytest.mark.models_performance_bare_metal
450
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
451
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
452
+ def test_optimized_multichip_baseline_1x1_prefill(mesh_device, reference, torch_weights):
453
+ """Stage-04's own copy of the warmed single-chip prefill baseline.
454
+
455
+ ``optimized_decoder.py`` is untouched by stage 04, so this re-measures the
456
+ same code stage 03 did; it is re-run rather than quoted so the speedup
457
+ columns in this stage's README are one CSV cell divided by another taken in
458
+ the same session."""
459
+ from ..tt import optimized_decoder as O
460
+
461
+ _, hf_config = reference
462
+ config = DecoderLayerConfig.from_hf(hf_config)
463
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
464
+ packed = O.upload_packed_expert_weights(torch_weights, mesh_device, config.moe)
465
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
466
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
467
+
468
+ def factory(seq_len):
469
+ torch.manual_seed(0)
470
+ tt_in = _to_device(torch.randn(1, 1, seq_len, hf_config.hidden_size) * 0.02, mesh_device)
471
+
472
+ def once():
473
+ out = O.decoder_layer_prefill_optimized(tt_in, weights, config, cos_cache, sin_cache, sparsity, packed)
474
+ ttnn.synchronize_device(mesh_device)
475
+ ttnn.deallocate(out)
476
+
477
+ return once
478
+
479
+ rows = _prefill_sweep(factory, label="BASELINE 1x1")
480
+ _write_rows(MC_DOC_DIR / "perf_baseline_1x1_prefill.csv", PREFILL_FIELDS, rows)
481
+ assert all(r["median_ms"] > 0 for r in rows)
482
+
483
+
484
+ @pytest.mark.models_performance_bare_metal
485
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
486
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
487
+ @pytest.mark.parametrize("context_len", [128, 1024, 4096], ids=["ctx128", "ctx1k", "ctx4k"])
488
+ def test_optimized_multichip_baseline_1x1_decode(mesh_device, reference, torch_weights, context_len):
489
+ """Stage-04's own copy of the warmed single-chip traced decode baseline."""
490
+ from ..tt import optimized_decoder as O
491
+
492
+ _, hf_config = reference
493
+ config = DecoderLayerConfig.from_hf(hf_config)
494
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
495
+ packed = O.upload_packed_expert_weights(torch_weights, mesh_device, config.moe)
496
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
497
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ, block_size=BLOCK_SIZE)
498
+
499
+ pos = context_len - 1
500
+ torch.manual_seed(0)
501
+ tt_in = _to_device(torch.randn(1, 1, 1, hf_config.hidden_size) * 0.02, mesh_device)
502
+ current_pos = ttnn.from_torch(torch.tensor([pos], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
503
+
504
+ median, lo, hi = _traced_decode_median(
505
+ mesh_device,
506
+ lambda: O.decoder_layer_decode_optimized(
507
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, pos, packed_experts=packed
508
+ ),
509
+ )
510
+ logger.info(f"BASELINE 1x1 traced decode ctx={context_len:>5}: median {median:7.4f} ms")
511
+ _write_decode_row(
512
+ MC_DOC_DIR / "perf_baseline_1x1_decode.csv",
513
+ {
514
+ "context_len": context_len,
515
+ "median_ms": round(median, 4),
516
+ "min_ms": round(lo, 4),
517
+ "max_ms": round(hi, 4),
518
+ "tok_per_s_per_layer": round(1e3 / median, 1),
519
+ "iters": DECODE_ITERS,
520
+ },
521
+ )
522
+ assert median > 0
523
+
524
+
525
+ @pytest.mark.models_performance_bare_metal
526
+ @pytest.mark.parametrize("device_params", [MC_DEVICE_PARAMS], indirect=True)
527
+ @pytest.mark.parametrize("mesh_device", [MC_MESH], ids=["1x4"], indirect=True)
528
+ def test_optimized_multichip_prefill_latency(mesh_device, reference, torch_weights):
529
+ from ..tt import multichip_decoder as MC
530
+
531
+ _, hf_config = reference
532
+ config = MC.MeshDecoderConfig.from_hf(hf_config)
533
+ ctx = MC.mesh_context(mesh_device)
534
+ weights = MC.upload_multichip_weights(torch_weights, mesh_device, config)
535
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
536
+ sparsity = MC.build_local_sparsity(mesh_device, config.local_moe)
537
+
538
+ def factory(seq_len):
539
+ torch.manual_seed(0)
540
+ tt_in = ttnn.from_torch(
541
+ torch.randn(1, 1, seq_len, hf_config.hidden_size) * 0.02,
542
+ dtype=ttnn.bfloat16,
543
+ layout=ttnn.TILE_LAYOUT,
544
+ device=mesh_device,
545
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
546
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
547
+ )
548
+
549
+ def once():
550
+ out = MC.decoder_layer_prefill_multichip(tt_in, weights, config, ctx, cos_cache, sin_cache, sparsity)
551
+ ttnn.synchronize_device(mesh_device)
552
+ ttnn.deallocate(out)
553
+
554
+ return once
555
+
556
+ rows = _prefill_sweep(factory, label="MULTICHIP 1x4")
557
+ _write_rows(MC_DOC_DIR / "perf_prefill.csv", PREFILL_FIELDS, rows)
558
+ assert all(r["median_ms"] > 0 for r in rows)
559
+
560
+
561
+ @pytest.mark.models_performance_bare_metal
562
+ @pytest.mark.parametrize("device_params", [MC_DEVICE_PARAMS], indirect=True)
563
+ @pytest.mark.parametrize("mesh_device", [MC_MESH], ids=["1x4"], indirect=True)
564
+ @pytest.mark.parametrize("context_len", [128, 1024, 4096], ids=["ctx128", "ctx1k", "ctx4k"])
565
+ def test_optimized_multichip_decode_latency_traced(mesh_device, reference, torch_weights, context_len):
566
+ """Warmed trace replay on the mesh. Same harness as the 1x1 baseline above."""
567
+ from ..tt import multichip_decoder as MC
568
+
569
+ _, hf_config = reference
570
+ config = MC.MeshDecoderConfig.from_hf(hf_config)
571
+ ctx = MC.mesh_context(mesh_device)
572
+ weights = MC.upload_multichip_weights(torch_weights, mesh_device, config)
573
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
574
+ kv_cache = MC.create_mesh_kv_cache(mesh_device, config, 1, MAX_SEQ, block_size=BLOCK_SIZE)
575
+
576
+ pos = context_len - 1
577
+ torch.manual_seed(0)
578
+ tt_in = ttnn.from_torch(
579
+ torch.randn(1, 1, 1, hf_config.hidden_size) * 0.02,
580
+ dtype=ttnn.bfloat16,
581
+ layout=ttnn.TILE_LAYOUT,
582
+ device=mesh_device,
583
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
584
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
585
+ )
586
+ current_pos = ttnn.from_torch(
587
+ torch.tensor([pos], dtype=torch.int32),
588
+ dtype=ttnn.int32,
589
+ device=mesh_device,
590
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
591
+ )
592
+
593
+ median, lo, hi = _traced_decode_median(
594
+ mesh_device,
595
+ lambda: MC.decoder_layer_decode_multichip(
596
+ tt_in, weights, config, ctx, cos_cache, sin_cache, kv_cache, current_pos, pos
597
+ ),
598
+ )
599
+ logger.info(f"MULTICHIP 1x4 traced decode ctx={context_len:>5}: median {median:7.4f} ms")
600
+ _write_decode_row(
601
+ MC_DOC_DIR / "perf_decode.csv",
602
+ {
603
+ "context_len": context_len,
604
+ "median_ms": round(median, 4),
605
+ "min_ms": round(lo, 4),
606
+ "max_ms": round(hi, 4),
607
+ "tok_per_s_per_layer": round(1e3 / median, 1),
608
+ "iters": DECODE_ITERS,
609
+ },
610
+ )
611
+ assert median > 0
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_precision_config.py ADDED
@@ -0,0 +1,369 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Gates for ``tt/precision.py`` -- the precision config stage 07 sweeps.
5
+
6
+ The stage-07 goal asks for a selected precision config that "later
7
+ full-model/vLLM construction paths actually consume by default", and says
8
+ explicitly that **a JSON field ignored by hard-coded model code does not satisfy
9
+ this requirement**. So this file is arranged around that sentence:
10
+
11
+ * :func:`test_default_is_the_shipped_policy` and the alias tests pin that the
12
+ *default* changed nothing -- every constant stages 02-06 measured at is still
13
+ what a no-argument construction produces;
14
+ * :func:`test_non_default_precision_reaches_the_device` builds a real model on
15
+ the mesh at a **non-default** value and asserts on what the device actually
16
+ holds: a different weight dtype, a different program-config block width, a
17
+ different compute-kernel fidelity, and a smaller per-die expert allocation.
18
+ That is the assertion the goal is really asking for;
19
+ * the round-trip tests pin that config -> JSON -> config is lossless and that
20
+ the JSON carries every field.
21
+
22
+ The host-only tests need no device. The two device tests build **two-layer**
23
+ models -- the observable is per-layer and a 48-layer load is three minutes --
24
+ on a module-scoped mesh, exactly as ``test_full_model.py`` does.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import dataclasses
30
+ import json
31
+
32
+ import pytest
33
+
34
+ import ttnn
35
+
36
+ # **Absolute imports, deliberately.** ``tt/generator.py`` imports ``tt.model``
37
+ # by absolute path, and under this repo's ``--import-mode=importlib`` a relative
38
+ # ``from ..tt import model`` here resolves to a *second* copy of the module (no
39
+ # ``models/__init__.py``, so pytest roots the package at this directory).
40
+ # The identity assertions below would then be comparing two different classes,
41
+ # and the device tests would be inspecting a model built by the other copy.
42
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt import model as M
43
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt import multichip_decoder as MC
44
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt import optimized_decoder as O
45
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator import build_generator
46
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.model import DEFAULT_TRACE_REGION_SIZE
47
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.precision import DEFAULT_PRECISION, PrecisionConfig
48
+
49
+ MODEL_DIR = "models/demos/blackhole/qwen3_coder_30b_a3b"
50
+
51
+ #: Every field the stage-07 goal enumerates, mapped to the config field(s) that
52
+ #: carry it. The test below fails if any of these stops being serialised, which
53
+ #: is the cheap way to notice a field being dropped from the artifact.
54
+ GOAL_FIELDS = {
55
+ "experts gate_up weight dtype": ["experts_gate_up_dtype"],
56
+ "experts down weight dtype": ["experts_down_dtype"],
57
+ "attention qkv weight dtype": ["attention_qkv_dtype"],
58
+ "attention wo weight dtype": ["attention_wo_dtype"],
59
+ "lm_head weight dtype": ["lm_head_dtype"],
60
+ "router weight dtype": ["router_dtype"],
61
+ "embedding weight dtype": ["embedding_dtype"],
62
+ "per-group compute fidelity": [
63
+ "experts_fidelity",
64
+ "attention_fidelity",
65
+ "router_window_fidelity",
66
+ "lm_head_fidelity",
67
+ "norm_fidelity",
68
+ ],
69
+ "activation/residual dtype": ["activation_dtype"],
70
+ "CCL dtype": ["ccl_dtype"],
71
+ "KV-cache dtype": ["kv_cache_dtype"],
72
+ "logits/sampling dtype": ["logits_dtype", "sampling_dtype"],
73
+ }
74
+
75
+
76
+ # -- the default is the shipped policy ----------------------------------------
77
+
78
+
79
+ def test_default_is_the_shipped_policy():
80
+ """The literal values stages 02-06 measured at, re-asserted here.
81
+
82
+ Written out rather than compared against the module constants, which are now
83
+ *derived* from this config -- comparing them to each other would be a
84
+ tautology. If a shipped value is ever changed, this test is the thing that
85
+ has to be changed with it, deliberately.
86
+ """
87
+ p = DEFAULT_PRECISION
88
+ assert p.experts_gate_up_dtype is ttnn.bfloat4_b
89
+ assert p.experts_down_dtype is ttnn.bfloat4_b
90
+ # Stage 07: retuned to the full-K ceilings the 48-layer sweep measured.
91
+ assert p.experts_gate_up_in0_block_w == 64
92
+ assert p.experts_down_in0_block_w == 24
93
+ assert p.experts_fidelity is ttnn.MathFidelity.LoFi
94
+ assert p.attention_qkv_dtype is ttnn.bfloat8_b
95
+ assert p.attention_wo_dtype is ttnn.bfloat8_b
96
+ assert p.attention_fidelity is None, "the projections take the op default; see _attention_compute_kernel_config"
97
+ assert p.lm_head_dtype is ttnn.bfloat8_b
98
+ assert p.lm_head_fidelity is ttnn.MathFidelity.HiFi2
99
+ assert p.router_dtype is ttnn.bfloat16
100
+ assert p.router_window_fidelity is ttnn.MathFidelity.HiFi4, "the one-hot window matmul must select, not approximate"
101
+ assert p.embedding_dtype is ttnn.bfloat16
102
+ assert p.norm_weight_dtype is ttnn.bfloat16
103
+ assert p.norm_fidelity is ttnn.MathFidelity.HiFi4
104
+ assert p.activation_dtype is ttnn.bfloat16
105
+ assert p.ccl_dtype is None and p.effective_ccl_dtype is ttnn.bfloat16
106
+ assert p.kv_cache_dtype is ttnn.bfloat16
107
+ assert p.logits_dtype is ttnn.bfloat16
108
+ assert p.sampling_dtype is ttnn.bfloat16
109
+
110
+
111
+ def test_module_constants_still_resolve_to_the_default():
112
+ """The stage-02/04 names are aliases now; they must still read the same.
113
+
114
+ Probes under ``doc/`` and several stage-02 tests import these, and the point
115
+ of keeping them was that nothing outside this file had to change.
116
+ """
117
+ assert O.EXPERT_WEIGHT_DTYPE is DEFAULT_PRECISION.experts_gate_up_dtype
118
+ assert O.EXPERT_IN0_BLOCK_W_GATE_UP == DEFAULT_PRECISION.experts_gate_up_in0_block_w
119
+ assert O.EXPERT_IN0_BLOCK_W_DOWN == DEFAULT_PRECISION.experts_down_in0_block_w
120
+ assert O.EXPERT_MATH_FIDELITY is DEFAULT_PRECISION.experts_fidelity
121
+ assert O.ATTENTION_WEIGHT_DTYPE is DEFAULT_PRECISION.attention_qkv_dtype
122
+ assert M.LM_HEAD_WEIGHT_DTYPE is DEFAULT_PRECISION.lm_head_dtype
123
+ assert M.EMBED_WEIGHT_DTYPE is DEFAULT_PRECISION.embedding_dtype
124
+
125
+
126
+ def test_default_construction_paths_all_resolve_to_the_same_object():
127
+ """``None`` means the shipped policy on every entry point that takes one."""
128
+ assert M._resolve_precision(None) is DEFAULT_PRECISION
129
+ assert M._resolve_precision(DEFAULT_PRECISION) is DEFAULT_PRECISION
130
+ assert M._resolve_precision(DEFAULT_PRECISION.to_dict()) == DEFAULT_PRECISION
131
+
132
+
133
+ # -- serialisation -------------------------------------------------------------
134
+
135
+
136
+ def test_json_round_trip_is_lossless():
137
+ for config in (
138
+ DEFAULT_PRECISION,
139
+ DEFAULT_PRECISION.with_overrides(experts_gate_up_dtype="bfloat8_b", experts_gate_up_in0_block_w=32),
140
+ DEFAULT_PRECISION.with_overrides(attention_fidelity="HiFi4", ccl_dtype="bfloat8_b"),
141
+ ):
142
+ assert PrecisionConfig.from_json(config.to_json()) == config
143
+ # and again through a second hop, so an asymmetric coercion cannot hide
144
+ assert PrecisionConfig.from_json(PrecisionConfig.from_json(config.to_json()).to_json()) == config
145
+
146
+
147
+ def test_json_carries_every_field_the_goal_lists():
148
+ payload = json.loads(DEFAULT_PRECISION.to_json())
149
+ declared = {f.name for f in dataclasses.fields(PrecisionConfig)}
150
+ assert set(payload) == declared, "to_dict() must emit exactly the dataclass fields"
151
+ for description, names in GOAL_FIELDS.items():
152
+ for name in names:
153
+ assert name in payload, f"{description} is missing from the serialised config"
154
+
155
+
156
+ def test_json_is_plain_names_not_repr():
157
+ """The artifact has to be readable and diffable, not ``DataType.BFLOAT4_B``."""
158
+ payload = json.loads(DEFAULT_PRECISION.to_json())
159
+ assert payload["experts_gate_up_dtype"] == "bfloat4_b"
160
+ assert payload["experts_fidelity"] == "LoFi"
161
+ assert payload["attention_fidelity"] is None
162
+ assert payload["ccl_dtype"] is None
163
+ assert payload["experts_gate_up_in0_block_w"] == 64
164
+
165
+
166
+ def test_write_and_read_json_file(tmp_path):
167
+ config = DEFAULT_PRECISION.with_overrides(lm_head_dtype="bfloat16")
168
+ path = config.write_json(tmp_path / "nested" / "selected_precision_config.json")
169
+ assert PrecisionConfig.read_json(path) == config
170
+ # the file form is what a construction path is handed
171
+ assert M._resolve_precision(str(path)) == config
172
+
173
+
174
+ def test_unknown_names_are_rejected(expect_error):
175
+ with expect_error(ValueError, "unknown dtype"):
176
+ PrecisionConfig(experts_gate_up_dtype="bfloat3_b")
177
+ with expect_error(ValueError, "unknown math fidelity"):
178
+ PrecisionConfig(experts_fidelity="HiFi9")
179
+ with expect_error(ValueError, "unknown precision fields"):
180
+ PrecisionConfig.from_dict({**DEFAULT_PRECISION.to_dict(), "expert_dtype": "bfloat8_b"})
181
+ with expect_error(ValueError, "unknown precision fields"):
182
+ DEFAULT_PRECISION.with_overrides(expert_weight_dtype="bfloat8_b")
183
+ with expect_error(ValueError, "may not be None"):
184
+ PrecisionConfig(activation_dtype=None)
185
+
186
+
187
+ def test_config_is_frozen(expect_error):
188
+ # the fixture requires a match string; frozen dataclasses name the field
189
+ with expect_error(dataclasses.FrozenInstanceError, "experts_gate_up_dtype"):
190
+ DEFAULT_PRECISION.experts_gate_up_dtype = ttnn.bfloat16
191
+
192
+
193
+ # -- the config is actually consumed (device) ----------------------------------
194
+
195
+
196
+ @pytest.fixture(scope="module")
197
+ def mesh_device():
198
+ ttnn.set_fabric_config(ttnn.FabricConfig.FABRIC_1D_RING)
199
+ mesh = ttnn.open_mesh_device(mesh_shape=ttnn.MeshShape(*MC.MESH_SHAPE), trace_region_size=DEFAULT_TRACE_REGION_SIZE)
200
+ yield mesh
201
+ ttnn.close_mesh_device(mesh)
202
+ ttnn.set_fabric_config(ttnn.FabricConfig.DISABLED)
203
+
204
+
205
+ #: A precision that differs from the shipped one in four independently
206
+ #: observable ways: a wider expert weight, a different block width for it, a
207
+ #: different expert fidelity, and a narrower lm_head.
208
+ NON_DEFAULT = DEFAULT_PRECISION.with_overrides(
209
+ experts_gate_up_dtype="bfloat8_b",
210
+ experts_gate_up_in0_block_w=32, # also a divisor of 2048/32 = 64
211
+ experts_fidelity="HiFi4",
212
+ lm_head_dtype="bfloat4_b",
213
+ attention_wo_dtype="bfloat16",
214
+ )
215
+
216
+
217
+ def _build(mesh_device, precision):
218
+ return build_generator(
219
+ MODEL_DIR,
220
+ mesh_device,
221
+ override_num_layers=2,
222
+ max_context_len=1024,
223
+ max_batch_size=1,
224
+ precision=precision,
225
+ )
226
+
227
+
228
+ def test_non_default_precision_reaches_the_device(mesh_device):
229
+ """Construct at ``NON_DEFAULT`` and assert on what the *device* holds.
230
+
231
+ This is the goal's "a JSON field ignored by hard-coded model code does not
232
+ satisfy this requirement" test. Nothing here reads ``model.precision``: every
233
+ assertion is against a dtype read back off an uploaded tensor, a block width
234
+ resolved by ``_tuned_sparse_matmul_config``, or a byte count computed from
235
+ the allocated shape.
236
+ """
237
+ default_gen = _build(mesh_device, None)
238
+ try:
239
+ base = default_gen.model.runtime_fallback_audit()
240
+ base_lm_head = str(default_gen.model.lm_head.dtype)
241
+ finally:
242
+ default_gen.teardown()
243
+
244
+ gen = _build(mesh_device, NON_DEFAULT)
245
+ try:
246
+ audit = gen.model.runtime_fallback_audit()
247
+ lm_head_dtype = str(gen.model.lm_head.dtype)
248
+
249
+ # 1. a different weight dtype reached the device
250
+ assert base["device_experts_gate_up_dtype"] == str(ttnn.bfloat4_b)
251
+ assert audit["device_experts_gate_up_dtype"] == str(ttnn.bfloat8_b)
252
+ assert audit["device_attention_wo_dtype"] == str(ttnn.bfloat16)
253
+ assert base["device_attention_wo_dtype"] == str(ttnn.bfloat8_b)
254
+ assert base_lm_head == str(ttnn.bfloat8_b)
255
+ assert lm_head_dtype == str(ttnn.bfloat4_b)
256
+
257
+ # 2. gate/up moved but down did **not** -- the two expert groups are
258
+ # genuinely separate fields, not one knob wearing two names
259
+ assert audit["device_experts_down_dtype"] == str(ttnn.bfloat4_b)
260
+
261
+ # 3. a different block width in the resolved program config
262
+ assert base["gate_up_in0_block_w"] == 64
263
+ assert audit["gate_up_in0_block_w"] == 32
264
+ assert audit["down_in0_block_w"] == 24, "down's width was not overridden and must not move"
265
+
266
+ # 4. a different fidelity in the compute kernel config
267
+ assert base["expert_math_fidelity"] == str(ttnn.MathFidelity.LoFi)
268
+ assert audit["expert_math_fidelity"] == str(ttnn.MathFidelity.HiFi4)
269
+
270
+ # 5. an allocation-size change: bfloat4_b -> bfloat8_b on gate/up
271
+ # roughly doubles the gate/up half of the per-die expert footprint
272
+ assert audit["device_expert_bytes_per_die"] > base["device_expert_bytes_per_die"]
273
+ grew = audit["device_expert_bytes_per_die"] - base["device_expert_bytes_per_die"]
274
+ assert grew > 20 * 1024 * 1024, f"expected tens of MB per die, got {grew}"
275
+
276
+ # 6. and it still runs -- a config that reaches the device but wedges it
277
+ # would not be sweepable
278
+ ids = gen.tokenizer("def fib(n):", add_special_tokens=False)["input_ids"]
279
+ out = gen.generate(ids, 4, enable_trace=True, sampling_mode="device", top_k=1)
280
+ assert len(out) == 4
281
+ finally:
282
+ gen.teardown()
283
+
284
+
285
+ def test_default_construction_audit_matches_the_shipped_values(mesh_device):
286
+ """The default path puts the shipped dtypes and widths on the device."""
287
+ gen = _build(mesh_device, None)
288
+ try:
289
+ audit = gen.model.runtime_fallback_audit()
290
+ assert audit["device_experts_gate_up_dtype"] == str(ttnn.bfloat4_b)
291
+ assert audit["device_experts_down_dtype"] == str(ttnn.bfloat4_b)
292
+ assert audit["device_attention_qkv_dtype"] == str(ttnn.bfloat8_b)
293
+ assert audit["device_attention_wo_dtype"] == str(ttnn.bfloat8_b)
294
+ assert audit["device_attention_qkv_decode_dtype"] == str(ttnn.bfloat8_b)
295
+ assert audit["device_router_dtype"] == str(ttnn.bfloat16)
296
+ assert audit["device_norm_weight_dtype"] == str(ttnn.bfloat16)
297
+ assert audit["gate_up_in0_block_w"] == 64
298
+ assert audit["down_in0_block_w"] == 24
299
+ assert audit["expert_math_fidelity"] == str(ttnn.MathFidelity.LoFi)
300
+ assert audit["attention_math_fidelity"] is None
301
+ assert audit["router_window_math_fidelity"] == str(ttnn.MathFidelity.HiFi4)
302
+ assert audit["ccl_dtype"] == str(ttnn.bfloat16)
303
+ assert audit["activation_dtype"] == str(ttnn.bfloat16)
304
+ assert str(gen.model.ensure_internal_kv_cache()[0].k.dtype) == str(ttnn.bfloat16)
305
+ # the audit also reports the config itself, which is what a sweep row
306
+ # records alongside its measurement
307
+ assert audit["precision"] == DEFAULT_PRECISION.to_dict()
308
+ finally:
309
+ gen.teardown()
310
+
311
+
312
+ #: The four fields stage 07's original selection proof could not see.
313
+ #:
314
+ #: They had **no audit entry at all**, which made "this lever does nothing" and
315
+ #: "this lever is not wired up" indistinguishable -- three sweep rows produced
316
+ #: ``device_audit`` blocks byte-identical to the baseline's. For
317
+ #: ``norm_fidelity`` it was the second: ``decode_residual_norm`` built its
318
+ #: compute config from the module default and never saw ``self.precision``, so
319
+ #: the field was a documented knob with no effect and ``R21_norm_hifi2``
320
+ #: measured nothing. This config moves all four away from their defaults.
321
+ TERMINAL_AND_NORM = DEFAULT_PRECISION.with_overrides(
322
+ norm_fidelity="HiFi2",
323
+ lm_head_fidelity="LoFi",
324
+ logits_dtype="bfloat8_b",
325
+ sampling_dtype="bfloat8_b",
326
+ )
327
+
328
+
329
+ def test_fidelity_and_terminal_dtypes_reach_the_device(mesh_device):
330
+ """The four fields that used to change nothing observable must change it now.
331
+
332
+ Regression test for a dead config field. Every assertion is against
333
+ something the *device* or the *ops* hold: the fidelities come off the
334
+ ``compute_kernel_config`` objects the norm and lm_head are handed, the two
335
+ dtypes off the tensors the terminal path actually produced during a real
336
+ traced decode. A default-vs-override diff is asserted for each, so a field
337
+ that silently stops being threaded fails here rather than in a sweep row
338
+ three stages later.
339
+ """
340
+ gen = _build(mesh_device, None)
341
+ try:
342
+ ids = gen.tokenizer("def fib(n):", add_special_tokens=False)["input_ids"]
343
+ gen.generate(ids, 4, enable_trace=True, sampling_mode="device", top_k=1)
344
+ base = gen.model.runtime_fallback_audit()
345
+ finally:
346
+ gen.teardown()
347
+
348
+ assert base["norm_math_fidelity"] == str(ttnn.MathFidelity.HiFi4)
349
+ assert base["lm_head_math_fidelity"] == str(ttnn.MathFidelity.HiFi2)
350
+ assert base["logits_dtype_observed"] == "bfloat16"
351
+ assert base["sampling_dtype_observed"] == "bfloat16"
352
+ assert base["terminal_dtype_source"] == "device_readback"
353
+
354
+ gen = _build(mesh_device, TERMINAL_AND_NORM)
355
+ try:
356
+ ids = gen.tokenizer("def fib(n):", add_special_tokens=False)["input_ids"]
357
+ out = gen.generate(ids, 4, enable_trace=True, sampling_mode="device", top_k=1)
358
+ assert len(out) == 4, "a config that reaches the device but wedges it is not sweepable"
359
+ audit = gen.model.runtime_fallback_audit()
360
+ finally:
361
+ gen.teardown()
362
+
363
+ # norm_fidelity: the field that was NOT threaded. It reaches only the decode
364
+ # residual norms -- the prefill norms pass no compute config at all -- which
365
+ # is why the audit name is norm_math_fidelity rather than something global.
366
+ assert audit["norm_math_fidelity"] == str(ttnn.MathFidelity.HiFi2)
367
+ assert audit["lm_head_math_fidelity"] == str(ttnn.MathFidelity.LoFi)
368
+ assert audit["logits_dtype_observed"] == "bfloat8_b"
369
+ assert audit["sampling_dtype_observed"] == "bfloat8_b"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_reference.py ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Validate the layer-only reference before any TTNN work is built against it.
5
+
6
+ The reference does one thing the checkpoint does not: it fuses each expert's
7
+ ``gate_proj``/``up_proj`` into a single ``gate_up_proj`` and stacks all experts.
8
+ If that fusion were reversed, the layer would still run and produce
9
+ plausible-looking numbers -- it would just be wrong. Everything downstream is
10
+ compared against this reference, so an error here is invisible forever after.
11
+
12
+ ``test_moe_matches_unfused_reimplementation`` therefore recomputes the MoE block
13
+ from the raw per-expert checkpoint tensors, following
14
+ ``Qwen3MoeSparseMoeBlock.forward`` literally, and requires the two to agree.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import pytest
20
+ import torch
21
+
22
+ from models.common.utility_functions import comp_pcc
23
+
24
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
25
+
26
+ LAYER_IDX = 0
27
+ SEQ_LEN = 32
28
+
29
+
30
+ @pytest.fixture(scope="module")
31
+ def reference():
32
+ layer, config = build_reference_layer(LAYER_IDX)
33
+ return layer, config
34
+
35
+
36
+ def _hidden(config, seq_len=SEQ_LEN, seed=0):
37
+ """Activations roughly matching what reaches a decoder layer post-embedding."""
38
+ torch.manual_seed(seed)
39
+ return torch.randn(1, seq_len, config.hidden_size, dtype=torch.float32) * 0.02
40
+
41
+
42
+ def test_layer_forward_is_finite_and_non_degenerate(reference):
43
+ layer, config = reference
44
+ hidden = _hidden(config)
45
+ cos, sin = rotary_embeddings(config, SEQ_LEN)
46
+
47
+ with torch.no_grad():
48
+ out = layer(hidden, position_embeddings=(cos, sin), attention_mask=None)
49
+ out = out[0] if isinstance(out, tuple) else out
50
+
51
+ assert out.shape == hidden.shape
52
+ assert torch.isfinite(out).all()
53
+ assert out.std() > 1e-6, "output is constant -- layer is not doing anything"
54
+ assert not torch.allclose(out, hidden), "output equals input -- residual-only path"
55
+
56
+
57
+ def test_layer_forward_is_deterministic(reference):
58
+ layer, config = reference
59
+ hidden = _hidden(config)
60
+ cos, sin = rotary_embeddings(config, SEQ_LEN)
61
+
62
+ with torch.no_grad():
63
+ a = layer(hidden, position_embeddings=(cos, sin), attention_mask=None)
64
+ b = layer(hidden, position_embeddings=(cos, sin), attention_mask=None)
65
+ a = a[0] if isinstance(a, tuple) else a
66
+ b = b[0] if isinstance(b, tuple) else b
67
+
68
+ assert torch.equal(a, b), "reference is not deterministic; PCC comparisons would be unstable"
69
+
70
+
71
+ def test_moe_matches_unfused_reimplementation(reference):
72
+ """Recompute the MoE block from raw per-expert tensors and require agreement.
73
+
74
+ This is the guard on the gate/up fusion order and on the router's
75
+ softmax -> top-k -> renormalise ordering.
76
+ """
77
+ layer, config = reference
78
+ # build_reference_layer upcasts the bf16 checkpoint into the fp32 module, so
79
+ # cast here too -- this test is about the fusion order, not about dtype.
80
+ sd = {k: v.float() for k, v in layer_state_dict(LAYER_IDX).items()}
81
+ hidden = _hidden(config)
82
+ flat = hidden.view(-1, config.hidden_size)
83
+
84
+ with torch.no_grad():
85
+ fused_out = layer.mlp(hidden).view(-1, config.hidden_size)
86
+
87
+ # --- independent implementation, straight from Qwen3MoeTopKRouter.forward ---
88
+ logits = torch.nn.functional.linear(flat, sd["mlp.gate.weight"])
89
+ probs = torch.softmax(logits, dim=-1, dtype=torch.float) # over ALL experts, fp32
90
+ top_w, top_i = torch.topk(probs, config.num_experts_per_tok, dim=-1)
91
+ if config.norm_topk_prob:
92
+ top_w = top_w / top_w.sum(dim=-1, keepdim=True)
93
+ top_w = top_w.to(logits.dtype)
94
+
95
+ # --- and from Qwen3MoeExperts.forward, but with unfused checkpoint tensors ---
96
+ manual = torch.zeros_like(flat)
97
+ for token in range(flat.shape[0]):
98
+ for slot in range(config.num_experts_per_tok):
99
+ e = int(top_i[token, slot])
100
+ x = flat[token]
101
+ gate = torch.nn.functional.linear(x, sd[f"mlp.experts.{e}.gate_proj.weight"])
102
+ up = torch.nn.functional.linear(x, sd[f"mlp.experts.{e}.up_proj.weight"])
103
+ h = torch.nn.functional.silu(gate) * up
104
+ manual[token] += torch.nn.functional.linear(h, sd[f"mlp.experts.{e}.down_proj.weight"]) * top_w[token, slot]
105
+
106
+ passing, message = comp_pcc(manual, fused_out, pcc=0.9999)
107
+ assert passing, f"fused MoE disagrees with unfused checkpoint math: {message}"
108
+
109
+
110
+ def test_router_selects_expected_expert_count(reference):
111
+ """top_k experts per token, weights renormalised to 1."""
112
+ layer, config = reference
113
+ hidden = _hidden(config)
114
+ flat = hidden.view(-1, config.hidden_size)
115
+
116
+ with torch.no_grad():
117
+ _, scores, indices = layer.mlp.gate(flat)
118
+
119
+ assert indices.shape == (flat.shape[0], config.num_experts_per_tok)
120
+ assert indices.min() >= 0 and indices.max() < config.num_experts
121
+ if config.norm_topk_prob:
122
+ sums = scores.float().sum(dim=-1)
123
+ assert torch.allclose(sums, torch.ones_like(sums), atol=1e-5), f"router weights not normalised: {sums[:4]}"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_rmsnorm.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """First on-device module for Qwen3-Coder-30B-A3B: RMSNorm vs the HF reference.
5
+
6
+ RMSNorm is the smallest piece of the decoder layer, so it is brought up first --
7
+ it proves the whole harness (open a 1x1 mesh, upload real checkpoint weights,
8
+ run, compare PCC against the layer-only reference) before GQA, RoPE and the MoE
9
+ block are layered on top. When attention PCC later misbehaves, normalisation is
10
+ already ruled out.
11
+
12
+ Two Qwen3-specific details are asserted rather than assumed:
13
+ * eps is 1e-6, not the module default of 1e-5 -- so the config path is used.
14
+ * the norm is the plain variant; ``add_unit_offset`` stays False. Qwen3.5/3.6
15
+ use a zero-centred RMSNorm that folds a "+1" into the weight; this
16
+ checkpoint does not.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ from pathlib import Path
22
+
23
+ import pytest
24
+ import torch
25
+ from loguru import logger
26
+
27
+ import ttnn
28
+ from models.common.auto_compose import to_torch_auto_compose
29
+ from models.common.modules.lazy_weight import LazyWeight
30
+ from models.common.modules.rmsnorm.rmsnorm_1d import RMSNorm1D, RMSNorm1DConfig
31
+ from models.common.utility_functions import comp_allclose, comp_pcc
32
+
33
+ from .reference import build_reference_layer
34
+
35
+ LAYER_IDX = 0
36
+ PCC_REQUIRED = 0.999 # normalisation alone should be near-exact; the layer bar is 0.995
37
+
38
+
39
+ @pytest.fixture(scope="module")
40
+ def reference():
41
+ return build_reference_layer(LAYER_IDX)
42
+
43
+
44
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
45
+ @pytest.mark.parametrize(
46
+ "norm_name,seq_len,mode",
47
+ [
48
+ ("input_layernorm", 32, "prefill"),
49
+ ("input_layernorm", 128, "prefill"),
50
+ ("post_attention_layernorm", 32, "prefill"),
51
+ ("input_layernorm", 1, "decode"),
52
+ ],
53
+ ids=["input_s32", "input_s128", "postattn_s32", "input_decode"],
54
+ )
55
+ def test_rmsnorm_vs_reference(
56
+ mesh_device: ttnn.MeshDevice,
57
+ reference,
58
+ norm_name: str,
59
+ seq_len: int,
60
+ mode: str,
61
+ ):
62
+ layer, config = reference
63
+ torch.manual_seed(0)
64
+
65
+ ref_norm = getattr(layer, norm_name)
66
+ dim = config.hidden_size
67
+ assert config.rms_norm_eps == 1e-6, f"unexpected eps {config.rms_norm_eps}"
68
+
69
+ # Activations scaled like what actually reaches a decoder layer, per the
70
+ # skill's guidance -- not arbitrary large randoms.
71
+ torch_input = (torch.randn(1, 1, seq_len, dim, dtype=torch.float32) * 0.02).to(torch.bfloat16)
72
+
73
+ ttnn.SetDefaultDevice(mesh_device)
74
+ try:
75
+ cache_dir = Path("model_cache/qwen3_coder_30b_a3b/rmsnorm")
76
+ lazy_weight = LazyWeight(
77
+ source=ref_norm.weight.data.clone(),
78
+ dtype=ttnn.bfloat16,
79
+ cache_dir_weight_name=(cache_dir, f"{norm_name}_L{LAYER_IDX}"),
80
+ )
81
+ tt_model = RMSNorm1D.from_config(
82
+ RMSNorm1DConfig(
83
+ weight=lazy_weight,
84
+ eps=config.rms_norm_eps,
85
+ add_unit_offset=False, # plain RMSNorm, not the zero-centred Qwen3.5/3.6 variant
86
+ )
87
+ )
88
+ tt_out = tt_model.forward(LazyWeight(source=torch_input, dtype=ttnn.bfloat16), mode=mode)
89
+ tt_out_torch = to_torch_auto_compose(tt_out)
90
+ finally:
91
+ ttnn.SetDefaultDevice(None)
92
+
93
+ with torch.no_grad():
94
+ ref_out = ref_norm(torch_input.to(torch.float32)).to(torch.bfloat16)
95
+
96
+ passing, pcc_message = comp_pcc(ref_out, tt_out_torch, PCC_REQUIRED)
97
+ logger.info(comp_allclose(ref_out, tt_out_torch))
98
+ logger.info(f"RMSNorm[{norm_name}] {mode} seq={seq_len}: {pcc_message}")
99
+ assert passing, f"RMSNorm {norm_name} ({mode}, seq={seq_len}) below {PCC_REQUIRED}: {pcc_message}"
code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_trace.py ADDED
@@ -0,0 +1,200 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Capture the decode step as a trace and replay it.
5
+
6
+ Tracing is the real test of every "on-device, trace-compatible" claim made
7
+ while building this layer. A trace records device commands once and replays
8
+ them, so anything that is not a pure device op -- a host round-trip, a
9
+ Python-side branch on tensor *values*, a shape that depends on the data --
10
+ either fails to capture or replays stale results. Notably, this is what makes
11
+ the router's design load-bearing: it keeps top-k selection and the scatter on
12
+ device precisely so this step is possible.
13
+
14
+ Two properties are checked, and the second is the one that catches real bugs:
15
+
16
+ 1. the traced output matches the eager output for the same input;
17
+ 2. replaying with *different* input produces *different*, still-correct output.
18
+
19
+ Property 2 is essential. A trace whose input tensor was captured by value
20
+ rather than written in place replays the original activations forever and
21
+ therefore passes property 1 perfectly, every time.
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ import pytest
27
+ import torch
28
+ from loguru import logger
29
+
30
+ import ttnn
31
+ from models.common.utility_functions import comp_allclose, comp_pcc
32
+
33
+ from ..tt.functional_decoder import (
34
+ DecoderLayerConfig,
35
+ build_expert_sparsity,
36
+ build_rope_cache,
37
+ create_kv_cache,
38
+ decoder_layer_decode,
39
+ decoder_layer_prefill,
40
+ upload_layer_weights,
41
+ )
42
+ from ..tt.weight_mapping import convert_layer_weights
43
+ from .reference import build_reference_layer, layer_state_dict, rotary_embeddings
44
+
45
+ LAYER_IDX = 0
46
+ PCC_REQUIRED = 0.99
47
+ MAX_SEQ = 256
48
+ PROMPT_LEN = 32
49
+ # Reserved at device open; the capture fails outright if the graph needs more.
50
+ TRACE_REGION_SIZE = 50331648
51
+
52
+
53
+ @pytest.fixture(scope="module")
54
+ def reference():
55
+ return build_reference_layer(LAYER_IDX)
56
+
57
+
58
+ @pytest.fixture(scope="module")
59
+ def torch_weights(reference):
60
+ _, hf_config = reference
61
+ return convert_layer_weights(layer_state_dict(LAYER_IDX), hf_config)
62
+
63
+
64
+ def _hidden(hf_config, seq_len, seed=0):
65
+ torch.manual_seed(seed)
66
+ return torch.randn(1, seq_len, hf_config.hidden_size, dtype=torch.float32) * 0.02
67
+
68
+
69
+ def _reference_layer(layer, hf_config, hidden):
70
+ seq_len = hidden.shape[1]
71
+ cos, sin = rotary_embeddings(hf_config, seq_len)
72
+ mask = torch.full((seq_len, seq_len), float("-inf")).triu(1).reshape(1, 1, seq_len, seq_len)
73
+ with torch.no_grad():
74
+ out = layer(hidden, position_embeddings=(cos, sin), attention_mask=mask)
75
+ return out[0] if isinstance(out, tuple) else out
76
+
77
+
78
+ def _to_device(t, mesh_device):
79
+ return ttnn.from_torch(
80
+ t, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG
81
+ )
82
+
83
+
84
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
85
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
86
+ def test_decode_step_is_traceable(mesh_device, reference, torch_weights):
87
+ layer, hf_config = reference
88
+ config = DecoderLayerConfig.from_hf(hf_config)
89
+
90
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
91
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
92
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
93
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ)
94
+
95
+ hidden_full = _hidden(hf_config, PROMPT_LEN + 2)
96
+ ref_out = _reference_layer(layer, hf_config, hidden_full)
97
+
98
+ decoder_layer_prefill(
99
+ _to_device(hidden_full[:, :PROMPT_LEN, :].unsqueeze(0), mesh_device),
100
+ weights,
101
+ config,
102
+ cos_cache,
103
+ sin_cache,
104
+ sparsity,
105
+ kv_cache=kv_cache,
106
+ )
107
+
108
+ # Persistent input buffers: a trace replays writes to the *same* addresses,
109
+ # so inputs must be updated in place rather than rebound each step.
110
+ tt_in = _to_device(hidden_full[:, PROMPT_LEN, :].reshape(1, 1, 1, hf_config.hidden_size), mesh_device)
111
+ current_pos = ttnn.from_torch(torch.tensor([PROMPT_LEN], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
112
+
113
+ def step():
114
+ return decoder_layer_decode(
115
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, token_index=PROMPT_LEN
116
+ )
117
+
118
+ # Warm up so program compilation happens outside the capture.
119
+ eager_out = ttnn.to_torch(step()).reshape(1, hf_config.hidden_size)
120
+
121
+ trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0)
122
+ traced_out = step()
123
+ ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0)
124
+
125
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
126
+ replayed = ttnn.to_torch(traced_out).reshape(1, hf_config.hidden_size)
127
+
128
+ passing, pcc_message = comp_pcc(eager_out, replayed, 0.999)
129
+ logger.info(comp_allclose(eager_out, replayed))
130
+ logger.info(f"traced vs eager: {pcc_message}")
131
+ assert passing, f"traced replay disagrees with eager execution: {pcc_message}"
132
+
133
+ passing, pcc_message = comp_pcc(ref_out[:, PROMPT_LEN, :], replayed, PCC_REQUIRED)
134
+ logger.info(f"traced vs reference: {pcc_message}")
135
+ assert passing, f"traced output below {PCC_REQUIRED} vs reference: {pcc_message}"
136
+
137
+ ttnn.release_trace(mesh_device, trace_id)
138
+
139
+
140
+ @pytest.mark.parametrize("device_params", [{"trace_region_size": TRACE_REGION_SIZE}], indirect=True)
141
+ @pytest.mark.parametrize("mesh_device", [(1, 1)], ids=["1x1"], indirect=True)
142
+ def test_traced_replay_follows_new_input(mesh_device, reference, torch_weights):
143
+ """Writing a new token into the input buffer must change the traced output.
144
+
145
+ Guards the failure mode a same-input trace test cannot see: if the capture
146
+ bound the input by value, replay reproduces the first token's result
147
+ forever and every equality check still passes.
148
+ """
149
+ layer, hf_config = reference
150
+ config = DecoderLayerConfig.from_hf(hf_config)
151
+
152
+ weights = upload_layer_weights(torch_weights, mesh_device, config)
153
+ cos_cache, sin_cache = build_rope_cache(hf_config, MAX_SEQ, mesh_device)
154
+ sparsity = build_expert_sparsity(mesh_device, config.moe.num_experts)
155
+ kv_cache = create_kv_cache(mesh_device, config.attention, max_batch=1, max_seq_len=MAX_SEQ)
156
+
157
+ hidden_full = _hidden(hf_config, PROMPT_LEN + 1)
158
+ decoder_layer_prefill(
159
+ _to_device(hidden_full[:, :PROMPT_LEN, :].unsqueeze(0), mesh_device),
160
+ weights,
161
+ config,
162
+ cos_cache,
163
+ sin_cache,
164
+ sparsity,
165
+ kv_cache=kv_cache,
166
+ )
167
+
168
+ token_a = hidden_full[:, PROMPT_LEN, :].reshape(1, 1, 1, hf_config.hidden_size)
169
+ token_b = (_hidden(hf_config, 1, seed=99)).reshape(1, 1, 1, hf_config.hidden_size)
170
+
171
+ tt_in = _to_device(token_a, mesh_device)
172
+ current_pos = ttnn.from_torch(torch.tensor([PROMPT_LEN], dtype=torch.int32), dtype=ttnn.int32, device=mesh_device)
173
+
174
+ def step():
175
+ return decoder_layer_decode(
176
+ tt_in, weights, config, cos_cache, sin_cache, kv_cache, current_pos, token_index=PROMPT_LEN
177
+ )
178
+
179
+ step() # warm up / compile
180
+
181
+ trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0)
182
+ traced_out = step()
183
+ ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0)
184
+
185
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
186
+ out_a = ttnn.to_torch(traced_out).reshape(-1).float().clone()
187
+
188
+ # Overwrite the captured input buffer in place, then replay.
189
+ ttnn.copy_host_to_device_tensor(ttnn.from_torch(token_b, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT), tt_in)
190
+ ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True)
191
+ out_b = ttnn.to_torch(traced_out).reshape(-1).float().clone()
192
+
193
+ delta = (out_a - out_b).abs().max().item()
194
+ logger.info(f"max|out_a - out_b| after swapping the input token = {delta:.6f}")
195
+ assert delta > 1e-3, (
196
+ "traced replay produced identical output for a different input token -- "
197
+ "the trace is not reading the live input buffer"
198
+ )
199
+
200
+ ttnn.release_trace(mesh_device, trace_id)
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt-model-localplugin.yaml ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # A/B VARIANT of tt-model.yaml: identical EXCEPT that the vLLM plugin comes from a local
2
+ # checkout rather than a cloned ref. Published as a separate repo so both packaging paths
3
+ # can be exercised end to end.
4
+ #
5
+ # This one file is the whole authoring interface:
6
+ # tt-model package --container models/demos/blackhole/qwen3_coder_30b_a3b/tt-model.yaml
7
+ #
8
+ # It lives next to the model on purpose: the serving recipe then travels through review
9
+ # in the same PR as the code it serves.
10
+ schema: "5.1"
11
+
12
+ repo: raahemnabeel/qwen3-coder-30b-a3b-localplugin
13
+ name: qwen3-coder-30b-a3b-localplugin
14
+ weights: Qwen/Qwen3-Coder-30B-A3B-Instruct # a POINTER — 57 GB, never baked into the image
15
+ kind: vllm-plugin # stock vLLM + the standalone tenstorrent/vllm-tt-plugin
16
+ arch: blackhole
17
+
18
+ source:
19
+ tt_metal: /home/raahem/metal-publish
20
+
21
+ # EXACTLY what ships — an allowlist, never a denylist. The image's tt-metal tree holds
22
+ # NO models/ except these, so an under-specified list fails the image's own build-time
23
+ # import check on the author's machine, not on a consumer's first boot.
24
+ code:
25
+ # tt/ imports models.common.modules.sampling.sampling_1d,
26
+ # models.common.modules.tt_ccl and models.common.readiness_check.contract
27
+ - models/common
28
+ # the whole model, 1.6 MB: tt/, vllm_bundle/, config/, tests/, README.md
29
+ - models/demos/blackhole/qwen3_coder_30b_a3b
30
+
31
+ ubuntu: "24.04"
32
+ python: "3.12"
33
+
34
+ runtime:
35
+ # The plugin monkeypatches vLLM internals, so this pin is load-bearing, not cosmetic.
36
+ # This box serves with 0.24.0 (reported 0.24.0+empty — the local tag comes from the
37
+ # VLLM_TARGET_DEVICE=empty sdist build, which is what the image does too).
38
+ vllm: {version: "0.24.0"}
39
+
40
+ # A/B VARIANT: the author's LOCAL plugin checkout, staged into the build context the
41
+ # same way source.tt_metal is — uncommitted work included, nothing fetched. The sibling
42
+ # manifest (tt-model.yaml) pins a pushed SHA and clones it during the build instead.
43
+ # Both produce an image whose venv holds the plugin; serving cannot tell them apart.
44
+ plugin:
45
+ path: /home/raahem/vllm-tt-plugin
46
+
47
+ # The directory the plugin SCANS — it walks the CHILDREN of this dir for
48
+ # vllm_metadata.json. Registration chain:
49
+ # register_tt_models() -> _register_models_from_extra_dir() -> this dir ->
50
+ # vllm_bundle/qwen3_coder_30b_a3b_instruct/vllm_metadata.json ->
51
+ # Qwen3MoeForCausalLM registered as TTQwen3MoeForCausalLM ->
52
+ # "tt_qwen3_coder_30b_a3b_instruct:Qwen3CoderForCausalLM" -> the shim ->
53
+ # models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator_vllm:Qwen3CoderForCausalLM
54
+ # The shim's parents[6] still resolves in-container: the tree lands at
55
+ # /opt/tt-metal/models/demos/blackhole/<model>/vllm_bundle/<bundle>/, so six levels up
56
+ # is /opt/tt-metal, which is also PYTHONPATH. Setting this also makes the build pass
57
+ # TT_VLLM_BUILTIN_MODELS=0, so registration comes solely from here.
58
+ extra_models_dir: models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle
59
+
60
+ # Uncomment once requirements.lock is committed next to this file, after the first
61
+ # successful build — later builds then resolve nothing at all.
62
+ # lock: requirements.lock
63
+
64
+ serve:
65
+ port: 8000
66
+
67
+ # 256000, NOT the 262144 HF advertises. config/context_contract.json records
68
+ # current_supported_context: 256000 with capability_reduction: true, and
69
+ # tt/generator_vllm.py reads that file as the source of truth and REFUSES to start
70
+ # above it. (Reason on record: an unresolved prefill cliff in the last ~3,100 tokens,
71
+ # and 17.5 min TTFT at that length even on the path that works.)
72
+ max_model_len: 256000
73
+
74
+ # 64, from the validated tt-inference-server P300X2 spec, which is authoritative here.
75
+ # config/context_contract.json's prose says "paging is unchanged (block_size 32)"; that
76
+ # text is older than the serving config and the spec supersedes it.
77
+ block_size: 64
78
+
79
+ # One configuration, measured. max_num_seqs only bounds what vLLM batches IN FLIGHT, and
80
+ # with the stage-11 decode-width ladder now the model default, batch-1 decode is already
81
+ # efficient at this cap — so a lower cap buys nothing. Measured on this box, 128 output
82
+ # tokens, tt-model 0.1.0:
83
+ #
84
+ # seqs=1 seqs=32
85
+ # single-stream TTFT 0.232 s 0.219 s
86
+ # single-stream decode 50.6 tok/s 49.2 tok/s
87
+ # 8 concurrent, total 46.7 tok/s 71.4 tok/s
88
+ # 8 concurrent, /user 5.8 tok/s 8.9 tok/s
89
+ # slowest of the 8 21.9 s 14.4 s
90
+ #
91
+ # seqs=32 is identical for one user and strictly better for several, so there is no
92
+ # latency-vs-capacity trade to expose as separate profiles. Untested above 8 concurrent,
93
+ # and TTFT-under-load is unmeasured; if either turns out to matter, `serve_profiles:`
94
+ # is how you would split this again.
95
+ hardware: p300x2
96
+ mesh_device: P300x2
97
+ max_num_seqs: 32
98
+
99
+ # Structured, not a JSON string: tt-model renders it into --additional-config.
100
+ additional_config:
101
+ tt:
102
+ sample_on_device_mode: all
103
+ trace_region_size: 50331648
104
+ fabric_config: FABRIC_1D_RING
105
+
106
+ # Tool calling. tt-model emits --enable-auto-tool-choice alongside the parser, because
107
+ # vLLM hard-errors on the parser flag without it.
108
+ capabilities:
109
+ tool_parser: qwen3_coder
110
+
111
+ # Runtime environment, from the validated tt-inference-server P300X2 spec.
112
+ # MESH_DEVICE is NOT listed: it is derived from each profile's mesh_device.
113
+ # VLLM_TARGET_DEVICE is NOT listed either: it is a BUILD-time variable for the vLLM
114
+ # fork, and this image builds stock vLLM with VLLM_TARGET_DEVICE=empty plus the plugin.
115
+ env:
116
+ ARCH_NAME: blackhole
117
+ VLLM_CONFIGURE_LOGGING: "1"
118
+ VLLM_RPC_TIMEOUT: "900000"
119
+ VLLM_ALLOW_LONG_MAX_MODEL_LEN: "1"
120
+ TORCHDYNAMO_DISABLE: "1"
121
+ TT_METAL_OPERATION_TIMEOUT_SECONDS: "120.0"
122
+ # The variable-width decode ladder (stage 11, doc/batch_scaling). This is the model's
123
+ # default now, so setting it changes nothing today — it is kept to pin the behaviour
124
+ # explicitly, so a future change to that default cannot silently move this model.
125
+ QWEN3_DECODE_WIDTHS: "1,2,4,8,16,32"
126
+
127
+ # Anything without a named field goes here, verbatim.
128
+ args:
129
+ - [--max-num-batched-tokens, "256000"]
130
+ - [--max-log-len, "32"]
131
+ - [--generation-config, vllm]
132
+ - [--seed, "9472"]
133
+
134
+
135
+ # Build-time assertions, run INSIDE the finished image after the USER switch. These make
136
+ # the code/ allowlist and the tree prune safe: an under-shipped image fails HERE.
137
+ verify:
138
+ # config/ is RUNTIME DATA, not documentation. tt/generator_vllm.py reads both files;
139
+ # _supported_context() swallows OSError and silently falls back to in-code defaults, so
140
+ # a missing file would serve at the wrong precision/context with no error at all.
141
+ - "from pathlib import Path; p = Path('/opt/tt-metal/models/demos/blackhole/qwen3_coder_30b_a3b'); assert (p/'config'/'selected_precision_config.json').is_file(), 'precision config missing — the model would silently use in-code defaults'; assert (p/'config'/'context_contract.json').is_file(), 'context contract missing — _supported_context() would silently fall back'"
142
+ # the real adapter. (The bundle shim is NOT asserted here: it is importable only once
143
+ # something has put its folder on sys.path, which is what the plugin does. tt-model
144
+ # verifies that resolution generically for every model — see launchers.py's
145
+ # RESOLVE_EXTRA_MODELS — so a per-model assertion would only duplicate it, badly.)
146
+ - "from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator_vllm import Qwen3CoderForCausalLM; assert Qwen3CoderForCausalLM"
147
+
148
+ card:
149
+ quickstart: |
150
+ ### Use it
151
+ An OpenAI-compatible server on `http://127.0.0.1:8000`. Point any OpenAI client at it
152
+ with model id `Qwen/Qwen3-Coder-30B-A3B-Instruct`.
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt-model.yaml ADDED
@@ -0,0 +1,242 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # tt-model container package (v5.1) — qwen3-coder-30b-a3b on QB2 (2x P300, 4 chips).
2
+ #
3
+ # This one file is the whole authoring interface:
4
+ # tt-model package --container models/demos/blackhole/qwen3_coder_30b_a3b/tt-model.yaml
5
+ #
6
+ # It lives next to the model on purpose: the serving recipe then travels through review
7
+ # in the same PR as the code it serves.
8
+ schema: "5.1"
9
+
10
+ repo: raahemnabeel/qwen3-coder-30b-a3b
11
+ name: qwen3-coder-30b-a3b
12
+ weights: Qwen/Qwen3-Coder-30B-A3B-Instruct # a POINTER — 57 GB, never baked into the image
13
+ kind: vllm-plugin # stock vLLM + the standalone tenstorrent/vllm-tt-plugin
14
+ arch: blackhole
15
+
16
+ source:
17
+ tt_metal: /home/raahem/metal-publish
18
+
19
+ # EXACTLY what ships — an allowlist, never a denylist. The image's tt-metal tree holds
20
+ # NO models/ except these, so an under-specified list fails the image's own build-time
21
+ # import check on the author's machine, not on a consumer's first boot.
22
+ code:
23
+ # tt/ imports models.common.modules.sampling.sampling_1d,
24
+ # models.common.modules.tt_ccl and models.common.readiness_check.contract
25
+ - models/common
26
+ # the whole model, 1.6 MB: tt/, vllm_bundle/, config/, tests/, README.md
27
+ - models/demos/blackhole/qwen3_coder_30b_a3b
28
+
29
+ ubuntu: "24.04"
30
+ python: "3.12"
31
+
32
+ runtime:
33
+ # The plugin monkeypatches vLLM internals, so this pin is load-bearing, not cosmetic.
34
+ # This box serves with 0.24.0 (reported 0.24.0+empty — the local tag comes from the
35
+ # VLLM_TARGET_DEVICE=empty sdist build, which is what the image does too).
36
+ vllm: {version: "0.24.0"}
37
+
38
+ # THE PLUGIN, bundled from your own checkout — the same hermetic treatment
39
+ # source.tt_metal gets. Whatever is in this directory ships, uncommitted work
40
+ # included; nothing is fetched at build time and nothing is resolved at serve time.
41
+ # `package` records its HEAD sha and flags the tree dirty if it is.
42
+ #
43
+ # Alternatives, when a local checkout is not the right input:
44
+ # plugin: {repo: https://github.com/tenstorrent/vllm-tt-plugin, ref: <pushed sha>}
45
+ # plugin: {version: "0.1.0"}
46
+ plugin:
47
+ path: /home/raahem/vllm-tt-plugin
48
+
49
+ # The directory the plugin SCANS — it walks the CHILDREN of this dir for
50
+ # vllm_metadata.json. Registration chain:
51
+ # register_tt_models() -> _register_models_from_extra_dir() -> this dir ->
52
+ # vllm_bundle/qwen3_coder_30b_a3b_instruct/vllm_metadata.json ->
53
+ # Qwen3MoeForCausalLM registered as TTQwen3MoeForCausalLM ->
54
+ # "tt_qwen3_coder_30b_a3b_instruct:Qwen3CoderForCausalLM" -> the shim ->
55
+ # models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator_vllm:Qwen3CoderForCausalLM
56
+ # The shim's parents[6] still resolves in-container: the tree lands at
57
+ # /opt/tt-metal/models/demos/blackhole/<model>/vllm_bundle/<bundle>/, so six levels up
58
+ # is /opt/tt-metal, which is also PYTHONPATH. Setting this also makes the build pass
59
+ # TT_VLLM_BUILTIN_MODELS=0, so registration comes solely from here.
60
+ extra_models_dir: models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle
61
+
62
+ # Uncomment once requirements.lock is committed next to this file, after the first
63
+ # successful build — later builds then resolve nothing at all.
64
+ # lock: requirements.lock
65
+
66
+ serve:
67
+ # The port the image binds under a bare `docker run`. `tt-model serve` does NOT use this
68
+ # as its seed — it opens 20000 and walks past busy ports — so leave it at vLLM's default.
69
+ port: 8000
70
+
71
+ # 256000, NOT the 262144 HF advertises. config/context_contract.json records
72
+ # current_supported_context: 256000 with capability_reduction: true, and
73
+ # tt/generator_vllm.py reads that file as the source of truth and REFUSES to start
74
+ # above it. (Reason on record: an unresolved prefill cliff in the last ~3,100 tokens,
75
+ # and 17.5 min TTFT at that length even on the path that works.)
76
+ max_model_len: 256000
77
+
78
+ # 64, from the validated tt-inference-server P300X2 spec, which is authoritative here.
79
+ # config/context_contract.json's prose says "paging is unchanged (block_size 32)"; that
80
+ # text is older than the serving config and the spec supersedes it.
81
+ block_size: 64
82
+
83
+ # One configuration, measured. max_num_seqs only bounds what vLLM batches IN FLIGHT, and
84
+ # with the stage-11 decode-width ladder now the model default, batch-1 decode is already
85
+ # efficient at this cap — so a lower cap buys nothing. Measured on this box, 128 output
86
+ # tokens, tt-model 0.1.0:
87
+ #
88
+ # (All figures below are p300x2. The p150x4 profile is validated by the author but its
89
+ # throughput was not separately measured, so do not read these numbers as applying to it.)
90
+ #
91
+ # seqs=1 seqs=32
92
+ # single-stream TTFT 0.232 s 0.219 s
93
+ # single-stream decode 50.6 tok/s 49.2 tok/s
94
+ # 8 concurrent, total 46.7 tok/s 71.4 tok/s
95
+ # 8 concurrent, /user 5.8 tok/s 8.9 tok/s
96
+ # slowest of the 8 21.9 s 14.4 s
97
+ #
98
+ # seqs=32 is identical for one user and strictly better for several, so there is no
99
+ # latency-vs-capacity trade to expose as separate profiles. Untested above 8 concurrent,
100
+ # and TTFT-under-load is unmeasured; if either turns out to matter, `serve_profiles:`
101
+ # is how you would split this again.
102
+ #
103
+ # Those decode figures are the ARGMAX path, which the --override-generation-config below
104
+ # now makes the default request — so they describe what an unparameterised client gets.
105
+ # A client that asks for sampling gets 38.6 tok/s instead; that gap is the strategy, not
106
+ # the configuration. Re-measured on this box after the warmup change: unparameterised
107
+ # 49.147 tok/s at 0.304 s TTFT first request, 8 concurrent 77.55 tok/s aggregate
108
+ # (9.69 tok/s/u), max TTFT 0.90 s.
109
+ # hardware/mesh_device are the ONLY per-profile fields — see serve_profiles below.
110
+ max_num_seqs: 32
111
+
112
+ # Structured, not a JSON string: tt-model renders it into --additional-config.
113
+ additional_config:
114
+ tt:
115
+ sample_on_device_mode: all
116
+ trace_region_size: 50331648
117
+ fabric_config: FABRIC_1D_RING
118
+
119
+ # Tool calling. tt-model emits --enable-auto-tool-choice alongside the parser, because
120
+ # vLLM hard-errors on the parser flag without it.
121
+ capabilities:
122
+ tool_parser: qwen3_coder
123
+
124
+ # Runtime environment, from the validated tt-inference-server P300X2 spec.
125
+ # MESH_DEVICE is NOT listed: it is derived from each profile's mesh_device.
126
+ # VLLM_TARGET_DEVICE is NOT listed either: it is a BUILD-time variable for the vLLM
127
+ # fork, and this image builds stock vLLM with VLLM_TARGET_DEVICE=empty plus the plugin.
128
+ env:
129
+ ARCH_NAME: blackhole
130
+ VLLM_CONFIGURE_LOGGING: "1"
131
+ VLLM_RPC_TIMEOUT: "900000"
132
+ VLLM_ALLOW_LONG_MAX_MODEL_LEN: "1"
133
+ TORCHDYNAMO_DISABLE: "1"
134
+ TT_METAL_OPERATION_TIMEOUT_SECONDS: "120.0"
135
+ # The variable-width decode ladder (stage 11, doc/batch_scaling). This is the model's
136
+ # default now, so setting it changes nothing today — it is kept to pin the behaviour
137
+ # explicitly, so a future change to that default cannot silently move this model.
138
+ QWEN3_DECODE_WIDTHS: "1,2,4,8,16,32"
139
+ # Warm a decode graph for every ladder width in BOTH sampling strategies at startup,
140
+ # rather than compiling one on whichever request first needs it. Costs 9.5 s of boot
141
+ # (decode warmup 2.41 s -> 11.89 s) and 24.8 MB/die; buys 1.64 s off the first
142
+ # request's TTFT (1.9416 s -> 0.2989 s) and removes the same stall from the first
143
+ # request at any other width or strategy. "both" is the model default; it is pinned
144
+ # here for the same reason QWEN3_DECODE_WIDTHS is.
145
+ QWEN3_WARMUP_SAMPLING: "both"
146
+ # Prefill bucketing (doc/prefill_buckets). A prefill program is compiled per exact
147
+ # sequence length, so without this every new prompt length paid a fresh compile —
148
+ # measured on 4 dies, cold cache: 4–11 s on top of the 0.89 ms/token floor, and
149
+ # recurring, not a startup cost (2443 tok cost 22.87 s, then 2444 tok cost 4.29 s
150
+ # again). Bucketing rounds each prefill to a ladder rung so the shape space is
151
+ # finite and warmable. "proportional" bounds the padding waste at 1/8th of the
152
+ # prefill at every length; measured steady-state tax is 2–3%, against 26% for a
153
+ # half-power-of-two ladder and ~111 s worst case for plain powers of two.
154
+ QWEN3_PREFILL_BUCKETS: "proportional"
155
+ # Warm the ladder once per HOST, not once per boot. The rungs and the cached-suffix
156
+ # shapes cost 217.9 s to compile on a cold cache and ~68 k tokens of prefill to run,
157
+ # so "auto" runs them only when its marker is absent from TT_METAL_CACHE — which
158
+ # tt-model mounts per model and keeps across container removal on purpose. Clearing
159
+ # that cache clears the marker, and the next boot re-warms. "full" ignores the
160
+ # marker; "off" skips it and leaves every shape to compile on a real request.
161
+ QWEN3_PREFILL_WARMUP: "auto"
162
+ # Warm every rung up to and including the one that covers 8192 tokens (23 rungs).
163
+ # Longer prompts still compile lazily, once each, and persist in the same cache.
164
+ # Raising this lengthens only the first boot on a fresh host; "0" warms all 52.
165
+ QWEN3_PREFILL_WARMUP_MAX: "8192"
166
+ # The cached-suffix (prefix-cache hit) path is a separate program set, keyed on the
167
+ # chunked-SDPA chunk size AND the suffix length. Reachable chunk sizes are just
168
+ # {64, 128, 256} at block_size 64, so this caps the second axis: suffix rungs up to
169
+ # 1024, which is where a hit's new tail actually lands.
170
+ QWEN3_PREFILL_WARMUP_SUFFIX_MAX: "1024"
171
+
172
+ # Anything without a named field goes here, verbatim.
173
+ args:
174
+ - [--max-num-batched-tokens, "256000"]
175
+ - [--max-log-len, "32"]
176
+ - [--generation-config, vllm]
177
+ # A GREEDY DEFAULT, measured. `--generation-config vllm` alone means an
178
+ # unparameterised request arrives at vLLM's own temperature=1.0 / top_p=1.0, and the
179
+ # model routes anything with k>1 or p>0 to split sampling: 25.88 ms/token (38.6 t/s/u)
180
+ # against 20.35 ms (49.1) for argmax, on 4 dies, 24-token completions. For a coding
181
+ # model that is the wrong default three times over — 21 % slower, non-deterministic
182
+ # across identical prompts, and full-distribution sampling at temperature 1.0 is the
183
+ # worst of the three available regimes for code. Overriding temperature to 0 makes the
184
+ # default request greedy (verified: 49.147 t/s/u unparameterised) while leaving
185
+ # repetition_penalty unset, so the decode graph stays on penalty mode 0. A client that
186
+ # wants sampling still sends temperature itself and gets the split path, warm — the
187
+ # model warms both strategies at startup (QWEN3_WARMUP_SAMPLING below).
188
+ #
189
+ # QUOTING IS LOAD-BEARING. launchers.py joins these with spaces into ONE
190
+ # Written BARE. The extra single quotes this used to carry were a workaround for the
191
+ # vllm-fork path, which joins args into one --additional-server-args string that the
192
+ # runner shlex-splits: bare JSON lost its double quotes there and reached vLLM as
193
+ # {temperature: 0}. This model is vllm-plugin, which passes argv loose -- compose_run
194
+ # hands docker the list and the entrypoint ends in `exec "$@"`, so nothing consumes a
195
+ # shell quote and the literal ' characters reached json.loads instead. tt-model now
196
+ # shlex-quotes the fork path itself, so bare is correct for both kinds.
197
+ - [--override-generation-config, '{"temperature": 0}']
198
+ - [--seed, "9472"]
199
+
200
+
201
+ # TWO BOARDS, ONE CONFIGURATION. p300x2 (2x P300) and p150x4 (4x P150) are both four
202
+ # Blackhole chips on a (1, 4) mesh — tt-model derives 4 from either label and
203
+ # parse_mesh_device resolves both SKUs to the same (1, 4) — so nothing above this line
204
+ # differs between them: same context bound, same block size, same slot count, same
205
+ # trace region, same fabric, same env, same args. Both were tested on hardware.
206
+ #
207
+ # They are separate profiles rather than one because `hardware` is what tt-model uses to
208
+ # state a package's device requirement, and a p150x4 box should not have to read "p300x2"
209
+ # and infer that the chip counts happen to match.
210
+ #
211
+ # p300x2 stays the default: it is the box every number in this file was measured on.
212
+ # `tt-model serve --profile p150x4` selects the other; `tt-model profiles` lists both.
213
+ default_profile: p300x2
214
+ serve_profiles:
215
+ - name: p300x2
216
+ hardware: p300x2
217
+ mesh_device: P300x2
218
+ - name: p150x4
219
+ hardware: p150x4
220
+ mesh_device: P150x4
221
+
222
+ # Build-time assertions, run INSIDE the finished image after the USER switch. These make
223
+ # the code/ allowlist and the tree prune safe: an under-shipped image fails HERE.
224
+ verify:
225
+ # config/ is RUNTIME DATA, not documentation. tt/generator_vllm.py reads both files;
226
+ # _supported_context() swallows OSError and silently falls back to in-code defaults, so
227
+ # a missing file would serve at the wrong precision/context with no error at all.
228
+ - "from pathlib import Path; p = Path('/opt/tt-metal/models/demos/blackhole/qwen3_coder_30b_a3b'); assert (p/'config'/'selected_precision_config.json').is_file(), 'precision config missing — the model would silently use in-code defaults'; assert (p/'config'/'context_contract.json').is_file(), 'context contract missing — _supported_context() would silently fall back'"
229
+ # the real adapter. (The bundle shim is NOT asserted here: it is importable only once
230
+ # something has put its folder on sys.path, which is what the plugin does. tt-model
231
+ # verifies that resolution generically for every model — see launchers.py's
232
+ # RESOLVE_EXTRA_MODELS — so a per-model assertion would only duplicate it, badly.)
233
+ - "from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator_vllm import Qwen3CoderForCausalLM; assert Qwen3CoderForCausalLM"
234
+
235
+ card:
236
+ description: >
237
+ Qwen3-Coder-30B-A3B-Instruct — a 30B mixture-of-experts coding model for agentic
238
+ development work, with tool calling and a 256K context, served on Blackhole via vLLM.
239
+ quickstart: |
240
+ ### Use it
241
+ Point any OpenAI client at the address `tt-model serve` prints (port 20000 unless it
242
+ was busy), with model id `Qwen/Qwen3-Coder-30B-A3B-Instruct`.
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/__init__.py ADDED
File without changes
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/functional_decoder.py ADDED
@@ -0,0 +1,1169 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """TTNN decoder layer for Qwen3-Coder-30B-A3B-Instruct.
5
+
6
+ Built bottom-up and validated piece by piece against the layer-only HuggingFace
7
+ reference in ``tests/reference.py``. Attention lands first; the MoE block and
8
+ the composed layer follow.
9
+
10
+ Shape of this model's attention
11
+ -------------------------------
12
+ hidden 2048 | 32 Q heads | 4 KV heads (8:1 GQA) | head_dim 128
13
+
14
+ ``32 * 128 = 4096 != hidden``, so ``o_proj`` is ``[2048, 4096]`` -- head_dim is
15
+ an independent config field here, not ``hidden / n_heads``. There is no
16
+ attention bias and no sliding window, so every layer is plain causal attention.
17
+
18
+ Two Qwen3-specific things the more common Llama-shaped ports do not have:
19
+
20
+ * **QK-norm.** An RMSNorm over ``head_dim`` is applied per head to Q and to K,
21
+ after the head split and *before* RoPE. Note it is Q and K only -- gemma4,
22
+ which this file follows closely, additionally norms V because Gemma has a
23
+ ``v_norm``. Qwen3 does not; adding it there would silently corrupt V.
24
+
25
+ * **RoPE theta lives in ``rope_scaling``.** ``config.rope_theta`` does not
26
+ exist on this checkpoint; the value (1e7, the long-context setting) is
27
+ nested at ``config.rope_scaling["rope_theta"]``. Rather than reach in and
28
+ risk a stale default, the cos/sin cache is produced by
29
+ ``Qwen3MoeRotaryEmbedding`` itself, which is correct by construction.
30
+
31
+ RoPE convention is HF-style throughout -- see the note in ``weight_mapping.py``.
32
+ Weights therefore keep their checkpoint channel order.
33
+ """
34
+
35
+ from __future__ import annotations
36
+
37
+ import math
38
+ from dataclasses import dataclass
39
+
40
+ import torch
41
+
42
+ import ttnn
43
+
44
+
45
+ @dataclass(frozen=True)
46
+ class AttentionConfig:
47
+ hidden_size: int
48
+ num_attention_heads: int
49
+ num_key_value_heads: int
50
+ head_dim: int
51
+ rms_norm_eps: float
52
+
53
+ @classmethod
54
+ def from_hf(cls, config) -> "AttentionConfig":
55
+ assert not getattr(config, "attention_bias", False), "attention bias is not wired up"
56
+ assert not getattr(config, "use_sliding_window", False), "sliding-window attention is not wired up"
57
+ return cls(
58
+ hidden_size=config.hidden_size,
59
+ num_attention_heads=config.num_attention_heads,
60
+ num_key_value_heads=config.num_key_value_heads,
61
+ head_dim=config.head_dim,
62
+ rms_norm_eps=config.rms_norm_eps,
63
+ )
64
+
65
+
66
+ @dataclass
67
+ class AttentionWeights:
68
+ """Device-resident attention weights, as produced by ``weight_mapping``."""
69
+
70
+ wqkv: ttnn.Tensor # [1, 1, hidden, (n_heads + 2*n_kv_heads) * head_dim]
71
+ wo: ttnn.Tensor # [1, 1, n_heads * head_dim, hidden]
72
+ q_norm: ttnn.Tensor # [1, 1, 1, head_dim]
73
+ k_norm: ttnn.Tensor # [1, 1, 1, head_dim]
74
+
75
+
76
+ def build_rope_cache(hf_config, max_seq_len: int, device) -> tuple[ttnn.Tensor, ttnn.Tensor]:
77
+ """Upload a ``[1, 1, max_seq_len, head_dim]`` cos/sin pair for HF-style RoPE.
78
+
79
+ Generated by the reference rotary module so ``rope_theta`` (and any future
80
+ scaling) can never drift from the checkpoint.
81
+ """
82
+ from ..tests.reference import rotary_embeddings
83
+
84
+ cos, sin = rotary_embeddings(hf_config, max_seq_len) # [1, S, head_dim]
85
+ out = []
86
+ for t in (cos, sin):
87
+ out.append(
88
+ ttnn.from_torch(
89
+ t.unsqueeze(0).float(),
90
+ dtype=ttnn.bfloat16,
91
+ layout=ttnn.TILE_LAYOUT,
92
+ device=device,
93
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
94
+ )
95
+ )
96
+ return out[0], out[1]
97
+
98
+
99
+ def _per_head_rms_norm(tensor: ttnn.Tensor, weight: ttnn.Tensor, eps: float) -> ttnn.Tensor:
100
+ """RMSNorm over the last dim, applied independently per head.
101
+
102
+ ``ttnn.rms_norm`` normalises the last dimension of a 2D-ish tile grid, so
103
+ the heads are folded into the row dimension and unfolded afterwards.
104
+ Input/output ``[1, n_heads, seq, head_dim]``.
105
+ """
106
+ shape = tensor.shape
107
+ flat = ttnn.reshape(tensor, (1, 1, shape[1] * shape[2], shape[3]))
108
+ normed = ttnn.rms_norm(flat, weight=weight, epsilon=eps)
109
+ return ttnn.reshape(normed, shape)
110
+
111
+
112
+ def _apply_rope(tensor: ttnn.Tensor, cos_cache, sin_cache, token_index=None) -> ttnn.Tensor:
113
+ """HF-style rotary embedding, restoring the logical dim-2 length afterwards.
114
+
115
+ The op pads dim 2 up to a tile multiple, and dim 2 means different things in
116
+ the two modes -- which makes this a hazard twice over:
117
+
118
+ * **decode**: dim 2 is the head count. Q already has 32 heads so nothing
119
+ happens, but K has only ``num_kv_heads`` = 4 and comes back 32-deep. The
120
+ reshape declares logical-vs-padded before slicing, following the
121
+ tt_transformers ``_hf_rope_decode`` pattern.
122
+
123
+ * **prefill**: dim 2 is the sequence. At a non-tile-aligned length (say 33)
124
+ Q and K return padded to 64 while **V never passes through RoPE** and
125
+ stays 33, so SDPA rejects them with "K and V sequence length must match".
126
+ Slicing back is what keeps odd prompt lengths working.
127
+ """
128
+ orig = tensor.shape
129
+ out = ttnn.experimental.rotary_embedding(tensor, cos_cache, sin_cache, token_index)
130
+ if out.shape[2] != orig[2]:
131
+ if token_index is not None:
132
+ out = ttnn.reshape(out, (orig[0], orig[1], orig[2], orig[3]), (orig[0], orig[1], 32, orig[3]))
133
+ out = out[:, :, : orig[2]]
134
+ else:
135
+ out = ttnn.slice(out, [0, 0, 0, 0], [orig[0], orig[1], orig[2], orig[3]])
136
+ return out
137
+
138
+
139
+ def rope_transformation_matrix() -> "torch.Tensor":
140
+ """The 32x32 matrix ``rotary_embedding_llama`` rotates a tile row with.
141
+
142
+ ``+1`` at ``(2i, 2i+1)`` and ``-1`` at ``(2i+1, 2i)`` -- i.e. the Meta
143
+ channel pairing, expressed as a matmul so the kernel needs no gather. This
144
+ is what makes the llama op cheaper than the HF one at the same core count:
145
+ the HF op reads a cos/sin row out of a DRAM cache per call, this one
146
+ multiplies by a resident 32x32 tile.
147
+ """
148
+ import torch
149
+
150
+ d = 32
151
+ m = torch.zeros(1, 1, d, d)
152
+ m[..., torch.arange(0, d, 2), torch.arange(1, d, 2)] = 1
153
+ m[..., torch.arange(1, d, 2), torch.arange(0, d, 2)] = -1
154
+ return m
155
+
156
+
157
+ def apply_rope_llama(tensor: ttnn.Tensor, cos_shard, sin_shard, trans_mat) -> ttnn.Tensor:
158
+ """Meta-style rotary embedding, decode mode.
159
+
160
+ ``tensor``, ``cos_shard`` and ``sin_shard`` must all be the height-sharded
161
+ ``[1, batch, 32, head_dim]`` L1 config ``nlp_create_qkv_heads_decode``
162
+ produces, and ``tensor`` must be in **Meta** channel order -- which is a
163
+ property of the *weights*, established once at upload by
164
+ ``weight_mapping.permute_wqkv_to_meta``, not of anything done here.
165
+
166
+ Measured against ``_apply_rope`` at the shipped per-die decode shape:
167
+ 3.84 -> 1.26 us, ``max|diff|`` exactly 0.0 and PCC 1.0000000
168
+ (``doc/optimized_multichip_decoder/probes/rope_probe.py``).
169
+
170
+ Unlike ``_apply_rope`` this needs **no** reshape-and-slice afterwards: the
171
+ op is shape-preserving in decode mode, so the dim-2 padding hazard that
172
+ ``_apply_rope`` documents does not arise.
173
+ """
174
+ return ttnn.experimental.rotary_embedding_llama(tensor, cos_shard, sin_shard, trans_mat, is_decode_mode=True)
175
+
176
+
177
+ def _concat_heads_decode(attn: ttnn.Tensor, config: AttentionConfig) -> ttnn.Tensor:
178
+ """Merge heads in decode mode. ``attn`` ``[1, batch, num_heads, head_dim]``.
179
+
180
+ ``nlp_concat_heads_decode`` is the multi-core decode variant and requires a
181
+ height-sharded input with one core per user, whereas SDPA-decode hands back
182
+ a DRAM interleaved tensor -- hence the reshard. The op also pads batch up to
183
+ a full tile, so the result is sliced back to the logical batch.
184
+ """
185
+ batch = attn.shape[1]
186
+
187
+ grid_x = min(batch, 8)
188
+ while batch % grid_x:
189
+ grid_x -= 1
190
+ grid_y = batch // grid_x
191
+ core_grid = ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(grid_x - 1, grid_y - 1))})
192
+ shard_cfg = ttnn.create_sharded_memory_config(
193
+ shape=(ttnn.TILE_SIZE, config.head_dim),
194
+ core_grid=core_grid,
195
+ strategy=ttnn.ShardStrategy.HEIGHT,
196
+ orientation=ttnn.ShardOrientation.ROW_MAJOR,
197
+ use_height_and_width_as_shard_shape=True,
198
+ )
199
+
200
+ sharded = ttnn.to_memory_config(attn, shard_cfg)
201
+ out = ttnn.experimental.nlp_concat_heads_decode(sharded, num_heads=config.num_attention_heads)
202
+ ttnn.deallocate(sharded)
203
+ out = ttnn.sharded_to_interleaved(out, ttnn.DRAM_MEMORY_CONFIG)
204
+ if out.shape[2] != batch:
205
+ out = out[:, :, :batch, :]
206
+ return out
207
+
208
+
209
+ @dataclass
210
+ class KVCache:
211
+ """K/V cache, either contiguous per user or paged through a block table.
212
+
213
+ Paged mode is what a serving stack actually uses: logical positions are
214
+ mapped to physical blocks by ``page_table``, so users can share a block pool
215
+ instead of each reserving ``max_seq_len``. Both modes are kept because they
216
+ exercise different kernels -- ``paged_*`` ops versus their contiguous
217
+ counterparts -- and the contiguous path is the simpler thing to bisect
218
+ against when a paged result looks wrong.
219
+ """
220
+
221
+ k: ttnn.Tensor
222
+ v: ttnn.Tensor
223
+ page_table: ttnn.Tensor | None = None
224
+ block_size: int = 0
225
+
226
+ @property
227
+ def is_paged(self) -> bool:
228
+ return self.page_table is not None
229
+
230
+
231
+ def create_kv_cache(
232
+ device,
233
+ config: AttentionConfig,
234
+ max_batch: int,
235
+ max_seq_len: int,
236
+ block_size: int | None = None,
237
+ ) -> KVCache:
238
+ """Allocate a KV cache.
239
+
240
+ ``block_size=None`` gives a contiguous cache of
241
+ ``[max_batch, n_kv_heads, max_seq_len, head_dim]``. Passing a block size
242
+ switches to a paged cache of ``[num_blocks, n_kv_heads, block_size,
243
+ head_dim]`` with an identity page table -- block ``b`` of user ``u`` lives
244
+ at physical block ``u * blocks_per_seq + b``. A real scheduler would hand
245
+ out blocks from a free pool; the identity mapping keeps the plumbing honest
246
+ (every op still goes through the table) without pulling a block allocator
247
+ into the decoder.
248
+ """
249
+ n_kv, head_dim = config.num_key_value_heads, config.head_dim
250
+
251
+ if block_size is None:
252
+ shape = (max_batch, n_kv, max_seq_len, head_dim)
253
+ page_table = None
254
+ else:
255
+ blocks_per_seq = math.ceil(max_seq_len / block_size)
256
+ shape = (max_batch * blocks_per_seq, n_kv, block_size, head_dim)
257
+ page_table = ttnn.from_torch(
258
+ torch.arange(max_batch * blocks_per_seq, dtype=torch.int32).reshape(max_batch, blocks_per_seq),
259
+ dtype=ttnn.int32,
260
+ layout=ttnn.ROW_MAJOR_LAYOUT,
261
+ device=device,
262
+ )
263
+
264
+ k, v = (
265
+ ttnn.from_torch(
266
+ torch.zeros(shape),
267
+ dtype=ttnn.bfloat16,
268
+ layout=ttnn.TILE_LAYOUT,
269
+ device=device,
270
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
271
+ )
272
+ for _ in range(2)
273
+ )
274
+ return KVCache(k=k, v=v, page_table=page_table, block_size=block_size or 0)
275
+
276
+
277
+ def match_cache_dtype(cache: ttnn.Tensor, x: ttnn.Tensor) -> ttnn.Tensor:
278
+ """Return ``x`` in the cache tensor's dtype, casting only if they differ.
279
+
280
+ **For the fill (prefill) writers only.** The two cache writers have
281
+ *opposite* dtype contracts, and neither says so out loud, which is how
282
+ stage 07's ``R19_kv_bfp8`` came to score at chance.
283
+ ``doc/datatype_sweep/probes/kv_bfp8_diagnosis.py`` measures both at the op
284
+ level, with no model:
285
+
286
+ ================== ============ ============ =============================
287
+ op cache input round-trip PCC
288
+ ================== ============ ============ =============================
289
+ paged_fill_cache bfloat16 bfloat16 1.0 (control)
290
+ paged_fill_cache bfloat8_b bfloat16 **NaN**
291
+ paged_fill_cache bfloat8_b bfloat8_b 1.0
292
+ paged_update_cache bfloat16 bfloat16 1.0 (control)
293
+ paged_update_cache bfloat8_b bfloat16 0.999969
294
+ paged_update_cache bfloat8_b bfloat8_b **rejected by the op**
295
+ ================== ============ ============ =============================
296
+
297
+ So ``paged_fill_cache`` needs the input cast **to** the cache dtype -- it
298
+ validates the input against a permissive ``OR`` that a mismatch satisfies,
299
+ and then writes NaN -- while ``paged_update_cache`` needs the input left
300
+ **alone**: it converts into the cache itself and hard-rejects a block-float
301
+ input (``paged_update_cache_device_operation.cpp:296``, *"Data type of input
302
+ tensor for update cache must be FLOAT32 or BFLOAT16"*). Casting at the
303
+ decode writer would turn silent corruption into a hard crash, which is why
304
+ this helper is applied at the fill sites and deliberately not at the update
305
+ sites.
306
+
307
+ Taking the dtype off the **cache tensor itself** rather than off
308
+ ``precision.kv_cache_dtype`` keeps the guarantee true by construction: the
309
+ thing the write must agree with is the allocated cache, and this cannot
310
+ drift from it even if a caller allocates a cache some other way.
311
+
312
+ The cast is a no-op in the shipped configuration (``kv_cache_dtype ==
313
+ activation_dtype == bfloat16``), so it costs nothing unless the cache dtype
314
+ is actually moved.
315
+ """
316
+ if x.dtype == cache.dtype:
317
+ return x
318
+ return ttnn.typecast(x, cache.dtype, memory_config=x.memory_config())
319
+
320
+
321
+ #: Which prefill attention branch ran, and at what chunk size. Instrumentation
322
+ #: only -- nothing reads it in production. It exists because a PCC of ~1.0 on a
323
+ #: split prefill has two explanations, "the arithmetic coincided" and "the
324
+ #: chunked branch never executed", and they are indistinguishable from the PCC
325
+ #: alone. A probe can assert `chunked > 0` and make the question un-askable.
326
+ PREFILL_ATTENTION_BRANCHES = {"standard": 0, "chunked": 0, "chunk_sizes": []}
327
+
328
+
329
+ def sdpa_chunk_size(chunk_start_idx: int) -> int:
330
+ """q_chunk_size == k_chunk_size for a chunked-prefill offset.
331
+
332
+ ``chunk_start_idx`` must be a multiple of BOTH q_chunk_size and k_chunk_size
333
+ (sdpa_nanobind.cpp:487-493), so the largest legal power of two is the offset's
334
+ own lowest set bit. Capped at 256: 512 overflows L1 -- measured
335
+ (1760704 B against a 1572864 B limit) and already recorded at
336
+ multichip_decoder.py:972 as rejected at every length.
337
+
338
+ Derivation copied from models/tt_transformers model_config.py:1552-1582
339
+ rather than hardcoded, so it tracks upstream.
340
+ """
341
+ return min(256, chunk_start_idx & -chunk_start_idx)
342
+
343
+
344
+ def _fill_cache(kv_cache: KVCache, k: ttnn.Tensor, v: ttnn.Tensor, user_id: int, fill_page_table=None) -> None:
345
+ """Write a whole prompt's K/V into the cache.
346
+
347
+ The paged kernel writes block-at-a-time, so a prompt that does not fill its
348
+ last block is zero-padded up to a block boundary first. Those trailing
349
+ positions are never read: decode passes ``cur_pos``, and SDPA only attends
350
+ up to it.
351
+
352
+ K/V are cast to the cache's own dtype before the write -- see
353
+ :func:`match_cache_dtype` for why the ops will not do it for us. The cast
354
+ goes *after* the pad, because ``ttnn.pad`` is a bfloat16/float32 op.
355
+ """
356
+ if not kv_cache.is_paged:
357
+ ttnn.fill_cache(kv_cache.k, match_cache_dtype(kv_cache.k, k), user_id)
358
+ ttnn.fill_cache(kv_cache.v, match_cache_dtype(kv_cache.v, v), user_id)
359
+ return
360
+
361
+ seq_len = k.shape[2]
362
+ padded = math.ceil(seq_len / kv_cache.block_size) * kv_cache.block_size
363
+ if padded != seq_len:
364
+ pad = [(0, 0), (0, 0), (0, padded - seq_len), (0, 0)]
365
+ k = ttnn.pad(k, pad, value=0.0)
366
+ v = ttnn.pad(v, pad, value=0.0)
367
+
368
+ # ``fill_page_table`` is the suffix write of a split prefill: a SINGLE-ROW
369
+ # table already sliced to the blocks the suffix occupies, so the op's
370
+ # block-0-relative write lands at the right absolute offset. batch_idx is 0
371
+ # because the row has already been selected. None => the shipped whole-prompt
372
+ # write, unchanged.
373
+ table = kv_cache.page_table if fill_page_table is None else fill_page_table
374
+ batch_idx = user_id if fill_page_table is None else 0
375
+ ttnn.experimental.paged_fill_cache(kv_cache.k, match_cache_dtype(kv_cache.k, k), table, batch_idx=batch_idx)
376
+ ttnn.experimental.paged_fill_cache(kv_cache.v, match_cache_dtype(kv_cache.v, v), table, batch_idx=batch_idx)
377
+
378
+
379
+ def attention_decode(
380
+ x: ttnn.Tensor,
381
+ weights: AttentionWeights,
382
+ config: AttentionConfig,
383
+ cos_cache: ttnn.Tensor,
384
+ sin_cache: ttnn.Tensor,
385
+ kv_cache: KVCache,
386
+ current_pos: ttnn.Tensor,
387
+ token_index: int,
388
+ compute_kernel_config=None,
389
+ ) -> ttnn.Tensor:
390
+ """Single-token attention against the KV cache.
391
+
392
+ ``x`` is ``[1, 1, batch, hidden]`` and the return matches. ``current_pos``
393
+ is an int32 tensor of shape ``[batch]`` holding each user's write position;
394
+ ``token_index`` is the same value as a Python int, needed because the
395
+ rotary op takes a scalar rather than a tensor.
396
+ """
397
+ k_cache, v_cache, page_table = kv_cache.k, kv_cache.v, kv_cache.page_table
398
+
399
+ xqkv = ttnn.linear(
400
+ x,
401
+ weights.wqkv,
402
+ dtype=ttnn.bfloat16,
403
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
404
+ compute_kernel_config=compute_kernel_config,
405
+ )
406
+
407
+ # Blackhole: nlp_create_qkv_heads_decode's interleaved reader zeroes
408
+ # odd-indexed Q rows when the fused input sits in DRAM, due to a NoC
409
+ # DRAM-read alignment violation (tt-metal #16667). Staging through L1 takes
410
+ # a different code path and is unaffected. No-op on Wormhole.
411
+ if xqkv.memory_config().buffer_type == ttnn.BufferType.DRAM:
412
+ xqkv = ttnn.to_memory_config(xqkv, ttnn.L1_MEMORY_CONFIG)
413
+
414
+ q, k, v = ttnn.experimental.nlp_create_qkv_heads_decode(
415
+ xqkv,
416
+ num_heads=config.num_attention_heads,
417
+ num_kv_heads=config.num_key_value_heads,
418
+ memory_config=ttnn.L1_HEIGHT_SHARDED_MEMORY_CONFIG,
419
+ )
420
+ ttnn.deallocate(xqkv)
421
+
422
+ # The decode ops disagree about layout: rms_norm wants interleaved DRAM,
423
+ # while paged_update_cache requires a *sharded* update tensor
424
+ # (paged_update_cache_device_operation.cpp:255). So remember the sharded
425
+ # config the split produced and restore it once the norm and RoPE are done.
426
+ kv_sharded_mem = k.memory_config()
427
+
428
+ q = _per_head_rms_norm(ttnn.to_memory_config(q, ttnn.DRAM_MEMORY_CONFIG), weights.q_norm, config.rms_norm_eps)
429
+ k = _per_head_rms_norm(ttnn.to_memory_config(k, ttnn.DRAM_MEMORY_CONFIG), weights.k_norm, config.rms_norm_eps)
430
+
431
+ q = _apply_rope(q, cos_cache, sin_cache, token_index=token_index)
432
+ k = _apply_rope(k, cos_cache, sin_cache, token_index=token_index)
433
+
434
+ k = ttnn.to_memory_config(k, kv_sharded_mem) # v never left the sharded layout
435
+ # page_table=None is the contiguous path; the same op serves both.
436
+ # NOT cast to the cache dtype -- unlike the fill writers above.
437
+ # ``paged_update_cache`` requires a FLOAT32/BFLOAT16 update and converts into
438
+ # the cache itself; handing it a block-float input is rejected outright
439
+ # (``paged_update_cache_device_operation.cpp:296``). Measured both ways in
440
+ # ``doc/datatype_sweep/probes/kv_bfp8_diagnosis.json``. See
441
+ # :func:`match_cache_dtype`.
442
+ ttnn.experimental.paged_update_cache(k_cache, k, update_idxs_tensor=current_pos, page_table=page_table)
443
+ ttnn.experimental.paged_update_cache(v_cache, v, update_idxs_tensor=current_pos, page_table=page_table)
444
+ ttnn.deallocate(k)
445
+ ttnn.deallocate(v)
446
+
447
+ if kv_cache.is_paged:
448
+ attn = ttnn.transformer.paged_scaled_dot_product_attention_decode(
449
+ q,
450
+ k_cache,
451
+ v_cache,
452
+ page_table_tensor=page_table,
453
+ cur_pos_tensor=current_pos,
454
+ scale=config.head_dim**-0.5,
455
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
456
+ )
457
+ else:
458
+ attn = ttnn.transformer.scaled_dot_product_attention_decode(
459
+ q,
460
+ k_cache,
461
+ v_cache,
462
+ cur_pos_tensor=current_pos,
463
+ scale=config.head_dim**-0.5,
464
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
465
+ )
466
+ ttnn.deallocate(q)
467
+
468
+ attn = _concat_heads_decode(attn, config)
469
+ out = ttnn.linear(
470
+ attn,
471
+ weights.wo,
472
+ dtype=ttnn.bfloat16,
473
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
474
+ compute_kernel_config=compute_kernel_config,
475
+ )
476
+ ttnn.deallocate(attn)
477
+ return out
478
+
479
+
480
+ def attention_prefill(
481
+ x: ttnn.Tensor,
482
+ weights: AttentionWeights,
483
+ config: AttentionConfig,
484
+ cos_cache: ttnn.Tensor,
485
+ sin_cache: ttnn.Tensor,
486
+ kv_cache: KVCache | None = None,
487
+ user_id: int = 0,
488
+ compute_kernel_config=None,
489
+ sdpa_program_config=None,
490
+ activation_dtype=ttnn.bfloat16,
491
+ start_pos: int = 0,
492
+ chunk_page_table=None,
493
+ fill_page_table=None,
494
+ fill_len: int | None = None,
495
+ ) -> ttnn.Tensor:
496
+ """Causal self-attention over a full sequence. ``x``/return ``[1, 1, S, hidden]``.
497
+
498
+ When ``kv_cache`` is given, the post-RoPE K and V for the whole sequence are
499
+ written into it so a decode pass can continue from position S.
500
+
501
+ ``sdpa_program_config`` is passed straight through to the prefill SDPA op and
502
+ defaults to ``None`` -- the op default, which is what **every** caller
503
+ currently uses, including the multichip path, and what every number in this
504
+ file was measured at. It is a seam, exactly like
505
+ ``attention_decode_optimized``'s: stage 06 used it to build and measure a
506
+ length-dependent chunking worth 6.3-6.8x on this op at S >= 4096, and then
507
+ **did not adopt it** because it costs a top-1 point on ``run_teacher_forcing``
508
+ and buys nothing at the 158-token prompt that gate uses. See
509
+ ``multichip_decoder._sdpa_prefill_program_config`` for the numbers and for
510
+ what it would take to adopt.
511
+
512
+ ``activation_dtype`` is the dtype the two projections emit. It defaults to
513
+ ``ttnn.bfloat16``, which is the literal that used to be written here, so
514
+ every existing caller is unchanged; the multichip prefill layer passes
515
+ ``precision.activation_dtype`` so that field reaches prefill attention and
516
+ not only decode.
517
+
518
+ ``fill_len`` is how many of ``S`` rows are REAL tokens, and exists for
519
+ bucketed prefill: the caller pads the prompt up to a bucket so the matmuls
520
+ and SDPA see one of a handful of shapes, but the cache write must stay at
521
+ the true length. vLLM allocates ``ceil(real/block_size)`` blocks and not one
522
+ more, so writing the padded length would run off the end of the user's page
523
+ table and into another request's pages. ``None`` means "all of ``S`` is
524
+ real" -- the unbucketed path, byte-identical to before this seam existed.
525
+ """
526
+ xqkv = ttnn.linear(
527
+ x,
528
+ weights.wqkv,
529
+ dtype=activation_dtype,
530
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
531
+ compute_kernel_config=compute_kernel_config,
532
+ )
533
+
534
+ q, k, v = ttnn.experimental.nlp_create_qkv_heads(
535
+ xqkv,
536
+ num_heads=config.num_attention_heads,
537
+ num_kv_heads=config.num_key_value_heads,
538
+ transpose_k_heads=False, # SDPA wants [.., S, head_dim], not a pre-transposed K
539
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
540
+ )
541
+ ttnn.deallocate(xqkv)
542
+
543
+ # QK-norm, then RoPE -- this order is what Qwen3MoeAttention.forward does.
544
+ q = _per_head_rms_norm(q, weights.q_norm, config.rms_norm_eps)
545
+ k = _per_head_rms_norm(k, weights.k_norm, config.rms_norm_eps)
546
+
547
+ q = _apply_rope(q, cos_cache, sin_cache)
548
+ k = _apply_rope(k, cos_cache, sin_cache)
549
+
550
+ # Seed the cache with the prompt's post-RoPE K/V so decode can continue.
551
+ if kv_cache is not None:
552
+ if fill_len is not None and int(fill_len) < int(k.shape[2]):
553
+ # Bucketed prefill: drop the padding rows before the write. SDPA
554
+ # above already ran at the padded length -- which is the point, it
555
+ # is the shape-hungry op -- and causality means the padding rows
556
+ # could not have influenced any real row's output.
557
+ real = int(fill_len)
558
+ k_fill = ttnn.slice(k, [0, 0, 0, 0], [k.shape[0], k.shape[1], real, k.shape[3]])
559
+ v_fill = ttnn.slice(v, [0, 0, 0, 0], [v.shape[0], v.shape[1], real, v.shape[3]])
560
+ _fill_cache(kv_cache, k_fill, v_fill, user_id=user_id, fill_page_table=fill_page_table)
561
+ ttnn.deallocate(k_fill)
562
+ ttnn.deallocate(v_fill)
563
+ else:
564
+ _fill_cache(kv_cache, k, v, user_id=user_id, fill_page_table=fill_page_table)
565
+
566
+ # GQA is handled inside SDPA: it broadcasts the 4 KV heads across 32 Q heads.
567
+ # Default scale is head_dim ** -0.5, which is what Qwen3 uses.
568
+ if start_pos > 0:
569
+ # Split prefill: this chunk's Q attends to the WHOLE cached prefix, which
570
+ # lives in the paged cache -- so the read goes through the paged kernel
571
+ # with the user's full page-table row, not the local k/v above.
572
+ chunk = sdpa_chunk_size(start_pos)
573
+ PREFILL_ATTENTION_BRANCHES["chunked"] += 1
574
+ PREFILL_ATTENTION_BRANCHES["chunk_sizes"].append(chunk)
575
+ prog = ttnn.SDPAProgramConfig(
576
+ compute_with_storage_grid_size=q.device().compute_with_storage_grid_size(),
577
+ q_chunk_size=chunk,
578
+ k_chunk_size=chunk,
579
+ exp_approx_mode=True,
580
+ )
581
+ attn = ttnn.transformer.chunked_scaled_dot_product_attention(
582
+ q,
583
+ kv_cache.k,
584
+ kv_cache.v,
585
+ chunk_page_table,
586
+ chunk_start_idx=start_pos,
587
+ program_config=prog,
588
+ compute_kernel_config=compute_kernel_config,
589
+ )
590
+ else:
591
+ PREFILL_ATTENTION_BRANCHES["standard"] += 1
592
+ attn = ttnn.transformer.scaled_dot_product_attention(
593
+ q, k, v, is_causal=True, program_config=sdpa_program_config
594
+ )
595
+ for t in (q, k, v):
596
+ ttnn.deallocate(t)
597
+
598
+ attn = ttnn.experimental.nlp_concat_heads(attn, memory_config=ttnn.DRAM_MEMORY_CONFIG)
599
+ out = ttnn.linear(
600
+ attn,
601
+ weights.wo,
602
+ dtype=activation_dtype,
603
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
604
+ compute_kernel_config=compute_kernel_config,
605
+ )
606
+ ttnn.deallocate(attn)
607
+ return out
608
+
609
+
610
+ @dataclass(frozen=True)
611
+ class MoEConfig:
612
+ hidden_size: int
613
+ num_experts: int
614
+ num_experts_per_tok: int
615
+ moe_intermediate_size: int
616
+ norm_topk_prob: bool
617
+
618
+ @classmethod
619
+ def from_hf(cls, config) -> "MoEConfig":
620
+ assert config.hidden_act == "silu", f"expected silu, got {config.hidden_act}"
621
+ assert config.decoder_sparse_step == 1 and not config.mlp_only_layers, "every layer is expected to be MoE"
622
+ return cls(
623
+ hidden_size=config.hidden_size,
624
+ num_experts=config.num_experts,
625
+ num_experts_per_tok=config.num_experts_per_tok,
626
+ moe_intermediate_size=config.moe_intermediate_size,
627
+ norm_topk_prob=config.norm_topk_prob,
628
+ )
629
+
630
+
631
+ def router_forward(x: ttnn.Tensor, w_router: ttnn.Tensor, config: MoEConfig) -> ttnn.Tensor:
632
+ """Dense routing weights ``[1, 1, S, num_experts]``: top-k weights, zeros elsewhere.
633
+
634
+ Returning a *dense* tensor rather than (values, indices) keeps everything on
635
+ device and trace-compatible -- the experts consume it as a sparsity pattern
636
+ and as the post-hoc scaling factor.
637
+
638
+ Selection happens on the **raw logits**, not on softmax probabilities, and
639
+ the softmax is taken over the 8 survivors only. That is algebraically the
640
+ same answer HF computes, and numerically a far better one.
641
+
642
+ Why it is the same answer. Softmax is monotonic, so the top-k of the
643
+ probabilities is the top-k of the logits -- the 128-wide softmax cannot
644
+ change *who* wins. And with ``norm_topk_prob`` the shared denominator
645
+ cancels out of the renormalisation::
646
+
647
+ w_i = [exp(x_i)/Z] / sum_{j in top8}[exp(x_j)/Z]
648
+ = exp(x_i) / sum_{j in top8} exp(x_j)
649
+
650
+ so ``Z``, the sum over all 128 experts, is never needed. (This cancellation
651
+ is what makes the rewrite legal; it does **not** hold if
652
+ ``norm_topk_prob`` is False, hence the assert.)
653
+
654
+ Why it is a better answer. Measured on this checkpoint with a 128-token
655
+ activation sample:
656
+
657
+ ttnn.topk on fp32 logits ............ 0/128 tokens misrouted (exact)
658
+ ttnn.topk on fp32 softmax probs ..... 0/128 (topk itself is fine)
659
+ ttnn.softmax fp32 vs torch fp32 ..... max abs error 3.3e-4
660
+
661
+ The last line is the problem. Softmax over 128 experts leaves the 8th-place
662
+ probability near 0.008 with only ~1.4e-5 separating it from the 9th, so a
663
+ 3.3e-4 error decides the cut by luck -- and ``ttnn.softmax`` carries that
664
+ error even when handed an fp32 tensor, so no dtype change rescues it.
665
+ Routing through the full softmax cost 34/128 misrouted tokens (83/128 with
666
+ a bf16 softmax, dense PCC 0.88). Comparing logits sidesteps it: the gaps
667
+ there are ordinary-sized.
668
+
669
+ The projection stays bf16 -- it is accurate to PCC 0.9999973 and costs only
670
+ 3/128 on its own -- but it accumulates into fp32 so the comparison is clean.
671
+ Weights are cast to bf16 for the scatter, which has no fp32 tiled support;
672
+ that is harmless, since representing a weight to 0.4% is not the same
673
+ problem as deciding which weights exist.
674
+ """
675
+ assert config.norm_topk_prob, (
676
+ "router selects on raw logits, which relies on the softmax denominator "
677
+ "cancelling during top-k renormalisation; that only holds when "
678
+ "norm_topk_prob is True"
679
+ )
680
+
681
+ logits = ttnn.linear(x, w_router, dtype=ttnn.float32, memory_config=ttnn.DRAM_MEMORY_CONFIG)
682
+
683
+ top_logits, top_indices = ttnn.topk(logits, k=config.num_experts_per_tok, dim=-1)
684
+
685
+ # softmax over the 8 survivors, written out rather than calling
686
+ # ttnn.softmax so the reduction is the ttnn.sum already exercised by
687
+ # test_router_weights_sum_to_one. Subtracting the max is for exp() range
688
+ # only; any shared shift cancels in the division.
689
+ top_max = ttnn.max(top_logits, dim=-1, keepdim=True)
690
+ exp_logits = ttnn.exp(ttnn.sub(top_logits, top_max))
691
+ total = ttnn.sum(exp_logits, dim=-1, keepdim=True)
692
+ top_values = ttnn.div(exp_logits, total)
693
+
694
+ zeros = ttnn.typecast(ttnn.zeros_like(logits), ttnn.bfloat16)
695
+ dense = ttnn.scatter(
696
+ zeros,
697
+ dim=-1,
698
+ index=top_indices,
699
+ src=ttnn.typecast(top_values, ttnn.bfloat16),
700
+ )
701
+ for t in (logits, top_logits, top_indices, top_max, exp_logits, total, top_values):
702
+ ttnn.deallocate(t)
703
+ return dense
704
+
705
+
706
+ def upload_router_weight(router: torch.Tensor, device) -> ttnn.Tensor:
707
+ """``[num_experts, hidden]`` checkpoint tensor -> ``[1, 1, hidden, num_experts]``."""
708
+ return ttnn.from_torch(
709
+ router.T.contiguous().reshape(1, 1, router.shape[1], router.shape[0]).float(),
710
+ dtype=ttnn.bfloat16,
711
+ layout=ttnn.TILE_LAYOUT,
712
+ device=device,
713
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
714
+ )
715
+
716
+
717
+ @dataclass
718
+ class ExpertWeights:
719
+ gate_proj: ttnn.Tensor # [1, num_experts, hidden, intermediate]
720
+ up_proj: ttnn.Tensor # [1, num_experts, hidden, intermediate]
721
+ down_proj: ttnn.Tensor # [1, num_experts, intermediate, hidden]
722
+
723
+
724
+ def upload_expert_weights(torch_weights: dict[str, torch.Tensor], device, config: MoEConfig) -> ExpertWeights:
725
+ """Unfuse and transpose the batched expert tensors into sparse_matmul layout.
726
+
727
+ ``weight_mapping`` keeps the checkpoint's fused ``[E, 2I, H]`` form because
728
+ that is what the HF module holds and what the reference test validates
729
+ against. ``sparse_matmul`` wants them separate and transposed, so the split
730
+ happens here rather than at conversion time -- the fusion order stays
731
+ guarded by ``test_moe_matches_unfused_reimplementation``.
732
+ """
733
+ fused = torch_weights["experts_gate_up"] # [E, 2I, H], gate first
734
+ inter = config.moe_intermediate_size
735
+
736
+ # bf16 while correctness is being established, so a PCC miss means a bug
737
+ # rather than quantisation. bfloat8_b is the intended production dtype
738
+ # (gemma4 and gpt_oss both ship it) and is a later, measurable step -- one
739
+ # layer's experts are ~604M params, 1.2GB at bf16, which fits comfortably.
740
+ def up(t: torch.Tensor) -> ttnn.Tensor:
741
+ return ttnn.from_torch(
742
+ t.contiguous().float(),
743
+ dtype=ttnn.bfloat16,
744
+ layout=ttnn.TILE_LAYOUT,
745
+ device=device,
746
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
747
+ )
748
+
749
+ return ExpertWeights(
750
+ gate_proj=up(fused[:, :inter, :].transpose(-2, -1).unsqueeze(0)),
751
+ up_proj=up(fused[:, inter:, :].transpose(-2, -1).unsqueeze(0)),
752
+ # checkpoint down_proj is [E, H, I]; sparse_matmul wants [1, E, I, H]
753
+ down_proj=up(torch_weights["experts_down"].transpose(-2, -1).unsqueeze(0)),
754
+ )
755
+
756
+
757
+ def build_expert_sparsity(device, num_experts: int) -> ttnn.Tensor:
758
+ """All-ones sparsity mask ``[1, 1, 1, E]``, ROW_MAJOR.
759
+
760
+ Every expert is computed for every token group; the routing weights zero
761
+ out the inactive ones after the down projection. That is the established
762
+ gpt_oss/gemma4 pattern -- it trades compute for a static, trace-friendly
763
+ shape rather than gathering per-token expert assignments.
764
+ """
765
+ return ttnn.from_torch(
766
+ torch.ones(1, 1, 1, num_experts, dtype=torch.bfloat16),
767
+ layout=ttnn.ROW_MAJOR_LAYOUT,
768
+ dtype=ttnn.bfloat16,
769
+ device=device,
770
+ )
771
+
772
+
773
+ def _expert_compute_kernel_config(device):
774
+ """HiFi4, and ``fp32_dest_acc_en`` deliberately OFF.
775
+
776
+ The functional path holds expert weights in **bf16** (see
777
+ ``upload_expert_weights``), so the accumulation over hidden/intermediate has
778
+ 16 bits of mantissa to preserve and LoFi -- the matmul default, which keeps
779
+ the top 5 -- is not enough. Enabling fp32 dest accumulation looks like the
780
+ natural next lever but must not be used here: it halves the matmul dest from
781
+ 8 tiles to 4, which corrupts expert output on Blackhole (tt-metal #49068,
782
+ hit on BH-QB-2). HiFi4 alone provides the accuracy.
783
+
784
+ The optimized decoder quantises the experts to bfloat4_b, where there is no
785
+ longer any low-order mantissa for HiFi4 to resolve, and measured LoFi as
786
+ both faster and marginally more accurate -- so it overrides this with its
787
+ own config rather than reusing it. See
788
+ ``optimized_decoder.EXPERT_MATH_FIDELITY``.
789
+ """
790
+ return ttnn.init_device_compute_kernel_config(
791
+ device.arch(),
792
+ math_fidelity=ttnn.MathFidelity.HiFi4,
793
+ math_approx_mode=False,
794
+ fp32_dest_acc_en=False,
795
+ packer_l1_acc=False,
796
+ )
797
+
798
+
799
+ def _sparse_matmul_config(m: int, n: int, in0_block_w: int = 1):
800
+ """Spread the N dimension over the largest usable slice of an 8x8 grid."""
801
+ n_tiles = math.ceil(n / 32)
802
+ best_cores, best_cx, best_cy = 1, 1, 1
803
+ for num_cores in range(1, min(65, n_tiles + 1)):
804
+ if n_tiles % num_cores:
805
+ continue
806
+ for cy in range(1, 9):
807
+ if num_cores % cy == 0:
808
+ cx = num_cores // cy
809
+ if cx <= 8 and num_cores > best_cores:
810
+ best_cores, best_cx, best_cy = num_cores, cx, cy
811
+ break
812
+ per_core_n = n_tiles // best_cores
813
+ return ttnn.MatmulMultiCoreReuseMultiCast1DProgramConfig(
814
+ compute_with_storage_grid_size=ttnn.CoreCoord(best_cx, best_cy),
815
+ in0_block_w=in0_block_w,
816
+ out_subblock_h=1,
817
+ out_subblock_w=1,
818
+ out_block_h=1,
819
+ out_block_w=per_core_n,
820
+ per_core_M=max(32, m) // 32,
821
+ per_core_N=per_core_n,
822
+ fuse_batch=False,
823
+ fused_activation=None,
824
+ mcast_in0=True,
825
+ )
826
+
827
+
828
+ # sparse_matmul folds the group dimension (chunk_len / 32) into M, which grows
829
+ # num_blocks_y and can overflow the core grid. Chunking at exactly one tile
830
+ # keeps group_size == 1 so all blocking comes from N.
831
+ EXPERT_CHUNK_SIZE = 32
832
+
833
+
834
+ def _experts_chunk(
835
+ hidden: ttnn.Tensor,
836
+ routing: ttnn.Tensor,
837
+ weights: ExpertWeights,
838
+ config: MoEConfig,
839
+ sparsity_base: ttnn.Tensor,
840
+ ) -> ttnn.Tensor:
841
+ """One 32-token chunk through all experts. ``hidden`` ``[1, 1, 32, H]``."""
842
+ chunk_len = hidden.shape[2]
843
+ n_experts = config.num_experts
844
+ hidden_size = config.hidden_size
845
+ group_size = chunk_len // EXPERT_CHUNK_SIZE
846
+
847
+ device = hidden.device()
848
+ compute_config = _expert_compute_kernel_config(device)
849
+ output_tile = ttnn.Tile([32, 32])
850
+ gate_up_config = _sparse_matmul_config(EXPERT_CHUNK_SIZE, config.moe_intermediate_size)
851
+ down_config = _sparse_matmul_config(EXPERT_CHUNK_SIZE, hidden_size)
852
+
853
+ hidden_grouped = ttnn.reshape(hidden, (1, group_size, EXPERT_CHUNK_SIZE, hidden_size))
854
+ sparsity = ttnn.repeat(sparsity_base, (1, 1, group_size, 1))
855
+ nnz = n_experts * group_size
856
+
857
+ def project(weight):
858
+ out = ttnn.sparse_matmul(
859
+ hidden_grouped,
860
+ weight,
861
+ sparsity=sparsity,
862
+ nnz=nnz,
863
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
864
+ output_tile=output_tile,
865
+ program_config=gate_up_config,
866
+ compute_kernel_config=compute_config,
867
+ dtype=ttnn.bfloat16,
868
+ )
869
+ inter = out.shape[-1]
870
+ return ttnn.reshape(ttnn.transpose(out, 1, 3), (1, n_experts, chunk_len, inter)), inter
871
+
872
+ gate, inter = project(weights.gate_proj)
873
+ up, _ = project(weights.up_proj)
874
+ ttnn.deallocate(hidden_grouped)
875
+
876
+ # SwiGLU -- Qwen3 is hidden_act="silu". gemma4, which this follows, uses
877
+ # GeGLU; swapping the activation runs fine and returns wrong numbers.
878
+ down_input = ttnn.reshape(ttnn.mul(ttnn.silu(gate), up), (1, n_experts, chunk_len, inter))
879
+ ttnn.deallocate(gate)
880
+ ttnn.deallocate(up)
881
+
882
+ down = ttnn.sparse_matmul(
883
+ down_input,
884
+ weights.down_proj,
885
+ sparsity=sparsity_base,
886
+ nnz=n_experts,
887
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
888
+ output_tile=output_tile,
889
+ program_config=down_config,
890
+ is_input_a_sparse=True,
891
+ compute_kernel_config=compute_config,
892
+ dtype=ttnn.bfloat16,
893
+ )
894
+ ttnn.deallocate(down_input)
895
+
896
+ # Scale each expert's contribution by its routing weight (zero for the 120
897
+ # experts this token did not select), then sum over the expert dimension.
898
+ states = ttnn.reshape(down, (1, n_experts, chunk_len, hidden_size))
899
+ states = ttnn.mul(states, ttnn.permute(routing, (0, 3, 2, 1))) # [1, E, S, 1]
900
+ states = ttnn.unsqueeze_to_4D(ttnn.experimental.fast_reduce_nc(states, dims=[1]))
901
+ return ttnn.reshape(states, (1, 1, chunk_len, hidden_size))
902
+
903
+
904
+ def moe_prefill(
905
+ x: ttnn.Tensor,
906
+ routing: ttnn.Tensor,
907
+ weights: ExpertWeights,
908
+ config: MoEConfig,
909
+ sparsity_base: ttnn.Tensor,
910
+ ) -> ttnn.Tensor:
911
+ """Full MoE expert pass over a sequence. ``x`` ``[1, 1, S, H]``, any S.
912
+
913
+ ``sparse_matmul`` works a tile at a time, so a sequence that is not a
914
+ multiple of 32 is zero-padded up to one and the extra rows are dropped from
915
+ the result. Padded rows carry an all-zero routing vector, so every expert's
916
+ contribution to them is scaled by zero -- they cost a little compute and
917
+ cannot perturb the real tokens.
918
+ """
919
+ seq_len = x.shape[2]
920
+ padded_len = math.ceil(seq_len / EXPERT_CHUNK_SIZE) * EXPERT_CHUNK_SIZE
921
+
922
+ if padded_len != seq_len:
923
+ pad = [(0, 0), (0, 0), (0, padded_len - seq_len), (0, 0)]
924
+ x = ttnn.pad(x, pad, value=0.0)
925
+ routing = ttnn.pad(routing, pad, value=0.0)
926
+
927
+ outputs = []
928
+ for start in range(0, padded_len, EXPERT_CHUNK_SIZE):
929
+ end = start + EXPERT_CHUNK_SIZE
930
+ outputs.append(
931
+ _experts_chunk(
932
+ ttnn.slice(x, [0, 0, start, 0], [1, 1, end, config.hidden_size]),
933
+ ttnn.slice(routing, [0, 0, start, 0], [1, 1, end, config.num_experts]),
934
+ weights,
935
+ config,
936
+ sparsity_base,
937
+ )
938
+ )
939
+ out = outputs[0] if len(outputs) == 1 else ttnn.concat(outputs, dim=2)
940
+ if padded_len != seq_len:
941
+ out = ttnn.slice(out, [0, 0, 0, 0], [1, 1, seq_len, config.hidden_size])
942
+ return out
943
+
944
+
945
+ def moe_decode(
946
+ x: ttnn.Tensor,
947
+ routing: ttnn.Tensor,
948
+ weights: ExpertWeights,
949
+ config: MoEConfig,
950
+ ) -> ttnn.Tensor:
951
+ """MoE for a single token per user. ``x`` ``[1, 1, batch, H]``.
952
+
953
+ Unlike the prefill path this uses **real** sparsity: the dense routing
954
+ tensor is handed to ``sparse_matmul`` directly, so only the selected experts
955
+ are computed rather than all 128. Prefill cannot do that -- across 32 tokens
956
+ the union of selected experts approaches the full set, so an all-ones mask
957
+ with a post-hoc mask-out is both simpler and no more expensive. In decode
958
+ the working set really is ``batch * top_k``, and skipping the other 120
959
+ experts per token is the entire reason decode is affordable.
960
+ """
961
+ batch = x.shape[2]
962
+ n_experts = config.num_experts
963
+ hidden_size = config.hidden_size
964
+ # True non-zero count of the sparsity tensor. gemma4 passes top_k because
965
+ # its decode is single-user; with several users each contributing top_k
966
+ # entries the kernel needs the full count or it under-sizes the work.
967
+ nnz = config.num_experts_per_tok * batch
968
+
969
+ sparsity = ttnn.to_layout(routing, ttnn.ROW_MAJOR_LAYOUT)
970
+ output_tile = ttnn.Tile([32, 32])
971
+ compute_config = _expert_compute_kernel_config(x.device())
972
+ gate_up_config = _sparse_matmul_config(batch, config.moe_intermediate_size)
973
+ down_config = _sparse_matmul_config(batch, hidden_size)
974
+
975
+ def project(weight):
976
+ out = ttnn.sparse_matmul(
977
+ x,
978
+ weight,
979
+ sparsity=sparsity,
980
+ nnz=nnz,
981
+ memory_config=ttnn.L1_MEMORY_CONFIG,
982
+ output_tile=output_tile,
983
+ program_config=gate_up_config,
984
+ compute_kernel_config=compute_config,
985
+ dtype=ttnn.bfloat16,
986
+ )
987
+ inter = out.shape[-1]
988
+ out = ttnn.transpose(ttnn.reshape(out, (batch, n_experts, 1, inter)), 1, 2)
989
+ return ttnn.reshape(out, (batch, n_experts, inter)), inter
990
+
991
+ gate, inter = project(weights.gate_proj)
992
+ up, _ = project(weights.up_proj)
993
+
994
+ down_input = ttnn.mul(ttnn.silu(gate), up) # SwiGLU, not GeGLU
995
+ ttnn.deallocate(gate)
996
+ ttnn.deallocate(up)
997
+ down_input = ttnn.reshape(ttnn.transpose(down_input, 1, 0), (1, n_experts, batch, inter))
998
+
999
+ down = ttnn.sparse_matmul(
1000
+ down_input,
1001
+ weights.down_proj,
1002
+ sparsity=sparsity,
1003
+ nnz=nnz,
1004
+ memory_config=ttnn.L1_MEMORY_CONFIG,
1005
+ output_tile=output_tile,
1006
+ program_config=down_config,
1007
+ is_input_a_sparse=True,
1008
+ compute_kernel_config=compute_config,
1009
+ dtype=ttnn.bfloat16,
1010
+ )
1011
+ ttnn.deallocate(down_input)
1012
+
1013
+ states = ttnn.reshape(ttnn.permute(down, (0, 2, 1, 3)), (batch, n_experts, hidden_size))
1014
+ states = ttnn.mul(states, ttnn.reshape(routing, (batch, n_experts, 1)))
1015
+ states = ttnn.unsqueeze_to_4D(ttnn.sum(states, dim=1))
1016
+ return ttnn.reshape(states, (1, 1, batch, hidden_size), (1, 1, max(32, batch), hidden_size))
1017
+
1018
+
1019
+ @dataclass(frozen=True)
1020
+ class DecoderLayerConfig:
1021
+ attention: AttentionConfig
1022
+ moe: MoEConfig
1023
+ rms_norm_eps: float
1024
+
1025
+ @classmethod
1026
+ def from_hf(cls, config) -> "DecoderLayerConfig":
1027
+ return cls(
1028
+ attention=AttentionConfig.from_hf(config),
1029
+ moe=MoEConfig.from_hf(config),
1030
+ rms_norm_eps=config.rms_norm_eps,
1031
+ )
1032
+
1033
+
1034
+ @dataclass
1035
+ class DecoderLayerWeights:
1036
+ input_layernorm: ttnn.Tensor
1037
+ post_attention_layernorm: ttnn.Tensor
1038
+ attention: AttentionWeights
1039
+ router: ttnn.Tensor
1040
+ experts: ExpertWeights
1041
+
1042
+
1043
+ def decoder_layer_prefill(
1044
+ x: ttnn.Tensor,
1045
+ weights: DecoderLayerWeights,
1046
+ config: DecoderLayerConfig,
1047
+ cos_cache: ttnn.Tensor,
1048
+ sin_cache: ttnn.Tensor,
1049
+ sparsity: ttnn.Tensor,
1050
+ kv_cache: KVCache | None = None,
1051
+ user_id: int = 0,
1052
+ ) -> ttnn.Tensor:
1053
+ """One full decoder layer. ``x`` / return ``[1, 1, S, hidden]``.
1054
+
1055
+ Passing ``kv_cache`` seeds it with the prompt's K/V so decode can continue
1056
+ from position S.
1057
+
1058
+ Pre-norm, matching ``Qwen3MoeDecoderLayer.forward``::
1059
+
1060
+ h = x + attn(norm1(x))
1061
+ out = h + moe(norm2(h))
1062
+
1063
+ Note the MoE consumes a single normed tensor for *both* the router and the
1064
+ experts. Some ports (gemma4) thread separate router/expert inputs because
1065
+ their router applies its own normalisation; Qwen3's does not, and feeding
1066
+ the router the un-normed residual instead would change every routing
1067
+ decision.
1068
+ """
1069
+ eps = config.rms_norm_eps
1070
+
1071
+ normed = ttnn.rms_norm(x, weight=weights.input_layernorm, epsilon=eps)
1072
+ attn_out = attention_prefill(normed, weights.attention, config.attention, cos_cache, sin_cache, kv_cache, user_id)
1073
+ ttnn.deallocate(normed)
1074
+ hidden = ttnn.add(x, attn_out)
1075
+ ttnn.deallocate(attn_out)
1076
+
1077
+ normed = ttnn.rms_norm(hidden, weight=weights.post_attention_layernorm, epsilon=eps)
1078
+ routing = router_forward(normed, weights.router, config.moe)
1079
+ moe_out = moe_prefill(normed, routing, weights.experts, config.moe, sparsity)
1080
+ ttnn.deallocate(normed)
1081
+ ttnn.deallocate(routing)
1082
+
1083
+ out = ttnn.add(hidden, moe_out)
1084
+ ttnn.deallocate(hidden)
1085
+ ttnn.deallocate(moe_out)
1086
+ return out
1087
+
1088
+
1089
+ def decoder_layer_decode(
1090
+ x: ttnn.Tensor,
1091
+ weights: DecoderLayerWeights,
1092
+ config: DecoderLayerConfig,
1093
+ cos_cache: ttnn.Tensor,
1094
+ sin_cache: ttnn.Tensor,
1095
+ kv_cache: KVCache,
1096
+ current_pos: ttnn.Tensor,
1097
+ token_index: int,
1098
+ ) -> ttnn.Tensor:
1099
+ """One decoder layer, single token per user. ``x`` / return ``[1, 1, batch, hidden]``.
1100
+
1101
+ Same graph as ``decoder_layer_prefill`` -- only the attention and expert
1102
+ kernels differ, because decode attends against the KV cache and can exploit
1103
+ real routing sparsity.
1104
+ """
1105
+ eps = config.rms_norm_eps
1106
+
1107
+ normed = ttnn.rms_norm(x, weight=weights.input_layernorm, epsilon=eps)
1108
+ attn_out = attention_decode(
1109
+ normed, weights.attention, config.attention, cos_cache, sin_cache, kv_cache, current_pos, token_index
1110
+ )
1111
+ ttnn.deallocate(normed)
1112
+ hidden = ttnn.add(x, attn_out)
1113
+ ttnn.deallocate(attn_out)
1114
+
1115
+ normed = ttnn.rms_norm(hidden, weight=weights.post_attention_layernorm, epsilon=eps)
1116
+ routing = router_forward(normed, weights.router, config.moe)
1117
+ moe_out = moe_decode(normed, routing, weights.experts, config.moe)
1118
+ ttnn.deallocate(normed)
1119
+ ttnn.deallocate(routing)
1120
+
1121
+ out = ttnn.add(hidden, moe_out)
1122
+ ttnn.deallocate(hidden)
1123
+ ttnn.deallocate(moe_out)
1124
+ return out
1125
+
1126
+
1127
+ def upload_layer_weights(torch_weights: dict[str, torch.Tensor], device, config: DecoderLayerConfig):
1128
+ """Everything one decoder layer needs, from ``convert_layer_weights`` output."""
1129
+
1130
+ def norm(t: torch.Tensor) -> ttnn.Tensor:
1131
+ return ttnn.from_torch(
1132
+ t.reshape(1, 1, 1, -1).float(),
1133
+ dtype=ttnn.bfloat16,
1134
+ layout=ttnn.TILE_LAYOUT,
1135
+ device=device,
1136
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1137
+ )
1138
+
1139
+ return DecoderLayerWeights(
1140
+ input_layernorm=norm(torch_weights["input_layernorm"]),
1141
+ post_attention_layernorm=norm(torch_weights["post_attention_layernorm"]),
1142
+ attention=upload_attention_weights(torch_weights, device),
1143
+ router=upload_router_weight(torch_weights["router"], device),
1144
+ experts=upload_expert_weights(torch_weights, device, config.moe),
1145
+ )
1146
+
1147
+
1148
+ def upload_attention_weights(torch_weights: dict[str, torch.Tensor], device) -> AttentionWeights:
1149
+ """Move the host-side tensors from ``weight_mapping`` onto the device."""
1150
+
1151
+ def up(t: torch.Tensor, pad_to_4d: bool = False) -> ttnn.Tensor:
1152
+ if pad_to_4d:
1153
+ t = t.reshape(1, 1, 1, -1)
1154
+ while t.dim() < 4:
1155
+ t = t.unsqueeze(0)
1156
+ return ttnn.from_torch(
1157
+ t.float(),
1158
+ dtype=ttnn.bfloat16,
1159
+ layout=ttnn.TILE_LAYOUT,
1160
+ device=device,
1161
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1162
+ )
1163
+
1164
+ return AttentionWeights(
1165
+ wqkv=up(torch_weights["wqkv"]),
1166
+ wo=up(torch_weights["wo"]),
1167
+ q_norm=up(torch_weights["q_norm"], pad_to_4d=True),
1168
+ k_norm=up(torch_weights["k_norm"], pad_to_4d=True),
1169
+ )
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator.py ADDED
@@ -0,0 +1,1637 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Metal-readiness generator for the 4-die Qwen3-Coder-30B-A3B-Instruct model.
5
+
6
+ Two API levels, as the readiness contract requires:
7
+
8
+ * **low level** -- ``prefill_forward`` / ``decode_forward``. The caller owns the
9
+ KV cache, the page table, the per-user prompt lengths and the per-user decode
10
+ positions, and threads them through each call. Mixed-length prompts, fixed
11
+ request slots and inactive rows (position ``-1``) are all expressible here.
12
+ This is the surface a serving adapter drives.
13
+ * **high level** -- ``generate``. Owns the cache and page table and loops
14
+ deterministically over the low-level calls.
15
+
16
+ The measured token-out path is **entirely on device**: the model trace produces
17
+ sampler-ready per-die logits, a second trace runs the split sampler, the sampled
18
+ token is written straight into the persistent decode token input through
19
+ ``tt_out_tok``, and the trace advances both position tensors itself with
20
+ ``ttnn.plus_one``. Between two steady-state tokens the host does exactly one
21
+ thing -- replay two traces -- plus whatever readback the caller asked for.
22
+ ``sampling_mode="host"`` is the explicit compatibility mode for tests that need
23
+ host sampling and is never used to produce a performance number.
24
+ """
25
+
26
+ from __future__ import annotations
27
+
28
+ import bisect
29
+ import contextlib
30
+ import math
31
+ import os
32
+ from pathlib import Path
33
+ from typing import Any, Optional, Sequence
34
+
35
+ import torch
36
+ from transformers import AutoTokenizer
37
+
38
+ import ttnn
39
+ from models.common.readiness_check.contract import Generator, NextInputFn
40
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.model import (
41
+ HF_MODEL_ID,
42
+ HF_REVISION,
43
+ MAX_CONTEXT,
44
+ NUM_LAYERS,
45
+ Qwen3CoderModel,
46
+ )
47
+
48
+ #: ``ttnn.sampling``/``nlp_create_qkv_heads_decode`` both work in 32-slot units.
49
+ SAMPLING_SLOTS = 32
50
+
51
+
52
+ #: Prefill bucket ladders, keyed by the value of ``QWEN3_PREFILL_BUCKETS``.
53
+ #:
54
+ #: A prefill program is compiled per *exact* sequence length: the model
55
+ #: deliberately runs each user at its own logical length so nothing needs a
56
+ #: mask (see :meth:`Generator.prefill_forward`). That is the right trade for
57
+ #: throughput and the wrong one for first-touch latency, because the shape
58
+ #: space is then 1..max_cache_len and a warmup cannot enumerate it -- every
59
+ #: new prompt length pays a fresh ~11 s compile.
60
+ #:
61
+ #: Bucketing collapses that space: the prompt is zero-padded up to the next
62
+ #: rung, the extra rows are dropped from the logits (they are past the selected
63
+ #: row) and from the cache write (``fill_len``), and causality means they could
64
+ #: not have influenced any real row. This is exactly what
65
+ #: ``models/tt_transformers`` does with ``get_padded_prefill_len`` -- powers of
66
+ #: two, warmed by ``warmup_model_prefill`` -- and the ladder below is that idea
67
+ #: with a rung added between each power so the padding tax is halved.
68
+ #:
69
+ #: The cost is wasted compute on the padding: at the measured 0.96 ms/token, a
70
+ #: worst-case ``pow2`` prompt pays ~100% of its prefill again, ``pow2_half``
71
+ #: ~50%, ``1k`` at most ~1 s. The benefit is that the whole ladder is finite
72
+ #: and therefore warmable.
73
+ DEFAULT_PREFILL_BUCKETS = "proportional"
74
+
75
+ #: The ``proportional`` ladder's gap, as a fraction of the rung it follows --
76
+ #: and therefore the worst-case padding tax at any length. An eighth was chosen
77
+ #: to sit well under the ~24% measured on the ``pow2_half`` ladder while
78
+ #: keeping the rung count near 50, which is small enough that the persistent
79
+ #: kernel cache saturates over a real session.
80
+ PROPORTIONAL_LADDER_STEP = 0.125
81
+
82
+
83
+ def prefill_bucket_ladder(name: str, max_len: int) -> tuple[int, ...]:
84
+ """Every prefill length the model will run under ladder ``name``.
85
+
86
+ Returned ascending, and the last rung is always ``max_len`` so that every
87
+ admissible prompt maps onto some rung. ``exact`` returns ``()`` -- the
88
+ unbucketed behaviour this model shipped with.
89
+ """
90
+ max_len = int(max_len)
91
+ if max_len < 1:
92
+ raise ValueError("max_len must be positive")
93
+ if name == "exact":
94
+ return ()
95
+ # Below 128 the projections are tile-bound anyway, so one small rung buys
96
+ # the whole short-prompt range; 512 keeps a 200-token prompt off the 1024
97
+ # rung, which is the common case for a chat turn. ``proportional`` derives
98
+ # its own bottom rungs and starts from 128 alone.
99
+ lengths = [128] if name == "proportional" else [128, 512]
100
+ if name == "pow2":
101
+ step = 1024
102
+ while step < max_len:
103
+ lengths.append(step)
104
+ step *= 2
105
+ elif name == "pow2_half":
106
+ step = 1024
107
+ while step < max_len:
108
+ lengths.append(step)
109
+ lengths.append(step + step // 2)
110
+ step *= 2
111
+ elif name == "1k":
112
+ step = 1024
113
+ while step < max_len:
114
+ lengths.append(step)
115
+ step += 1024
116
+ elif name == "proportional":
117
+ # Gaps that grow with the rung, so the padding tax is bounded as a
118
+ # FRACTION of the prefill rather than as a token count. That is the
119
+ # invariant a user actually feels, and neither fixed-step nor doubling
120
+ # ladders have it: ``1k`` wastes 0.9 s on a 200-token prompt whose real
121
+ # prefill is 0.18 s, while ``pow2`` charges a prompt just past 131072
122
+ # about 111 s of padding. This caps the overhead at
123
+ # ``PROPORTIONAL_LADDER_STEP`` of the prefill at *every* length, for 52
124
+ # rungs -- few enough that the persistent kernel cache saturates over a
125
+ # session, and the low rungs (the ones a warmup can afford) are the
126
+ # ones a coding client spends most of its prompts on.
127
+ #
128
+ # The 128 floor keeps the count finite near the bottom, where an eighth
129
+ # of the length is a handful of tokens; it also means the shortest
130
+ # prompts pay at most 127 padded tokens, ~0.11 s.
131
+ step = 128
132
+ while step < max_len:
133
+ lengths.append(step)
134
+ grow = math.ceil(step * PROPORTIONAL_LADDER_STEP / 128) * 128
135
+ step += max(128, grow)
136
+ else:
137
+ raise ValueError("QWEN3_PREFILL_BUCKETS must be exact, proportional, pow2, pow2_half or 1k; " f"got {name!r}")
138
+ lengths.append(max_len)
139
+ return tuple(sorted({v for v in lengths if v <= max_len}))
140
+
141
+
142
+ def _first_device_to_torch(tensor) -> torch.Tensor:
143
+ shards = ttnn.get_device_tensors(tensor)
144
+ return ttnn.to_torch(shards[0] if shards else tensor)
145
+
146
+
147
+ class Qwen3CoderGenerator(Generator):
148
+ """Caller-owned cache/page-table state plus traced on-device token feedback."""
149
+
150
+ def __init__(self, model: Qwen3CoderModel, tokenizer):
151
+ self.model = model
152
+ self.mesh_device = model.mesh_device
153
+ self.tokenizer = tokenizer
154
+ self.batch = model.max_batch_size
155
+ self.page_block_size = model.page_block_size
156
+ self.pages_per_user = math.ceil(model.max_cache_len / self.page_block_size)
157
+ self.num_blocks = self.batch * self.pages_per_user
158
+
159
+ self._kv_cache: list | None = None
160
+ self._trace_model_id = None
161
+ #: Resolved on first use from ``QWEN3_PREFILL_BUCKETS``.
162
+ self._prefill_buckets: tuple[int, ...] | None = None
163
+ self._trace_sampling_id = None
164
+ self._trace_inputs = None
165
+ self._trace_logits = None
166
+ self._trace_sampled = None
167
+ self._trace_kv_cache = None
168
+ self._trace_page_table_snapshot: torch.Tensor | None = None
169
+ self._trace_active_batch = None
170
+ #: Width of the *captured* decode graph. Equal to ``self.batch`` unless
171
+ #: the caller asked for a narrower one via ``decode_forward(graph_width=)``.
172
+ self._trace_graph_width = None
173
+ #: Highest rotary position the *next* trace replay will gather at. The
174
+ #: trace advances ``rotary_position`` on device with ``ttnn.plus_one``
175
+ #: and nothing on device clamps it, so this host-side mirror is the only
176
+ #: thing that can tell a replay it is about to index past the cos/sin
177
+ #: table. See ``decode_forward``.
178
+ self._trace_rotary_position: int | None = None
179
+ self._decode_warm_key = None
180
+ #: Decode graph keys whose programs are already in the program cache.
181
+ #: Unlike ``_decode_warm_key`` this **survives a trace release**: the
182
+ #: eager warm pass exists to get every program compiled before capture,
183
+ #: and a program stays compiled after the trace that used it is freed.
184
+ #: Serving releases and re-captures the decode traces on every prefill
185
+ #: (a new request is admitted while other slots decode), so without this
186
+ #: each admission would pay a full eager decode forward it does not need.
187
+ #: The key is ``_decode_graph_key``, which includes ``rope_cache_len``:
188
+ #: growing the rotary tables changes the graph's shapes, so those
189
+ #: programs are *not* already compiled and the warm pass must run.
190
+ self._decode_compiled_keys: set = set()
191
+ self._sampling_params = None
192
+ self._sampling_snapshot = None
193
+ self._sampling_stochastic = False
194
+
195
+ #: Sampling-penalty state. ``_penalty_mode`` is a *graph* property (see
196
+ #: ``_WatcherCleanSampling1D``'s penalty section): 0 means the penalty ops
197
+ #: are not in the captured decode trace at all, so an unpenalised request
198
+ #: pays nothing. ``_penalty_host`` are persistent full-vocabulary staging
199
+ #: buffers; ``_penalty_prev_*`` remember which columns each row last wrote
200
+ #: so a step resets only those instead of the whole 151936-wide row.
201
+ self._penalty_mode = 0
202
+ self._penalty_host = None
203
+ self._penalty_local_vocab = None
204
+ self._penalty_prev_add: list = []
205
+ self._penalty_prev_rep: list = []
206
+
207
+ #: Steady-state host-work counters. Everything except ``replays`` and
208
+ #: ``caller_token_readbacks`` must stay flat while tokens are produced.
209
+ self.trace_stats = {
210
+ "captures": 0,
211
+ "replays": 0,
212
+ "releases": 0,
213
+ "decode_warmups": 0,
214
+ "token_host_copies": 0,
215
+ "token_device_copies": 0,
216
+ "position_host_copies": 0,
217
+ "rotary_position_host_copies": 0,
218
+ "page_table_host_copies": 0,
219
+ "sampling_param_host_copies": 0,
220
+ "penalty_host_copies": 0,
221
+ "caller_token_readbacks": 0,
222
+ "explicit_synchronizations": 0,
223
+ "resets": 0,
224
+ }
225
+ self._allocate_persistent_inputs()
226
+
227
+ # -- persistent device state ---------------------------------------------
228
+
229
+ def _replicated_host_tensor(self, host: torch.Tensor, *, dtype):
230
+ return ttnn.from_torch(
231
+ host.contiguous(),
232
+ device=None,
233
+ dtype=dtype,
234
+ layout=ttnn.ROW_MAJOR_LAYOUT,
235
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
236
+ )
237
+
238
+ def _replicated_device_tensor(self, host: torch.Tensor, *, dtype):
239
+ return ttnn.from_torch(
240
+ host.contiguous(),
241
+ device=self.mesh_device,
242
+ dtype=dtype,
243
+ layout=ttnn.ROW_MAJOR_LAYOUT,
244
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
245
+ )
246
+
247
+ def _copy_host(self, host: torch.Tensor, device, *, dtype) -> None:
248
+ ttnn.copy_host_to_device_tensor(self._replicated_host_tensor(host, dtype=dtype), device)
249
+
250
+ def _allocate_persistent_inputs(self) -> None:
251
+ """Allocate every stable decode input **before** any trace is captured."""
252
+ self._prefill_page_table = self._replicated_device_tensor(
253
+ torch.full((self.batch, self.pages_per_user), -1, dtype=torch.int32), dtype=ttnn.int32
254
+ )
255
+ self._prefill_sampled = self._replicated_device_tensor(
256
+ torch.zeros((1, 1, 1, SAMPLING_SLOTS), dtype=torch.int32), dtype=ttnn.uint32
257
+ )
258
+ self._width_pools = {}
259
+ self._decode_trace_input_pool = self._decode_input_pool(self.batch)
260
+
261
+ def _decode_input_pool(self, width: int) -> tuple:
262
+ """The four persistent decode inputs, ``width`` rows wide.
263
+
264
+ One pool per captured graph width. The **token** tensor is
265
+ ``[1,1,1,32]`` at every width because that is ``tt_out_tok``'s shape --
266
+ the sampler always addresses 32 slots and ``embed_decode`` slices the
267
+ embedding down to ``model.decode_width``. The other three are the only
268
+ things that bind a request to a row, and they are exactly what
269
+ compaction permutes.
270
+ """
271
+ width = int(width)
272
+ pool = self._width_pools.get(width)
273
+ if pool is not None:
274
+ return pool
275
+ pool = (
276
+ # token: [1,1,1,32] uint32, the tensor ``tt_out_tok`` writes into
277
+ self._replicated_device_tensor(
278
+ torch.zeros((1, 1, 1, SAMPLING_SLOTS), dtype=torch.int32), dtype=ttnn.uint32
279
+ ),
280
+ # current_pos: [width] int32, consumed by paged_update_cache and SDPA
281
+ self._replicated_device_tensor(torch.full((width,), -1, dtype=torch.int32), dtype=ttnn.int32),
282
+ # rotary_position: [1, width] uint32, the cos/sin gather index
283
+ self._replicated_device_tensor(torch.zeros((1, width), dtype=torch.int32), dtype=ttnn.uint32),
284
+ # page_table: [width, pages_per_user] int32
285
+ self._replicated_device_tensor(
286
+ torch.full((width, self.pages_per_user), -1, dtype=torch.int32), dtype=ttnn.int32
287
+ ),
288
+ )
289
+ self._width_pools[width] = pool
290
+ return pool
291
+
292
+ def _ensure_kv_cache(self):
293
+ if self._kv_cache is None:
294
+ self._kv_cache = self.model.allocate_kv_cache(num_blocks=self.num_blocks)
295
+ return self._kv_cache
296
+
297
+ def configure_paging(self, *, page_block_size: int, pages_per_user: int, num_blocks: int) -> None:
298
+ """Adopt a **caller-owned** paging geometry (the vLLM serving mode).
299
+
300
+ Standalone mode derives ``page_block_size`` / ``pages_per_user`` /
301
+ ``num_blocks`` from the model, allocates its own cache and builds its own
302
+ page tables. Serving inverts that: vLLM picks the cache block size and
303
+ the block count, and every page table it hands over is
304
+ ``[batch, max_num_blocks_per_req]`` at *its* width. The two persistent
305
+ page-table tensors were sized for the standalone geometry in
306
+ ``_allocate_persistent_inputs``, so adopting vLLM's means reallocating
307
+ them -- which is only safe before any trace exists, hence the guard.
308
+
309
+ ``tt/generator_vllm.py`` calls this from ``allocate_kv_cache``, which the
310
+ TT plugin invokes once, before warmup and before any forward.
311
+ """
312
+ page_block_size = int(page_block_size)
313
+ pages_per_user = int(pages_per_user)
314
+ num_blocks = int(num_blocks)
315
+ if min(page_block_size, pages_per_user, num_blocks) < 1:
316
+ raise ValueError("paging geometry must be positive")
317
+ if self._trace_model_id is not None or self._trace_sampling_id is not None:
318
+ raise RuntimeError("configure_paging must run before any decode trace is captured")
319
+ if self._kv_cache is not None:
320
+ raise RuntimeError("configure_paging must run before the generator allocates its own cache")
321
+ if (page_block_size, pages_per_user, num_blocks) == (
322
+ self.page_block_size,
323
+ self.pages_per_user,
324
+ self.num_blocks,
325
+ ):
326
+ return
327
+ self.page_block_size = page_block_size
328
+ self.model.page_block_size = page_block_size
329
+ self.pages_per_user = pages_per_user
330
+ self.num_blocks = num_blocks
331
+ self._allocate_persistent_inputs()
332
+
333
+ def decode_device_state(self) -> dict[str, torch.Tensor] | None:
334
+ """The authoritative per-slot decode state that lives **on device**.
335
+
336
+ The traced decode path writes the sampled token straight into the
337
+ persistent token input and advances ``current_pos`` with
338
+ ``ttnn.plus_one``, so after step *N* the device -- not the host -- holds
339
+ the token and position step *N+1* must use. A serving scheduler under
340
+ async scheduling can be a step behind that, and re-installing its host
341
+ view would re-decode a position or feed a stale token. This exposes the
342
+ device view (plus the page table the live trace was captured against) so
343
+ ``tt/generator_vllm.py`` can keep it for slots that are simply
344
+ continuing and take the host's only for slots that changed hands.
345
+
346
+ Returns ``None`` when no trace is live. Costs two small device reads and
347
+ is called only on scheduler-layout changes, never per token.
348
+ """
349
+ if self._trace_model_id is None or self._trace_inputs is None:
350
+ return None
351
+ token, current_pos, _rotary, _page_table = self._trace_inputs
352
+ # The live trace may be **narrower** than the configured slot count, so
353
+ # everything is reported at the graph's width and ``width`` says what
354
+ # that is. Row *i* here is graph row *i*, not necessarily vLLM slot *i* --
355
+ # the caller owns the mapping (``Qwen3CoderForCausalLM._compaction``).
356
+ width = self._trace_graph_width or self.batch
357
+ return {
358
+ "width": width,
359
+ "tokens": _first_device_to_torch(token).reshape(-1)[:width].to(torch.int64),
360
+ "positions": _first_device_to_torch(current_pos).reshape(-1)[:width].to(torch.int64),
361
+ "page_table": (
362
+ None if self._trace_page_table_snapshot is None else self._trace_page_table_snapshot.clone()
363
+ ),
364
+ }
365
+
366
+ def read_sampled_tokens(self, sampled, count: int | None = None) -> torch.Tensor:
367
+ """Host copy of a sampled-token tensor. The only readback on the token path."""
368
+ tokens = self._sampled_to_torch(sampled)
369
+ return tokens if count is None else tokens[: int(count)]
370
+
371
+ def _synchronize(self) -> None:
372
+ ttnn.synchronize_device(self.mesh_device)
373
+ self.trace_stats["explicit_synchronizations"] += 1
374
+
375
+ # -- page tables ----------------------------------------------------------
376
+
377
+ def _page_table_to_torch(self, page_table) -> torch.Tensor:
378
+ if isinstance(page_table, torch.Tensor):
379
+ host = page_table.detach().cpu().to(torch.int32)
380
+ elif isinstance(page_table, ttnn.Tensor):
381
+ host = _first_device_to_torch(page_table).to(torch.int32)
382
+ else:
383
+ raise TypeError("page_table must be a torch or TTNN tensor")
384
+ if host.ndim != 2:
385
+ raise ValueError(f"page_table must be rank two, got {tuple(host.shape)}")
386
+ return host
387
+
388
+ def _normalise_page_table(self, page_table, active_batch: int, width: int | None = None) -> torch.Tensor:
389
+ """Trim/pad a caller's block table to ``width`` rows x ``pages_per_user``.
390
+
391
+ ``width`` defaults to the configured slot count and is the *graph* width
392
+ -- the number of rows the captured decode trace has. A narrow graph is
393
+ handed the first ``width`` rows, which is why the caller must compact its
394
+ live requests into them first.
395
+ """
396
+ width = self.batch if width is None else int(width)
397
+ host = self._page_table_to_torch(page_table)
398
+ if host.shape[0] < active_batch or host.shape[0] > self.batch:
399
+ raise ValueError("page table does not match the configured/active batch")
400
+ if host.shape[1] < self.pages_per_user:
401
+ host = torch.nn.functional.pad(host, (0, self.pages_per_user - host.shape[1]), value=-1)
402
+ elif host.shape[1] > self.pages_per_user:
403
+ host = host[:, : self.pages_per_user]
404
+ if host.shape[0] < width:
405
+ host = torch.nn.functional.pad(host, (0, 0, 0, width - host.shape[0]), value=-1)
406
+ elif host.shape[0] > width:
407
+ host = host[:width]
408
+ return host.contiguous()
409
+
410
+ def _sdpa_rounded_page_count(self, token_count: int) -> int:
411
+ """Physical pages the paged decode SDPA kernel actually reads.
412
+
413
+ The kernel rounds a short sequence up to a power-of-two tile count and a
414
+ long one up to a multiple of eight, and it reads the whole rounded
415
+ window before causal masking. Every rounded tail page therefore needs a
416
+ valid mapping even though it holds no live token yet -- allocating only
417
+ ``ceil(len/block)`` pages produces top-k misses that cliff at exactly
418
+ those boundaries and look like dtype drift.
419
+ """
420
+ if token_count < 1 or token_count > self.model.max_cache_len:
421
+ raise ValueError("SDPA token count is outside the supported context")
422
+ logical_pages = math.ceil(token_count / self.page_block_size)
423
+ if logical_pages <= 8:
424
+ return 1 << (logical_pages - 1).bit_length()
425
+ return 8 * math.ceil(logical_pages / 8)
426
+
427
+ def make_page_table(self, lengths: Sequence[int]) -> torch.Tensor:
428
+ """A disjoint physical-block assignment covering each user's horizon."""
429
+ if len(lengths) > self.batch:
430
+ raise ValueError(f"{len(lengths)} prompts exceed configured batch {self.batch}")
431
+ table = torch.full((self.batch, self.pages_per_user), -1, dtype=torch.int32)
432
+ next_block = 0
433
+ for user, length in enumerate(lengths):
434
+ blocks = self._sdpa_rounded_page_count(int(length))
435
+ if blocks > self.pages_per_user or next_block + blocks > self.num_blocks:
436
+ raise ValueError("paged KV-cache capacity is insufficient for the requested prompts")
437
+ table[user, :blocks] = torch.arange(next_block, next_block + blocks, dtype=torch.int32)
438
+ next_block += blocks
439
+ return table
440
+
441
+ def _validate_page_coverage(self, page_table: torch.Tensor, positions: torch.Tensor, active_batch: int) -> None:
442
+ assigned: set[int] = set()
443
+ for slot, position in enumerate(positions.reshape(-1).tolist()[:active_batch]):
444
+ if position < 0: # inactive row
445
+ continue
446
+ logical_pages = math.ceil((int(position) + 1) / self.page_block_size)
447
+ rounded_pages = self._sdpa_rounded_page_count(int(position) + 1)
448
+ if rounded_pages > page_table.shape[1]:
449
+ raise ValueError(f"slot {slot} page table is too narrow for decode position {position}")
450
+ physical = [int(v) for v in page_table[slot, :rounded_pages].tolist()]
451
+ if any(v < 0 or v >= self.num_blocks for v in physical):
452
+ raise ValueError(f"slot {slot} lacks valid physical pages for the rounded SDPA read at {position}")
453
+ live = physical[:logical_pages]
454
+ if len(set(live)) != len(live) or assigned.intersection(live):
455
+ raise ValueError("active page-table rows must map disjoint physical cache pages")
456
+ assigned.update(live)
457
+
458
+ # -- prefill --------------------------------------------------------------
459
+
460
+ @property
461
+ def prefill_buckets(self) -> tuple[int, ...]:
462
+ """The prefill length ladder in force, ascending. ``()`` means exact."""
463
+ if self._prefill_buckets is None:
464
+ name = os.getenv("QWEN3_PREFILL_BUCKETS", DEFAULT_PREFILL_BUCKETS).strip().lower()
465
+ self._prefill_buckets = prefill_bucket_ladder(name, self.model.max_cache_len)
466
+ return self._prefill_buckets
467
+
468
+ def prefill_padded_len(self, real_len: int, *, start: int = 0) -> int:
469
+ """Round a suffix length up to its bucket, or return it unchanged.
470
+
471
+ Clamped so ``start + result`` never leaves the supported context: near
472
+ the very top of the window there may be no rung left to round to, and
473
+ such a prompt simply keeps its exact length and pays a compile. Every
474
+ other length lands on a rung.
475
+ """
476
+ real_len = int(real_len)
477
+ ladder = self.prefill_buckets
478
+ if not ladder:
479
+ return real_len
480
+ index = bisect.bisect_left(ladder, real_len)
481
+ if index >= len(ladder):
482
+ return real_len
483
+ return max(real_len, min(ladder[index], self.model.max_cache_len - int(start)))
484
+
485
+ @staticmethod
486
+ def _prefill_starts(active_batch: int, start_pos) -> list[int]:
487
+ """How many tokens of each row are already in the cache. None/0 => all new."""
488
+ if start_pos is None:
489
+ return [0] * active_batch
490
+ if isinstance(start_pos, int):
491
+ return [int(start_pos)] * active_batch
492
+ return [int(v) for v in start_pos]
493
+
494
+ def _prefill_rope_horizon(self, prompt_lens, start_pos) -> int:
495
+ """The highest absolute position any row of this prefill will rotate at.
496
+
497
+ Bucketing pads past the prompt, so the RoPE tables have to cover the
498
+ PADDED horizon. Getting this wrong is not a slow path but a crash:
499
+ ``prefill_forward`` releases the decode traces when the tables grow,
500
+ precisely because a captured trace holds the old tables' identities,
501
+ and a growth that happened later -- inside ``prefill_hidden``, under a
502
+ live trace -- would slip past that guard.
503
+ """
504
+ starts = self._prefill_starts(len(prompt_lens), start_pos)
505
+ return max(
506
+ start + self.prefill_padded_len(int(length) - start, start=start)
507
+ for length, start in zip(prompt_lens, starts)
508
+ )
509
+
510
+ def _release_decode_traces_before_allocating(self) -> None:
511
+ """Prefill is eager and allocates; a live trace makes that unsafe."""
512
+ if self._trace_model_id is None and self._trace_sampling_id is None:
513
+ return
514
+ self._synchronize()
515
+ self._release_decode_traces()
516
+
517
+ def prefill_forward(
518
+ self,
519
+ tokens: torch.Tensor,
520
+ *,
521
+ page_table,
522
+ kv_cache: Any,
523
+ prompt_lens: Sequence[int],
524
+ return_all_logits: bool = False,
525
+ sampling_mode: str = "host",
526
+ preserve_decode_traces: bool = False,
527
+ start_pos: Sequence[int] | int | None = None,
528
+ **kwargs: Any,
529
+ ):
530
+ """Prefill arbitrary logical lengths, one user at a time into the cache.
531
+
532
+ ``tokens`` is ``[active_batch, width]`` and ``prompt_lens`` gives each
533
+ row's **real** length. Rows may differ in length, and no mask is needed
534
+ at any length: nothing is padded to a chunk, tile or page boundary, and
535
+ the returned logits are always sliced back to the logical prompt length.
536
+
537
+ Each row is prefilled at its bucket rather than at its exact length --
538
+ see :func:`prefill_bucket_ladder`, and ``QWEN3_PREFILL_BUCKETS=exact``
539
+ to turn that off. Bucketing is invisible here: the padding rows are
540
+ zeros appended *after* the prompt, so causality keeps them out of every
541
+ real row's attention, ``fill_len`` keeps them out of the KV cache, and
542
+ the row this method selects is the real last token either way.
543
+
544
+ ``preserve_decode_traces`` keeps a captured decode trace alive across
545
+ this prefill. Standalone callers never need it -- ``generate`` prefills
546
+ once, before any decode trace exists. **Serving does**: vLLM admits a new
547
+ request by prefilling it while other slots are mid-decode, and releasing
548
+ the decode traces there would re-capture them on the very next token,
549
+ putting a multi-second stall inside the measured inter-token latency of
550
+ every other in-flight request. It is safe because prefill's allocations
551
+ never touch the trace region and every tensor a captured trace holds --
552
+ ``_decode_trace_input_pool``, ``_trace_logits``, ``_trace_sampled`` --
553
+ is owned by this object and therefore never freed underneath it. The
554
+ page table is the one shared binding, and this method rebinds the cache
555
+ back to the live trace's page-table tensor before it returns.
556
+ """
557
+ if sampling_mode not in {"host", "device"}:
558
+ raise ValueError("sampling_mode must be 'host' or 'device'")
559
+ if sampling_mode == "device" and return_all_logits:
560
+ raise ValueError("return_all_logits is incompatible with device sampling")
561
+ if tokens.ndim != 2:
562
+ raise ValueError(f"tokens must be [batch,seq], got {tuple(tokens.shape)}")
563
+ active_batch, logical_width = int(tokens.shape[0]), int(tokens.shape[1])
564
+ if not 1 <= active_batch <= self.batch:
565
+ raise ValueError(f"active batch must be in [1,{self.batch}]")
566
+ if len(prompt_lens) != active_batch or any(not 1 <= int(n) <= logical_width for n in prompt_lens):
567
+ raise ValueError("prompt_lens must contain one valid logical length per input row")
568
+ if max(prompt_lens) > self.model.max_cache_len:
569
+ raise ValueError("prompt exceeds the supported context")
570
+
571
+ if preserve_decode_traces:
572
+ if self._trace_model_id is not None:
573
+ # The replayed trace is asynchronous; prefill is eager. Let the
574
+ # queue drain before eager work reads or writes the same cache.
575
+ self._synchronize()
576
+ if self.model.ensure_rope_capacity(self._prefill_rope_horizon(prompt_lens, start_pos)):
577
+ # Growing the tables moves them, and a captured trace holds the
578
+ # old identities. Nothing can preserve a trace across that.
579
+ self._release_decode_traces()
580
+ else:
581
+ self._release_decode_traces_before_allocating()
582
+ self.model.ensure_rope_capacity(self._prefill_rope_horizon(prompt_lens, start_pos))
583
+ caches = self._ensure_kv_cache() if kv_cache is None else kv_cache
584
+ page_host = self._normalise_page_table(page_table, active_batch)
585
+ self._copy_host(page_host, self._prefill_page_table, dtype=ttnn.int32)
586
+ self.model.bind_page_table(caches, self._prefill_page_table)
587
+ try:
588
+ return self._prefill_body(
589
+ tokens,
590
+ caches,
591
+ active_batch=active_batch,
592
+ logical_width=logical_width,
593
+ prompt_lens=prompt_lens,
594
+ return_all_logits=return_all_logits,
595
+ sampling_mode=sampling_mode,
596
+ start_pos=start_pos,
597
+ page_table=page_host,
598
+ )
599
+ finally:
600
+ if self._trace_inputs is not None:
601
+ # Hand the cache back to the tensor the live decode trace was
602
+ # captured against, so the next replay writes through the page
603
+ # table the scheduler owns rather than the prefill scratch one.
604
+ self.model.bind_page_table(caches, self._trace_inputs[3])
605
+
606
+ def _prefill_body(
607
+ self,
608
+ tokens: torch.Tensor,
609
+ caches,
610
+ *,
611
+ active_batch: int,
612
+ logical_width: int,
613
+ prompt_lens: Sequence[int],
614
+ return_all_logits: bool,
615
+ sampling_mode: str,
616
+ start_pos: Sequence[int] | int | None = None,
617
+ page_table=None,
618
+ ):
619
+ # ``start_pos`` is how many tokens of each row are ALREADY in the cache.
620
+ # None/0 is the shipped whole-prompt prefill and takes an identical path.
621
+ starts = self._prefill_starts(active_batch, start_pos)
622
+ per_user_logits: list[torch.Tensor] = []
623
+ selected_rows = []
624
+ for user in range(active_batch):
625
+ prompt_len = int(prompt_lens[user])
626
+ start = starts[user]
627
+ # Bucketed prefill: pad the suffix up to its rung so the shape-hungry
628
+ # ops (QKV/wo projections, MoE, SDPA) see one of a handful of
629
+ # lengths instead of one per prompt. ``fill_len`` keeps the cache
630
+ # write at the real length, ``select_prefill_rows`` indexes the real
631
+ # last row, and the padding rows are causally invisible to it.
632
+ token_host = tokens[user : user + 1, start:prompt_len].to(torch.int32)
633
+ real_len = int(token_host.shape[1])
634
+ padded_len = self.prefill_padded_len(real_len, start=start)
635
+ fill_len = None
636
+ if padded_len > real_len:
637
+ token_host = torch.nn.functional.pad(token_host, (0, padded_len - real_len))
638
+ fill_len = real_len
639
+
640
+ chunk_pt = fill_pt = None
641
+ if start:
642
+ block = int(caches[0].block_size)
643
+ if start % block:
644
+ raise ValueError(f"start_pos {start} is not a multiple of the block size {block}")
645
+ if not 0 < start < prompt_len:
646
+ raise ValueError(f"start_pos {start} must be inside (0, prompt_len={prompt_len})")
647
+ if page_table is None:
648
+ raise ValueError("a split prefill needs a page table")
649
+ row = torch.as_tensor(page_table)[user : user + 1].to(torch.int32)
650
+ # Two different tables, and the difference is the whole trick:
651
+ # chunk_pt -- the user's FULL row, so chunked SDPA can read the
652
+ # cached prefix from absolute block 0;
653
+ # fill_pt -- a window over the suffix's blocks, because
654
+ # paged_fill_cache writes relative to block 0 of the
655
+ # table it is handed.
656
+ chunk_pt = ttnn.from_torch(
657
+ row,
658
+ device=self.mesh_device,
659
+ dtype=ttnn.int32,
660
+ layout=ttnn.ROW_MAJOR_LAYOUT,
661
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
662
+ )
663
+ # WIDTH IS DERIVED FROM THE BUCKET, NOT FROM ``start``. The
664
+ # obvious ``row[:, start // block :]`` is one column narrower
665
+ # per cached block, so its shape -- and therefore the
666
+ # ``paged_fill_cache`` program -- varies with the length of the
667
+ # prefix, which is unbounded. Measured: an unwarmed prefix
668
+ # offset cost 0.976 s against 0.178 s for a warmed one, on
669
+ # identical suffix and chunk sizes. Sizing the window by the
670
+ # suffix's rung instead leaves one program per rung.
671
+ #
672
+ # Trailing ``-1`` entries are never read: ``k`` is sliced to
673
+ # ``fill_len`` above, so the op writes only
674
+ # ``ceil(fill_len / block)`` blocks from the front of this table.
675
+ fill_blocks = math.ceil(padded_len / block)
676
+ fill_window = row[:, start // block : start // block + fill_blocks]
677
+ if fill_window.shape[1] < fill_blocks:
678
+ fill_window = torch.nn.functional.pad(
679
+ fill_window, (0, fill_blocks - fill_window.shape[1]), value=-1
680
+ )
681
+ fill_pt = ttnn.from_torch(
682
+ fill_window.contiguous(),
683
+ device=self.mesh_device,
684
+ dtype=ttnn.int32,
685
+ layout=ttnn.ROW_MAJOR_LAYOUT,
686
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
687
+ )
688
+ token_device = ttnn.from_torch(
689
+ token_host,
690
+ device=self.mesh_device,
691
+ dtype=ttnn.uint32,
692
+ layout=ttnn.ROW_MAJOR_LAYOUT,
693
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
694
+ )
695
+ hidden = self.model.prefill_hidden(
696
+ token_device,
697
+ kv_cache=caches,
698
+ user_id=user,
699
+ start_pos=start,
700
+ chunk_page_table=chunk_pt,
701
+ fill_page_table=fill_pt,
702
+ fill_len=fill_len,
703
+ )
704
+ ttnn.deallocate(token_device, True)
705
+ for t in (chunk_pt, fill_pt):
706
+ if t is not None:
707
+ ttnn.deallocate(t, True)
708
+ if return_all_logits:
709
+ normed = self.model.prefill_norm(hidden)
710
+ ttnn.deallocate(hidden, True)
711
+ local = self.model.local_logits(normed)
712
+ ttnn.deallocate(normed, True)
713
+ host = self.model.gather_logits_to_torch(local)[0, 0, : prompt_len - start, :] # padding rows dropped
714
+ ttnn.deallocate(local, True)
715
+ per_user_logits.append(
716
+ torch.nn.functional.pad(host, (0, 0, 0, logical_width - (prompt_len - start))).unsqueeze(0)
717
+ )
718
+ else:
719
+ # The last row of THIS chunk: absolute prompt_len-1 is row
720
+ # prompt_len-1-start within the suffix.
721
+ selected_rows.append(self.model.select_prefill_rows(hidden, [prompt_len - 1 - start]))
722
+ ttnn.deallocate(hidden, True)
723
+
724
+ if return_all_logits:
725
+ return torch.cat(per_user_logits, dim=0)[:, :logical_width]
726
+
727
+ selected = (
728
+ selected_rows[0]
729
+ if len(selected_rows) == 1
730
+ else ttnn.concat(selected_rows, dim=2, memory_config=ttnn.DRAM_MEMORY_CONFIG)
731
+ )
732
+ normed = self.model.prefill_norm(selected)
733
+ if selected is not selected_rows[0] or len(selected_rows) > 1:
734
+ for row in selected_rows:
735
+ ttnn.deallocate(row, True)
736
+ else:
737
+ ttnn.deallocate(selected, True)
738
+ if sampling_mode == "device":
739
+ padded = self._pad_rows_to_sampling_slots(normed, active_batch)
740
+ local = self.model.local_logits(padded)
741
+ ttnn.deallocate(padded, True)
742
+ with self._penalties_suspended():
743
+ sampled = self._sample_device(local, tt_out_tok=self._prefill_sampled)
744
+ ttnn.deallocate(local, True)
745
+ return sampled
746
+ local = self.model.local_logits(normed)
747
+ ttnn.deallocate(normed, True)
748
+ host = self.model.gather_logits_to_torch(local, valid_rows=active_batch)[0, 0]
749
+ ttnn.deallocate(local, True)
750
+ return host.unsqueeze(1)
751
+
752
+ def _pad_rows_to_sampling_slots(self, normed, active_batch: int):
753
+ """``ttnn.sampling`` works in 32 fixed slots; pad the selected rows up."""
754
+ rows = int(normed.shape[-2])
755
+ if rows >= SAMPLING_SLOTS:
756
+ return normed
757
+ padded = ttnn.pad(
758
+ normed,
759
+ [(0, 0), (0, 0), (0, SAMPLING_SLOTS - rows), (0, 0)],
760
+ value=0.0,
761
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
762
+ )
763
+ ttnn.deallocate(normed, True)
764
+ return padded
765
+
766
+ # -- sampling parameters --------------------------------------------------
767
+
768
+ def _ensure_sampling_params(self):
769
+ if self._sampling_params is None:
770
+ self._sampling_params = (
771
+ self._replicated_device_tensor(torch.ones(SAMPLING_SLOTS, dtype=torch.int32), dtype=ttnn.uint32),
772
+ self._replicated_device_tensor(torch.zeros(SAMPLING_SLOTS, dtype=torch.bfloat16), dtype=ttnn.bfloat16),
773
+ self._replicated_device_tensor(torch.ones(SAMPLING_SLOTS, dtype=torch.bfloat16), dtype=ttnn.bfloat16),
774
+ )
775
+ self._sampling_snapshot = ((1,) * SAMPLING_SLOTS, (0.0,) * SAMPLING_SLOTS, (1.0,) * SAMPLING_SLOTS)
776
+ return self._sampling_params
777
+
778
+ @staticmethod
779
+ def _expand(value, *, active_batch: int, inactive, name: str):
780
+ active = [value] * active_batch if isinstance(value, (int, float)) else list(value)
781
+ if len(active) != active_batch:
782
+ raise ValueError(f"{name} must be scalar or contain {active_batch} values")
783
+ return active + [inactive] * (SAMPLING_SLOTS - active_batch)
784
+
785
+ def set_sampling_params(self, *, top_k=1, top_p=0.0, temperature=1.0, active_batch: int = 1) -> None:
786
+ """Set per-slot ``(k, p, temperature)``. ``k=1`` is exactly greedy."""
787
+ if not 1 <= active_batch <= self.batch:
788
+ raise ValueError(f"active_batch must be in [1,{self.batch}]")
789
+ k = [int(v) for v in self._expand(top_k, active_batch=active_batch, inactive=1, name="top_k")]
790
+ p = [float(v) for v in self._expand(top_p, active_batch=active_batch, inactive=0.0, name="top_p")]
791
+ temp = [
792
+ float(v) for v in self._expand(temperature, active_batch=active_batch, inactive=0.0, name="temperature")
793
+ ]
794
+ if any(not 0.0 <= v <= 1.0 for v in p[:active_batch]):
795
+ raise ValueError("top_p must be in [0,1]")
796
+ if any(v < 0.0 for v in temp[:active_batch]):
797
+ raise ValueError("temperature must be non-negative")
798
+ for slot in range(active_batch):
799
+ # A serving stack spells greedy as temperature=0 / top_k=0.
800
+ if temp[slot] == 0.0 and k[slot] == 0:
801
+ k[slot] = 1
802
+ if any(not 1 <= v <= SAMPLING_SLOTS for v in k[:active_batch]):
803
+ raise ValueError("top_k must be in [1,32] (0 accepted only with temperature=0)")
804
+ device_temp = []
805
+ for slot, value in enumerate(temp):
806
+ if value == 0.0:
807
+ k[slot], p[slot], value = 1, 0.0, 1.0
808
+ device_temp.append(1.0 / value)
809
+ stochastic = any(v > 1 for v in k[:active_batch]) or any(v > 0.0 for v in p[:active_batch])
810
+ if stochastic != self._sampling_stochastic and self._trace_model_id is not None:
811
+ self._release_decode_traces()
812
+ self._sampling_stochastic = stochastic
813
+ params = self._ensure_sampling_params()
814
+ snapshot = (tuple(k), tuple(p), tuple(device_temp))
815
+ if snapshot == self._sampling_snapshot:
816
+ return
817
+ for host, device, dtype in (
818
+ (torch.tensor(k, dtype=torch.int32), params[0], ttnn.uint32),
819
+ (torch.tensor(p, dtype=torch.bfloat16), params[1], ttnn.bfloat16),
820
+ (torch.tensor(device_temp, dtype=torch.bfloat16), params[2], ttnn.bfloat16),
821
+ ):
822
+ self._copy_host(host, device, dtype=dtype)
823
+ self.trace_stats["sampling_param_host_copies"] += 1
824
+ self._sampling_snapshot = snapshot
825
+
826
+ # -- sampling penalties ---------------------------------------------------
827
+
828
+ @staticmethod
829
+ def _row_token_ids(history, row: int) -> torch.Tensor:
830
+ """Row ``row`` of a vLLM ``[rows, L]`` history tensor, -1 padding dropped.
831
+
832
+ vLLM pads both ``prompt_tokens`` and ``output_tokens`` with **-1**
833
+ (``input_batch.make_prompt_token_ids_tensor``: "TT device sampling relies
834
+ on -1 as the padding sentinel"), and pads the *batch* to ``max_num_reqs``
835
+ with all-(-1) rows. Dropping every negative entry handles both.
836
+ """
837
+ if history is None:
838
+ return torch.empty(0, dtype=torch.int64)
839
+ tensor = torch.as_tensor(history)
840
+ if tensor.ndim == 1:
841
+ tensor = tensor.reshape(1, -1)
842
+ if row >= tensor.shape[0]:
843
+ return torch.empty(0, dtype=torch.int64)
844
+ ids = tensor[row].reshape(-1).to(torch.int64)
845
+ return ids[ids >= 0]
846
+
847
+ def _ensure_penalty_host(self, slots: int, vocab: int) -> dict:
848
+ """Per-die staging buffers, **already contiguous in the shard layout**.
849
+
850
+ Not one ``[1,1,32,151936]`` tensor. Handing a full-width host tensor to
851
+ ``ttnn.ShardTensorToMesh(dim=-1)`` makes it re-slice a strided view into
852
+ four contiguous copies on every decode step, and that reshard -- not
853
+ tilization, not the wire -- was **6.601 ms of a 6.897 ms** upload.
854
+ Keeping the four ``[1,1,32,37984]`` buffers contiguous from the start and
855
+ assembling them with ``ttnn.from_host_shards`` costs **2.049 ms**
856
+ end to end, 3.4x less.
857
+
858
+ The trade is that the global -> (die, local) split now happens here, in
859
+ host Python, instead of being implied by the mesh mapper. That is the
860
+ one piece of index arithmetic in this feature, so it is checked rather
861
+ than trusted: ``penalty_shard_boundary_probe.py``'s
862
+ ``fast_staging_matches_shard_mapper`` leg builds the same operand both
863
+ ways and requires the two device tensors to be bit-identical, and its
864
+ cross-die/boundary legs would fail first if the split were wrong.
865
+ """
866
+ if self._penalty_host is None:
867
+ devices, local = self.model.sampler.penalty_shard_geometry()
868
+ if devices * local != vocab:
869
+ raise RuntimeError(f"penalty shard geometry {devices}x{local} does not cover {vocab}")
870
+ self._penalty_local_vocab = local
871
+ self._penalty_host = {
872
+ "rep_neg": [torch.ones((1, 1, slots, local), dtype=torch.bfloat16) for _ in range(devices)],
873
+ "add": [torch.zeros((1, 1, slots, local), dtype=torch.bfloat16) for _ in range(devices)],
874
+ }
875
+ self._penalty_prev_add = [None] * slots
876
+ self._penalty_prev_rep = [None] * slots
877
+ return self._penalty_host
878
+
879
+ def _penalty_split(self, ids: torch.Tensor):
880
+ """Global token ids -> ``[(die, local_ids_on_that_die, selector), ...]``.
881
+
882
+ ``die = t // local_vocab``, ``local = t % local_vocab`` -- the same
883
+ contiguous ascending decomposition ``_dist_die_offset`` is built from,
884
+ with ``local_vocab`` read off the sampler rather than re-derived here.
885
+ """
886
+ local_vocab = self._penalty_local_vocab
887
+ die = torch.div(ids, local_vocab, rounding_mode="floor")
888
+ local = ids - die * local_vocab
889
+ out = []
890
+ for index in range(len(self._penalty_host["rep_neg"])):
891
+ selector = die == index
892
+ if bool(selector.any()):
893
+ out.append((index, local[selector], selector))
894
+ return out
895
+
896
+ def set_penalty_params(
897
+ self,
898
+ *,
899
+ presence=None,
900
+ frequency=None,
901
+ repetition=None,
902
+ prompt_tokens=None,
903
+ output_tokens=None,
904
+ active_batch: int = 1,
905
+ ) -> tuple[bool, bool]:
906
+ """Stage the three vLLM sampling penalties for the next decode step.
907
+
908
+ Returns ``(live, graph_changed)``: whether the penalty stage runs this
909
+ step, and whether the decode graph changed shape -- the caller must
910
+ reinstall the trace when it did, because a mode change releases it. Everything
911
+ here is per **global** token id; the global -> die mapping is done by the
912
+ same ``ShardTensorToMesh(dim=-1)`` split the logits themselves live under,
913
+ so no index arithmetic can put a penalty on the wrong die. The argument is
914
+ in ``_WatcherCleanSampling1D``'s penalty section.
915
+
916
+ Semantics are vLLM's ``model_executor/layers/utils.py::apply_penalties``,
917
+ including its order: repetition (over prompt+output) multiplies the raw
918
+ logit, then frequency (output counts) and presence (output mask) subtract.
919
+ """
920
+ sampler = self.model.sampler
921
+ slots, vocab = sampler.penalty_buffer_shape()
922
+ rows = max(0, min(int(active_batch), slots))
923
+
924
+ def _row_values(values, neutral):
925
+ if values is None:
926
+ return [neutral] * rows
927
+ if isinstance(values, (int, float)):
928
+ return [float(values)] * rows
929
+ listed = [float(v) for v in list(values)[:rows]]
930
+ return listed + [neutral] * (rows - len(listed))
931
+
932
+ presence = _row_values(presence, 0.0)
933
+ frequency = _row_values(frequency, 0.0)
934
+ repetition = _row_values(repetition, 1.0)
935
+
936
+ rep_rows = [r for r in range(rows) if repetition[r] != 1.0]
937
+ add_rows = [r for r in range(rows) if presence[r] != 0.0 or frequency[r] != 0.0]
938
+ if (rep_rows or add_rows) and prompt_tokens is None and output_tokens is None:
939
+ # vLLM only sends the history when a penalty is live (and only on
940
+ # decode). Without it there is nothing to key a penalty on; run the
941
+ # unpenalised graph rather than invent one.
942
+ rep_rows, add_rows = [], []
943
+ mode = (1 if rep_rows else 0) | (2 if add_rows else 0)
944
+
945
+ if mode != self._penalty_mode:
946
+ # A graph change, exactly like the argmax/split flip: the ops either
947
+ # are or are not in the captured trace, so the trace must go.
948
+ if self._trace_model_id is not None or self._trace_sampling_id is not None:
949
+ self._release_decode_traces()
950
+ self._decode_warm_key = None
951
+ sampler.allocate_penalty_buffers(mode)
952
+ self._penalty_mode = mode
953
+ graph_changed = True
954
+ else:
955
+ graph_changed = False
956
+ if mode == 0:
957
+ return False, graph_changed
958
+
959
+ host = self._ensure_penalty_host(slots, vocab)
960
+ add, rep_neg = host["add"], host["rep_neg"]
961
+
962
+ # Reset only what this row wrote last step, not the whole 151936-wide
963
+ # row: the history is at most the context length and is usually far
964
+ # shorter, so this is O(history) rather than O(vocabulary).
965
+ for row in range(slots):
966
+ if mode & 2:
967
+ previous = self._penalty_prev_add[row]
968
+ if previous is not None:
969
+ for die, local, _ in self._penalty_split(previous):
970
+ add[die][0, 0, row].index_fill_(0, local, 0.0)
971
+ self._penalty_prev_add[row] = None
972
+ if mode & 1:
973
+ previous = self._penalty_prev_rep[row]
974
+ if previous is not None:
975
+ for die, local, _ in self._penalty_split(previous):
976
+ rep_neg[die][0, 0, row].index_fill_(0, local, 1.0)
977
+ self._penalty_prev_rep[row] = None
978
+
979
+ for row in add_rows:
980
+ out_ids = self._row_token_ids(output_tokens, row)
981
+ if out_ids.numel() == 0:
982
+ continue
983
+ unique, counts = torch.unique(out_ids, return_counts=True)
984
+ # f * count(t in output) + q * (count > 0), summed on the host so the
985
+ # device sees one additive tensor rather than two.
986
+ values = (counts.to(torch.float32) * frequency[row] + presence[row]).to(torch.bfloat16)
987
+ for die, local, selector in self._penalty_split(unique):
988
+ add[die][0, 0, row].index_copy_(0, local, values[selector])
989
+ self._penalty_prev_add[row] = unique
990
+
991
+ for row in rep_rows:
992
+ ids = torch.cat((self._row_token_ids(prompt_tokens, row), self._row_token_ids(output_tokens, row)))
993
+ if ids.numel() == 0:
994
+ continue
995
+ unique = torch.unique(ids)
996
+ # Only ``p`` is staged; ``1/p - p`` is derived on device from it. See
997
+ # ``_WatcherCleanSampling1D._apply_penalties``.
998
+ for die, local, _ in self._penalty_split(unique):
999
+ rep_neg[die][0, 0, row].index_fill_(0, local, repetition[row])
1000
+ self._penalty_prev_rep[row] = unique
1001
+
1002
+ buffers = sampler.penalty_device_buffers()
1003
+ for name in (("rep_neg",) if mode & 1 else ()) + (("add",) if mode & 2 else ()):
1004
+ self._upload_penalty_tensor(host[name], buffers[name])
1005
+ self.trace_stats["penalty_host_copies"] += 1
1006
+ return True, graph_changed
1007
+
1008
+ def _upload_penalty_tensor(self, shards: list, device) -> None:
1009
+ """Four contiguous ``[1,1,32,37984]`` host buffers -> the four die shards.
1010
+
1011
+ ``ttnn.from_host_shards`` assembles them into the multi-device host
1012
+ tensor directly, so nothing re-slices a 9.7 MB strided view per step.
1013
+ Shard ``d`` is die ``d``'s columns by the ordering
1014
+ ``ShardTensorToMesh(dim=-1)`` uses, which the probe pins by building the
1015
+ same operand both ways and requiring bit-identical device tensors.
1016
+ """
1017
+ ttnn.copy_host_to_device_tensor(
1018
+ ttnn.from_host_shards(
1019
+ [ttnn.from_torch(shard, device=None, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT) for shard in shards],
1020
+ self.mesh_device.shape,
1021
+ ),
1022
+ device,
1023
+ )
1024
+
1025
+ @contextlib.contextmanager
1026
+ def _penalties_suspended(self):
1027
+ """Run the enclosed sampling without the penalty stage.
1028
+
1029
+ Prefill and the eager host/device decode compatibility paths sample rows
1030
+ that are **not** the decode trace's slots (a prefill's row *i* is the
1031
+ *i*-th admitted request, not slot *i*), and vLLM does not send a token
1032
+ history for them -- it populates ``prompt_tokens``/``output_tokens``
1033
+ "if penalties are needed (decode only)". Applying another slot's staged
1034
+ penalty row to them would penalise the wrong tokens, so the stage is off.
1035
+ """
1036
+ sampler = self.model.sampler
1037
+ saved = sampler._penalty_mode
1038
+ sampler._penalty_mode = 0
1039
+ try:
1040
+ yield
1041
+ finally:
1042
+ sampler._penalty_mode = saved
1043
+
1044
+ def _sample_device(self, logits, *, tt_out_tok=None):
1045
+ """Greedy takes the argmax strategy, anything sampled takes the split one.
1046
+
1047
+ Both are ``Sampling1D``, both traced, both write ``tt_out_tok``. The
1048
+ split is by *request*, not by convenience: greedy is exactly top-1, and
1049
+ at this vocabulary the argmax strategy computes it 6.6x faster
1050
+ (0.928 ms against 6.155 ms, ``doc/optimized_full_model/README.md``,
1051
+ "The sampler comparison"). Changing
1052
+ between the two releases the decode traces, which ``set_sampling_params``
1053
+ already does when ``_sampling_stochastic`` flips.
1054
+ """
1055
+ k, p, temp = self._ensure_sampling_params()
1056
+ if not self._sampling_stochastic:
1057
+ return self.model.sample_greedy_argmax(logits, tt_out_tok=tt_out_tok)
1058
+ return self.model.sample_split(logits, k=k, p=p, temp=temp, tt_out_tok=tt_out_tok)
1059
+
1060
+ def _sampled_to_torch(self, sampled) -> torch.Tensor:
1061
+ self.trace_stats["caller_token_readbacks"] += 1
1062
+ return _first_device_to_torch(sampled).reshape(-1)[: self.batch].to(torch.long)
1063
+
1064
+ # -- decode trace ---------------------------------------------------------
1065
+
1066
+ def _prepare_decode_host_inputs(
1067
+ self, tokens: torch.Tensor, positions: torch.Tensor, page_table: torch.Tensor, width: int | None = None
1068
+ ):
1069
+ width = self.batch if width is None else int(width)
1070
+ tokens = tokens.reshape(-1).to(torch.int64)
1071
+ positions = positions.reshape(-1).to(torch.int64)
1072
+ if tokens.numel() > width or positions.numel() > width:
1073
+ raise ValueError("decode batch exceeds the graph width")
1074
+ padded_tokens = torch.zeros(SAMPLING_SLOTS, dtype=torch.int32)
1075
+ padded_tokens[: tokens.numel()] = tokens.to(torch.int32)
1076
+ padded_positions = torch.full((width,), -1, dtype=torch.int32)
1077
+ padded_positions[: positions.numel()] = positions.to(torch.int32)
1078
+ rotary = torch.clamp(padded_positions, min=0).reshape(1, width)
1079
+ return (
1080
+ self._replicated_host_tensor(padded_tokens.reshape(1, 1, 1, SAMPLING_SLOTS), dtype=ttnn.uint32),
1081
+ self._replicated_host_tensor(padded_positions, dtype=ttnn.int32),
1082
+ self._replicated_host_tensor(rotary, dtype=ttnn.uint32),
1083
+ self._replicated_host_tensor(page_table, dtype=ttnn.int32),
1084
+ )
1085
+
1086
+ def _restore_trace_inputs(self, host_inputs, *, include_page_table: bool, token_device=None) -> None:
1087
+ count = 4 if include_page_table else 3
1088
+ start = 0
1089
+ if token_device is not None:
1090
+ ttnn.copy(token_device, self._trace_inputs[0])
1091
+ self.trace_stats["token_device_copies"] += 1
1092
+ start = 1
1093
+ for index in range(start, count):
1094
+ ttnn.copy_host_to_device_tensor(host_inputs[index], self._trace_inputs[index])
1095
+ if token_device is None:
1096
+ self.trace_stats["token_host_copies"] += 1
1097
+ self.trace_stats["position_host_copies"] += 1
1098
+ self.trace_stats["rotary_position_host_copies"] += 1
1099
+ if include_page_table:
1100
+ self.trace_stats["page_table_host_copies"] += 1
1101
+
1102
+ def _decode_graph_key(self, kv_cache, graph_width: int) -> tuple:
1103
+ """Everything that changes which programs the decode graph needs.
1104
+
1105
+ ``rope_cache_len`` is part of it because ``_ensure_decode_rope_capacity``
1106
+ reallocates the cos/sin tables at a *new length* when the horizon grows,
1107
+ which changes the shapes ``ttnn.embedding`` and the untilize behind it
1108
+ run at. Those programs are not in the cache yet, and a trace capture
1109
+ cannot compile them ("Cannot load new binaries during trace capture").
1110
+ Without the length in the key, ``_decode_compiled_keys`` would claim the
1111
+ graph was already warm and skip the eager pass that compiles them.
1112
+ """
1113
+ return (
1114
+ id(kv_cache),
1115
+ graph_width,
1116
+ self._sampling_stochastic,
1117
+ self._penalty_mode,
1118
+ self.model.rope_cache_len,
1119
+ # ``Qwen3CoderModel.active_row_gating`` adds four small ops per step
1120
+ # and one broadcast multiply per layer; flipping it is a different
1121
+ # program set, so a stale "already compiled" claim would try to load
1122
+ # binaries inside an open capture. Same failure mode
1123
+ # ``rope_cache_len`` was added for -- see the docstring above and
1124
+ # ``doc/vllm_integration/work_log.md`` §12.
1125
+ self.model.active_row_gating,
1126
+ )
1127
+
1128
+ def _warm_decode_graphs(self, host_inputs, kv_cache, *, graph_width: int, initial_token_device=None) -> None:
1129
+ """Compile every program once eagerly.
1130
+
1131
+ Non-negotiable rather than merely tidy: ``_decode_ccl_buffers``
1132
+ allocates the two persistent collective buffers on the first call at a
1133
+ shape, and ``ttnn.from_torch`` inside ``begin_trace_capture`` raises and
1134
+ leaves the capture open -- a hung mesh (stage-04 ``work_log.md`` §6).
1135
+ """
1136
+ self._trace_inputs = self._decode_input_pool(graph_width)
1137
+ self._restore_trace_inputs(host_inputs, include_page_table=True, token_device=initial_token_device)
1138
+ token, current_pos, rotary_pos, page_table = self._trace_inputs
1139
+ self.model.bind_page_table(kv_cache, page_table)
1140
+ with self.model.decode_width_scope(graph_width):
1141
+ logits = self.model.decode_forward_from_ttnn_inputs(
1142
+ # ``advance_position=True`` here as well as in the capture: every
1143
+ # op the traced graph contains must already be in the program
1144
+ # cache, and that includes the two ``ttnn.plus_one`` calls. The
1145
+ # positions this leaves behind are overwritten by the restore below.
1146
+ token,
1147
+ current_pos,
1148
+ rotary_position=rotary_pos,
1149
+ kv_cache=kv_cache,
1150
+ advance_position=True,
1151
+ )
1152
+ self._sample_device(logits, tt_out_tok=token)
1153
+ ttnn.deallocate(logits, True)
1154
+ self._synchronize()
1155
+ self._restore_trace_inputs(host_inputs, include_page_table=True, token_device=initial_token_device)
1156
+ self._synchronize()
1157
+ self._decode_warm_key = self._decode_graph_key(kv_cache, graph_width)
1158
+ self._decode_compiled_keys.add(self._decode_warm_key)
1159
+ self.trace_stats["decode_warmups"] += 1
1160
+
1161
+ def _capture_decode_traces(
1162
+ self, host_inputs, kv_cache, *, graph_width: int, active_batch: int, initial_token_device=None
1163
+ ) -> None:
1164
+ self._trace_inputs = self._decode_input_pool(graph_width)
1165
+ model_trace_id = sampling_trace_id = None
1166
+ model_open = sampling_open = False
1167
+ try:
1168
+ warm_key = self._decode_graph_key(kv_cache, graph_width)
1169
+ if self._decode_warm_key != warm_key and warm_key not in self._decode_compiled_keys:
1170
+ self._warm_decode_graphs(
1171
+ host_inputs, kv_cache, graph_width=graph_width, initial_token_device=initial_token_device
1172
+ )
1173
+ self._restore_trace_inputs(host_inputs, include_page_table=True, token_device=initial_token_device)
1174
+ self._synchronize()
1175
+ token, current_pos, rotary_pos, page_table = self._trace_inputs
1176
+ self.model.bind_page_table(kv_cache, page_table)
1177
+
1178
+ model_trace_id = ttnn.begin_trace_capture(self.mesh_device, cq_id=0)
1179
+ model_open = True
1180
+ with self.model.decode_width_scope(graph_width):
1181
+ logits = self.model.decode_forward_from_ttnn_inputs(
1182
+ token, current_pos, rotary_position=rotary_pos, kv_cache=kv_cache, advance_position=True
1183
+ )
1184
+ ttnn.end_trace_capture(self.mesh_device, model_trace_id, cq_id=0)
1185
+ model_open = False
1186
+ self._synchronize()
1187
+
1188
+ sampling_trace_id = ttnn.begin_trace_capture(self.mesh_device, cq_id=0)
1189
+ sampling_open = True
1190
+ sampled = self._sample_device(logits, tt_out_tok=token)
1191
+ ttnn.end_trace_capture(self.mesh_device, sampling_trace_id, cq_id=0)
1192
+ sampling_open = False
1193
+ self._synchronize()
1194
+ except Exception:
1195
+ if sampling_open:
1196
+ ttnn.end_trace_capture(self.mesh_device, sampling_trace_id, cq_id=0)
1197
+ if model_open:
1198
+ ttnn.end_trace_capture(self.mesh_device, model_trace_id, cq_id=0)
1199
+ for trace_id in (sampling_trace_id, model_trace_id):
1200
+ if trace_id is not None:
1201
+ try:
1202
+ ttnn.release_trace(self.mesh_device, trace_id)
1203
+ except Exception:
1204
+ pass
1205
+ raise
1206
+
1207
+ self._trace_model_id = model_trace_id
1208
+ self._trace_sampling_id = sampling_trace_id
1209
+ self._trace_logits = logits
1210
+ self._trace_sampled = sampled
1211
+ self._trace_kv_cache = kv_cache
1212
+ self._trace_page_table_snapshot = self._page_table_to_torch(host_inputs[3]).clone()
1213
+ self._trace_active_batch = active_batch
1214
+ self._trace_graph_width = graph_width
1215
+ self.trace_stats["captures"] += 1
1216
+ self._restore_trace_inputs(host_inputs, include_page_table=True, token_device=initial_token_device)
1217
+ self._synchronize()
1218
+
1219
+ def _refresh_trace_state(
1220
+ self, host_inputs, kv_cache, *, graph_width: int, active_batch: int, initial_token_device=None
1221
+ ) -> None:
1222
+ new_page_table = self._page_table_to_torch(host_inputs[3])
1223
+ shape_changed = (
1224
+ self._trace_page_table_snapshot is not None
1225
+ and new_page_table.shape != self._trace_page_table_snapshot.shape
1226
+ )
1227
+ if self._trace_model_id is not None and (
1228
+ kv_cache is not self._trace_kv_cache
1229
+ or graph_width != self._trace_graph_width
1230
+ or active_batch != self._trace_active_batch
1231
+ or shape_changed
1232
+ ):
1233
+ self._release_decode_traces()
1234
+ if self._trace_model_id is None:
1235
+ self._capture_decode_traces(
1236
+ host_inputs,
1237
+ kv_cache,
1238
+ graph_width=graph_width,
1239
+ active_batch=active_batch,
1240
+ initial_token_device=initial_token_device,
1241
+ )
1242
+ return
1243
+ self._restore_trace_inputs(host_inputs, include_page_table=False, token_device=initial_token_device)
1244
+ if not torch.equal(new_page_table, self._trace_page_table_snapshot):
1245
+ ttnn.copy_host_to_device_tensor(host_inputs[3], self._trace_inputs[3])
1246
+ self._trace_page_table_snapshot = new_page_table.clone()
1247
+ self.trace_stats["page_table_host_copies"] += 1
1248
+
1249
+ def _refresh_persistent_page_table(self, page_table, kv_cache, *, active_batch: int) -> None:
1250
+ if self._trace_model_id is None:
1251
+ raise RuntimeError("decode trace is not initialized")
1252
+ if kv_cache is not self._trace_kv_cache:
1253
+ raise RuntimeError("KV-cache identity changed; initialize a new trace")
1254
+ if active_batch != self._trace_active_batch:
1255
+ raise RuntimeError("fixed active slots changed; initialize a new trace")
1256
+ if page_table is None:
1257
+ return
1258
+ new_page_table = self._normalise_page_table(page_table, active_batch, width=self._trace_graph_width)
1259
+ if torch.equal(new_page_table, self._trace_page_table_snapshot):
1260
+ return # unchanged page table costs zero host copies
1261
+ ttnn.copy_host_to_device_tensor(
1262
+ self._replicated_host_tensor(new_page_table, dtype=ttnn.int32), self._trace_inputs[3]
1263
+ )
1264
+ self._trace_page_table_snapshot = new_page_table.clone()
1265
+ self.trace_stats["page_table_host_copies"] += 1
1266
+
1267
+ def _copy_forced_tokens(self, tokens: torch.Tensor) -> None:
1268
+ """Teacher forcing: overwrite the fed-back token, keep everything else."""
1269
+ values = tokens.reshape(-1).to(torch.int64)
1270
+ if values.numel() != self._trace_active_batch:
1271
+ raise ValueError(f"expected {self._trace_active_batch} forced tokens, got {values.numel()}")
1272
+ host = torch.zeros(SAMPLING_SLOTS, dtype=torch.int32)
1273
+ host[: values.numel()] = values.to(torch.int32)
1274
+ ttnn.copy_host_to_device_tensor(
1275
+ self._replicated_host_tensor(host.reshape(1, 1, 1, SAMPLING_SLOTS), dtype=ttnn.uint32),
1276
+ self._trace_inputs[0],
1277
+ )
1278
+ self.trace_stats["token_host_copies"] += 1
1279
+
1280
+ def _replay_split_sampling(self):
1281
+ ttnn.execute_trace(self.mesh_device, self._trace_model_id, cq_id=0, blocking=False)
1282
+ ttnn.execute_trace(self.mesh_device, self._trace_sampling_id, cq_id=0, blocking=False)
1283
+ self.trace_stats["replays"] += 1
1284
+ return self._trace_sampled
1285
+
1286
+ def _release_decode_traces(self) -> None:
1287
+ released = self._trace_model_id is not None or self._trace_sampling_id is not None
1288
+ for trace_id in (self._trace_model_id, self._trace_sampling_id):
1289
+ if trace_id is not None:
1290
+ ttnn.release_trace(self.mesh_device, trace_id)
1291
+ if released:
1292
+ self.trace_stats["releases"] += 1
1293
+ self._trace_model_id = None
1294
+ self._trace_sampling_id = None
1295
+ self._trace_inputs = None
1296
+ self._trace_logits = None
1297
+ self._trace_sampled = None
1298
+ self._trace_kv_cache = None
1299
+ self._trace_page_table_snapshot = None
1300
+ self._trace_active_batch = None
1301
+ self._trace_graph_width = None
1302
+ self._trace_rotary_position = None
1303
+ self._decode_warm_key = None
1304
+
1305
+ def _ensure_decode_rope_capacity(self, required_len: int) -> None:
1306
+ """Grow the cos/sin tables for decode, releasing traces if they move.
1307
+
1308
+ ``Qwen3CoderModel.ensure_rope_capacity`` reallocates ``cos_table`` and
1309
+ ``sin_table`` when it grows them, and a captured trace holds the *old*
1310
+ tensor identities -- replaying it afterwards would gather from freed
1311
+ DRAM. ``prefill_forward`` and ``generate`` are safe because both release
1312
+ the decode traces before they call it; the low-level ``decode_forward``
1313
+ has no such release, so it does one here and only when the tables
1314
+ actually moved.
1315
+ """
1316
+ if required_len > self.model.max_cache_len:
1317
+ raise ValueError(f"decode horizon {required_len} exceeds the supported context {self.model.max_cache_len}")
1318
+ if self.model.ensure_rope_capacity(required_len):
1319
+ self._release_decode_traces()
1320
+
1321
+ def decode_forward(
1322
+ self,
1323
+ tokens: torch.Tensor | None,
1324
+ start_pos: torch.Tensor | None,
1325
+ *,
1326
+ page_table,
1327
+ kv_cache: Any,
1328
+ sampling_mode: str = "host",
1329
+ enable_trace: bool = False,
1330
+ active_batch: int | None = None,
1331
+ graph_width: int | None = None,
1332
+ decode_horizon: int | None = None,
1333
+ validate_page_coverage: bool = True,
1334
+ **kwargs: Any,
1335
+ ):
1336
+ """One decode step.
1337
+
1338
+ With ``enable_trace=True, sampling_mode="device"`` this is the delivered
1339
+ path: pass ``start_pos``/``page_table`` on the first step to install the
1340
+ trace, then call with ``tokens=None, start_pos=None, page_table=None``
1341
+ and the traces replay over persistent state -- the sampled token from
1342
+ step *N* is already the token input of step *N+1*, and both position
1343
+ tensors were advanced on device inside the model trace.
1344
+
1345
+ **Rotary capacity.** The cos/sin tables are sized lazily, and the traced
1346
+ loop advances ``rotary_position`` with ``ttnn.plus_one`` with nothing on
1347
+ device to clamp it, so a replay past the table length would gather
1348
+ out of range and silently rotate at the wrong position. Pass
1349
+ ``decode_horizon`` -- the highest position this trace will ever decode
1350
+ at, i.e. ``prompt_len + max_new_tokens - 1`` -- on the installing call
1351
+ and the tables are grown once, up front, to cover the whole run. Without
1352
+ it the tables are sized only for ``start_pos`` and a replay that would
1353
+ step past them raises instead of returning a wrong answer.
1354
+ ``generate`` sizes for its own horizon and never hits either path.
1355
+ """
1356
+ if sampling_mode not in {"host", "device"}:
1357
+ raise ValueError("sampling_mode must be 'host' or 'device'")
1358
+ caches = self._ensure_kv_cache() if kv_cache is None else kv_cache
1359
+ inferred = self._trace_active_batch if tokens is None else int(tokens.numel())
1360
+ active_batch = inferred if active_batch is None else int(active_batch)
1361
+ if active_batch is None or not 1 <= active_batch <= self.batch:
1362
+ raise ValueError(f"active_batch must be in [1,{self.batch}]")
1363
+ if tokens is not None and tokens.numel() != active_batch:
1364
+ raise ValueError("tokens do not match active_batch")
1365
+ if start_pos is not None and start_pos.numel() != active_batch:
1366
+ raise ValueError("start_pos does not match active_batch")
1367
+ # ``graph_width`` is how many rows the captured decode graph has;
1368
+ # ``active_batch`` is how many of them the caller is filling. They are
1369
+ # the same on the shipped path. A caller that has compacted its live
1370
+ # requests into rows ``0..active_batch-1`` may ask for a narrower graph,
1371
+ # which is the whole point of ``doc/batch_scaling``: expert, router and
1372
+ # SDPA cost is paid per row *configured*, so the only way to stop paying
1373
+ # for 32 rows when one is live is to capture a graph that has fewer.
1374
+ if graph_width is None:
1375
+ graph_width = self._trace_graph_width if start_pos is None and self._trace_graph_width else self.batch
1376
+ graph_width = int(graph_width)
1377
+ if not active_batch <= graph_width <= self.batch:
1378
+ raise ValueError(f"graph_width must be in [{active_batch},{self.batch}], got {graph_width}")
1379
+
1380
+ if enable_trace and sampling_mode == "device":
1381
+ if start_pos is not None:
1382
+ if page_table is None:
1383
+ raise ValueError("initial trace state requires positions and page_table")
1384
+ highest = int(start_pos.reshape(-1).max().item())
1385
+ horizon = highest + 1 if decode_horizon is None else int(decode_horizon)
1386
+ if horizon < highest + 1:
1387
+ raise ValueError("decode_horizon is below the requested start_pos")
1388
+ self._ensure_decode_rope_capacity(horizon)
1389
+ page_host = self._normalise_page_table(page_table, active_batch, width=graph_width)
1390
+ if validate_page_coverage:
1391
+ self._validate_page_coverage(page_host, start_pos, active_batch)
1392
+ initial_token_device = self._prefill_sampled if tokens is None else None
1393
+ host_tokens = torch.zeros(active_batch, dtype=torch.long) if tokens is None else tokens
1394
+ host_inputs = self._prepare_decode_host_inputs(host_tokens, start_pos, page_host, width=graph_width)
1395
+ self._refresh_trace_state(
1396
+ host_inputs,
1397
+ caches,
1398
+ graph_width=graph_width,
1399
+ active_batch=active_batch,
1400
+ initial_token_device=initial_token_device,
1401
+ )
1402
+ # The installing call also replays once, at ``highest``.
1403
+ self._trace_rotary_position = highest
1404
+ else:
1405
+ self._refresh_persistent_page_table(page_table, caches, active_batch=active_batch)
1406
+ if tokens is not None:
1407
+ self._copy_forced_tokens(tokens)
1408
+ if self._trace_rotary_position is not None:
1409
+ # ``ttnn.plus_one`` already moved the device tensor on the
1410
+ # previous replay; this replay gathers at that position.
1411
+ self._trace_rotary_position += 1
1412
+ if self._trace_rotary_position >= self.model.rope_cache_len:
1413
+ raise RuntimeError(
1414
+ f"decode position {self._trace_rotary_position} is past the rotary table "
1415
+ f"({self.model.rope_cache_len} entries). ``ttnn.embedding`` would gather out of "
1416
+ "range inside the replayed trace and rotate at a wrong position without "
1417
+ "raising. Re-install the trace with decode_horizon= set to the highest "
1418
+ "position this run will reach."
1419
+ )
1420
+ return self._replay_split_sampling()
1421
+
1422
+ if tokens is None or start_pos is None or page_table is None:
1423
+ raise ValueError("eager/host decode requires tokens, start_pos and page_table")
1424
+ self._release_decode_traces_before_allocating()
1425
+ # Eager decode gathers cos/sin at ``start_pos`` too, and holds no trace
1426
+ # by this point, so the tables can simply be grown to fit.
1427
+ self._ensure_decode_rope_capacity(int(start_pos.reshape(-1).max().item()) + 1)
1428
+ page_host = self._normalise_page_table(page_table, active_batch)
1429
+ if validate_page_coverage:
1430
+ self._validate_page_coverage(page_host, start_pos, active_batch)
1431
+ host_inputs = self._prepare_decode_host_inputs(tokens, start_pos, page_host)
1432
+ device_inputs = [
1433
+ ttnn.to_device(tensor, self.mesh_device, memory_config=ttnn.DRAM_MEMORY_CONFIG) for tensor in host_inputs
1434
+ ]
1435
+ self.model.bind_page_table(caches, device_inputs[3])
1436
+ logits = self.model.decode_forward_from_ttnn_inputs(
1437
+ device_inputs[0],
1438
+ device_inputs[1],
1439
+ rotary_position=device_inputs[2],
1440
+ kv_cache=caches,
1441
+ advance_position=False,
1442
+ )
1443
+ if sampling_mode == "device":
1444
+ with self._penalties_suspended():
1445
+ sampled = self._sample_device(logits)
1446
+ ttnn.deallocate(logits, True)
1447
+ return sampled
1448
+ host = self.model.gather_logits_to_torch(logits, valid_rows=active_batch)[0, 0]
1449
+ ttnn.deallocate(logits, True)
1450
+ return host
1451
+
1452
+ # -- high level -----------------------------------------------------------
1453
+
1454
+ def _generate_host_compat(
1455
+ self, prompt_token_ids: list[int], max_new_tokens: int, *, next_input: Optional[NextInputFn]
1456
+ ) -> list[int]:
1457
+ """Explicit host-sampling compatibility mode. Never a measured path."""
1458
+ self._release_decode_traces()
1459
+ kv_cache = self._ensure_kv_cache()
1460
+ horizon = len(prompt_token_ids) + max_new_tokens - 1
1461
+ page_host = self.make_page_table([horizon])
1462
+ logits = self.prefill_forward(
1463
+ torch.tensor([prompt_token_ids]),
1464
+ page_table=page_host,
1465
+ kv_cache=kv_cache,
1466
+ prompt_lens=[len(prompt_token_ids)],
1467
+ sampling_mode="host",
1468
+ )
1469
+ predicted = int(logits[0, 0].argmax().item())
1470
+ outputs: list[int] = []
1471
+ for step in range(max_new_tokens):
1472
+ outputs.append(predicted)
1473
+ next_token = next_input(step, predicted) if next_input is not None else predicted
1474
+ if step + 1 == max_new_tokens:
1475
+ break
1476
+ decoded = self.decode_forward(
1477
+ torch.tensor([[next_token]]),
1478
+ torch.tensor([len(prompt_token_ids) + step]),
1479
+ page_table=page_host,
1480
+ kv_cache=kv_cache,
1481
+ sampling_mode="host",
1482
+ enable_trace=False,
1483
+ )
1484
+ predicted = int(decoded[0].argmax().item())
1485
+ return outputs
1486
+
1487
+ def generate(
1488
+ self,
1489
+ prompt_token_ids: list[int],
1490
+ max_new_tokens: int,
1491
+ *,
1492
+ next_input: Optional[NextInputFn] = None,
1493
+ enable_trace: bool = True,
1494
+ sampling_mode: str = "device",
1495
+ stop_on_eos: bool = False,
1496
+ top_k=1,
1497
+ top_p=0.0,
1498
+ temperature=1.0,
1499
+ **kwargs: Any,
1500
+ ) -> list[int]:
1501
+ """Prefill, then loop the traced split-sampling decode path."""
1502
+ if not prompt_token_ids or max_new_tokens < 1:
1503
+ return []
1504
+ horizon = len(prompt_token_ids) + max_new_tokens - 1
1505
+ if horizon > self.model.max_cache_len:
1506
+ raise ValueError("prompt plus requested output exceeds the supported context")
1507
+ self._release_decode_traces_before_allocating()
1508
+ self.model.ensure_rope_capacity(horizon)
1509
+ if sampling_mode == "host":
1510
+ return self._generate_host_compat(prompt_token_ids, max_new_tokens, next_input=next_input)
1511
+ if sampling_mode != "device":
1512
+ raise ValueError("sampling_mode must be 'device' or 'host'")
1513
+ if not enable_trace and max_new_tokens > 1:
1514
+ raise ValueError("the optimized token-out path requires enable_trace=True")
1515
+ self.set_sampling_params(top_k=top_k, top_p=top_p, temperature=temperature, active_batch=1)
1516
+
1517
+ kv_cache = self._ensure_kv_cache()
1518
+ page_host = self.make_page_table([horizon])
1519
+ sampled = self.prefill_forward(
1520
+ torch.tensor([prompt_token_ids]),
1521
+ page_table=page_host,
1522
+ kv_cache=kv_cache,
1523
+ prompt_lens=[len(prompt_token_ids)],
1524
+ sampling_mode="device",
1525
+ )
1526
+ predicted = int(self._sampled_to_torch(sampled)[0].item())
1527
+ outputs: list[int] = []
1528
+ for step in range(max_new_tokens):
1529
+ outputs.append(predicted)
1530
+ forced = next_input(step, predicted) if next_input is not None else predicted
1531
+ if step + 1 == max_new_tokens:
1532
+ break
1533
+ if stop_on_eos and next_input is None and predicted == self.tokenizer.eos_token_id:
1534
+ break
1535
+ initial = step == 0
1536
+ sampled = self.decode_forward(
1537
+ (torch.tensor([[forced]]) if next_input is not None else None),
1538
+ torch.tensor([len(prompt_token_ids)]) if initial else None,
1539
+ page_table=page_host if initial else None,
1540
+ kv_cache=kv_cache,
1541
+ sampling_mode="device",
1542
+ enable_trace=True,
1543
+ active_batch=1,
1544
+ decode_horizon=horizon,
1545
+ )
1546
+ predicted = int(self._sampled_to_torch(sampled)[0].item())
1547
+ return outputs
1548
+
1549
+ def reset(self) -> None:
1550
+ """Wipe per-prompt state, keeping weights, buffers and program cache."""
1551
+ if self._trace_model_id is not None or self._trace_sampling_id is not None:
1552
+ self._synchronize()
1553
+ self._release_decode_traces()
1554
+ if self._kv_cache is not None:
1555
+ self.model.reset_kv_cache(self._kv_cache)
1556
+ empty = torch.full((self.batch, self.pages_per_user), -1, dtype=torch.int32)
1557
+ self._copy_host(empty, self._prefill_page_table, dtype=ttnn.int32)
1558
+ self._copy_host(
1559
+ torch.zeros((1, 1, 1, SAMPLING_SLOTS), dtype=torch.int32), self._prefill_sampled, dtype=ttnn.uint32
1560
+ )
1561
+ for width, pool in self._width_pools.items():
1562
+ token, current_pos, rotary_pos, page_table = pool
1563
+ self._copy_host(torch.zeros((1, 1, 1, SAMPLING_SLOTS), dtype=torch.int32), token, dtype=ttnn.uint32)
1564
+ self._copy_host(torch.full((width,), -1, dtype=torch.int32), current_pos, dtype=ttnn.int32)
1565
+ self._copy_host(torch.zeros((1, width), dtype=torch.int32), rotary_pos, dtype=ttnn.uint32)
1566
+ self._copy_host(empty[:width], page_table, dtype=ttnn.int32)
1567
+ self._trace_rotary_position = None
1568
+ self.trace_stats["resets"] += 1
1569
+ self._synchronize()
1570
+
1571
+ def teardown(self) -> None:
1572
+ self._release_decode_traces()
1573
+ #: Not cleared by ``_release_decode_traces`` on purpose -- a released
1574
+ #: trace leaves its programs compiled, which is the whole point of the
1575
+ #: set. Cleared here because teardown deallocates the KV cache, and the
1576
+ #: key holds ``id(kv_cache)``: a later allocation could land on the same
1577
+ #: address with a different shape and falsely claim to be warm.
1578
+ self._decode_compiled_keys.clear()
1579
+ if self._kv_cache is not None:
1580
+ for cache in self._kv_cache:
1581
+ ttnn.deallocate(cache.k, True)
1582
+ ttnn.deallocate(cache.v, True)
1583
+ self._kv_cache = None
1584
+
1585
+
1586
+ def _resolve_snapshot(model_path: str | Path | None = None) -> Path:
1587
+ if model_path is not None:
1588
+ path = Path(model_path)
1589
+ if not path.exists():
1590
+ raise FileNotFoundError(path)
1591
+ return path
1592
+ hf_home = Path(os.getenv("HF_HOME", Path.home() / ".cache" / "huggingface"))
1593
+ snapshot = hf_home / "hub" / "models--Qwen--Qwen3-Coder-30B-A3B-Instruct" / "snapshots" / HF_REVISION
1594
+ if (snapshot / "model.safetensors.index.json").is_file():
1595
+ return snapshot
1596
+ from huggingface_hub import snapshot_download
1597
+
1598
+ return Path(snapshot_download(HF_MODEL_ID, revision=HF_REVISION))
1599
+
1600
+
1601
+ def build_generator(model_dir: str | Path, mesh_device, **kwargs) -> Generator:
1602
+ """Readiness discovery factory. See ``models/common/readiness_check/contract.py``.
1603
+
1604
+ **This is the construction path the precision config has to reach.** The
1605
+ readiness runners, the qualitative suite and (later) vLLM all arrive here
1606
+ and none of them can pass a Python object, so ``precision`` is accepted as a
1607
+ kwarg *and* read from ``QWEN3_PRECISION_CONFIG`` in the environment as a
1608
+ path to a ``selected_precision_config.json``. Unset -- which is every run to
1609
+ date -- means ``DEFAULT_PRECISION``, i.e. the shipped policy, so this
1610
+ default is the one the stage-07 goal asks for rather than a JSON field that
1611
+ hard-coded model code ignores.
1612
+ """
1613
+ snapshot = _resolve_snapshot(kwargs.pop("model_path", os.getenv("QWEN3_CODER_30B_MODEL_PATH")))
1614
+ tokenizer = AutoTokenizer.from_pretrained(snapshot)
1615
+ max_batch_size = int(kwargs.pop("max_batch_size", 1))
1616
+ max_context_len = int(kwargs.pop("max_context_len", MAX_CONTEXT))
1617
+ override_num_layers = kwargs.pop("override_num_layers", None)
1618
+ num_layers = NUM_LAYERS if override_num_layers is None else int(override_num_layers)
1619
+ page_block_size = int(kwargs.pop("page_block_size", 32))
1620
+ rope_cache_len = int(kwargs.pop("rope_cache_len", 8192))
1621
+ precision = kwargs.pop("precision", os.getenv("QWEN3_PRECISION_CONFIG") or None)
1622
+ if kwargs:
1623
+ raise TypeError(f"unsupported build_generator kwargs: {sorted(kwargs)}")
1624
+ model = Qwen3CoderModel.from_checkpoint(
1625
+ snapshot,
1626
+ mesh_device=mesh_device,
1627
+ max_batch_size=max_batch_size,
1628
+ max_cache_len=max_context_len,
1629
+ num_layers=num_layers,
1630
+ page_block_size=page_block_size,
1631
+ rope_cache_len=rope_cache_len,
1632
+ precision=precision,
1633
+ )
1634
+ return Qwen3CoderGenerator(model, tokenizer)
1635
+
1636
+
1637
+ __all__ = ["Qwen3CoderGenerator", "build_generator"]
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator_vllm.py ADDED
@@ -0,0 +1,1568 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """vLLM serving adapter for Qwen3-Coder-30B-A3B-Instruct on 4 Blackhole dies.
5
+
6
+ This file is **translation only**. Every device operation it causes is a call
7
+ into ``tt/generator.py``'s low-level surface -- ``prefill_forward`` /
8
+ ``decode_forward`` / ``set_sampling_params`` / ``configure_paging`` -- which is
9
+ the same surface the standalone readiness runners drive. There is no model code
10
+ here, no second sampler, no host argmax, no full-logits readback on the measured
11
+ path and no Python readback/writeback token-feedback loop.
12
+
13
+ The three things it actually has to reconcile
14
+ ---------------------------------------------
15
+
16
+ **1. Who owns the cache.** Standalone, the generator allocates its own paged
17
+ cache and builds its own page tables. Serving, vLLM owns both: it picks the
18
+ block size and the block count, calls ``allocate_kv_cache`` once, and hands that
19
+ exact cache and a ``[batch, max_num_blocks_per_req]`` block table into every
20
+ forward. ``allocate_kv_cache`` therefore installs vLLM's geometry through
21
+ ``Qwen3CoderGenerator.configure_paging`` and allocates the cache through
22
+ ``Qwen3CoderModel.allocate_kv_cache(num_blocks=...)``; the generator's own
23
+ ``_kv_cache`` is never created, so no standalone-cache assumption can survive.
24
+
25
+ **2. Who owns the token.** The generator's traced decode path writes the sampled
26
+ token straight back into the persistent decode token input with ``tt_out_tok``
27
+ and advances ``current_pos``/``rotary_position`` on device with
28
+ ``ttnn.plus_one``. So after step *N* the **device**, not the host, holds the
29
+ token and position step *N+1* needs. vLLM's scheduler also tracks them, and
30
+ under ``--async-scheduling`` its copy is a step behind. On a steady decode step
31
+ this adapter therefore passes *nothing* -- ``decode_forward(None, None,
32
+ page_table=..., ...)`` replays the two traces over state the device already
33
+ owns. Only when vLLM says the slot layout changed (``reset_batch``, or the step
34
+ right after a prefill) does it reinstall host state, and even then it keeps the
35
+ device's token and position for slots that are merely continuing
36
+ (``_merge_scheduler_view``), so an async-ahead scheduler cannot re-decode a
37
+ position or feed back a stale token.
38
+
39
+ **3. Who owns sampling.** vLLM would rather hand us logits. It only stops doing
40
+ that when the model declares ``supports_sample_on_device`` and the server runs
41
+ with ``sample_on_device_mode: all`` -- then it sends per-row
42
+ ``(temperature, top_k, top_p)`` and expects token ids back. That is the measured
43
+ path and it is exactly the full model's canonical split sampling: greedy takes
44
+ ``Qwen3CoderModel.sample_greedy_argmax``, anything sampled takes
45
+ ``sample_split``, both are ``_WatcherCleanSampling1D``, both are traced, both
46
+ write ``tt_out_tok``. The TT plugin still routes a few request shapes to host
47
+ sampling on its own (logprobs on a mesh that is not 8 or 32 dies, ``min_p``,
48
+ ``bad_words``, ``logit_bias``, structured output) -- for those it passes
49
+ ``sampling_params=None`` and wants logits. That is served by the generator's
50
+ pre-existing, explicit ``sampling_mode="host"`` compatibility mode. It is opt-in
51
+ per request by vLLM, never used for a performance number, and it never displaces
52
+ the traced path. The eager decode does release the captured decode traces (it
53
+ allocates, so ``generator.decode_forward`` calls
54
+ ``_release_decode_traces_before_allocating``); the adapter therefore sets
55
+ ``_needs_decode_install`` and the next device-sampled step re-captures through
56
+ ``_refresh_trace_state`` rather than replaying a stale trace.
57
+ """
58
+
59
+ from __future__ import annotations
60
+
61
+ import json
62
+ import math
63
+ import os
64
+ from collections import deque
65
+ from pathlib import Path
66
+ from typing import Any, Optional
67
+
68
+ import torch
69
+ from loguru import logger
70
+
71
+ import ttnn
72
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.functional_decoder import sdpa_chunk_size
73
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator import SAMPLING_SLOTS, Qwen3CoderGenerator, build_generator
74
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.model import HF_MODEL_ID, MAX_CONTEXT
75
+
76
+ #: This port's directory. Everything below reads its policy from here, not from
77
+ #: a vLLM flag, so serving cannot silently run a different model than readiness.
78
+ MODEL_DIR = Path(__file__).resolve().parents[1]
79
+
80
+ #: The datatype-sweep selection (stage 07). ``build_generator`` accepts a path
81
+ #: and threads it into ``Qwen3CoderModel``, so weight groups, activation dtype,
82
+ #: CCL dtype, KV-cache dtype, compute fidelities and layer exceptions all come
83
+ #: from this one file on the serving path exactly as they do on the readiness
84
+ #: path. ``QWEN3_PRECISION_CONFIG`` still overrides it, for sweeps.
85
+ SELECTED_PRECISION_CONFIG = MODEL_DIR / "config" / "selected_precision_config.json"
86
+
87
+ #: ``config/context_contract.json`` is the single source of truth for served
88
+ #: context. ``get_max_tokens_all_users`` and ``initialize_vllm_model`` both read
89
+ #: it rather than trusting a CLI value, so a ``--max-model-len`` above the
90
+ #: recorded capability fails loudly instead of serving a quietly-clipped model.
91
+ CONTEXT_CONTRACT = MODEL_DIR / "config" / "context_contract.json"
92
+
93
+ #: Whether a serving prefill may run while the decode traces stay captured.
94
+ #: **Off by default, and that is a measurement, not caution.** Keeping them alive
95
+ #: hangs the NoC: prefill's collectives share the sampler's persistent CCL
96
+ #: buffers and semaphores with the captured graph, and after a few admissions a
97
+ #: replay waits on a semaphore value an eager collective already consumed.
98
+ #: ``doc/vllm_integration/triage/tt-triage-preserve-traces-hang.txt`` is that
99
+ #: hang, caught with ``dump_running_operations`` reporting ``NOC0 CB0..3 active
100
+ #: (0xFFFFFFFF). NoC is likely hung.`` on device 0. Releasing on prefill --
101
+ #: which is what ``Qwen3CoderGenerator.prefill_forward`` has always done -- makes
102
+ #: the next capture re-establish that state, and
103
+ #: ``Qwen3CoderGenerator._decode_compiled_keys`` keeps the re-capture from
104
+ #: paying for a second eager warm pass. Set ``QWEN3_VLLM_PRESERVE_DECODE_TRACES=1``
105
+ #: to reproduce the hang.
106
+ PRESERVE_DECODE_TRACES = os.getenv("QWEN3_VLLM_PRESERVE_DECODE_TRACES", "0") not in ("0", "", "false", "no")
107
+
108
+
109
+ #: Prefix caching is ON by default from phase 3. This is no longer a feature gate
110
+ #: but a kill switch: ``QWEN3_PREFIX_CACHING=0`` restores the pre-phase-3 refusal
111
+ #: without touching ``model_capabilities``, so the two can be reverted separately.
112
+ _PREFIX_CACHING_ENABLED = os.getenv("QWEN3_PREFIX_CACHING", "1") not in ("0", "", "false", "no")
113
+
114
+
115
+ def _supported_context() -> int:
116
+ try:
117
+ contract = json.loads(CONTEXT_CONTRACT.read_text())
118
+ except (OSError, ValueError):
119
+ return MAX_CONTEXT
120
+ return int(contract.get("current_supported_context") or MAX_CONTEXT)
121
+
122
+
123
+ def _reorder_history(history, order, picks):
124
+ """Reorder a vLLM ``[rows, L]`` token history into graph-row order.
125
+
126
+ Type-preserving **on purpose**. ``Generator._row_token_ids`` consumes these
127
+ with ``torch.as_tensor(history)``, which raises ``TypeError: only integer
128
+ tensors of a single element can be converted to an index`` when handed a
129
+ *list of 1-D tensors*. Rebuilding a tensor history with a list comprehension
130
+ therefore turns a working penalised decode into a crash -- and only when the
131
+ width ladder is active, since with the ladder off ``order`` is ``None`` and
132
+ the history is passed through untouched. That is exactly the shape of bug
133
+ that reaches production: invisible on the default path, fatal on the new one.
134
+
135
+ Anything that supports fancy indexing (torch tensors, numpy arrays) is
136
+ indexed so its type survives; genuine python sequences keep the list
137
+ comprehension, which is already what ``_row_token_ids`` expects from them.
138
+ """
139
+ if history is None:
140
+ return None
141
+ if isinstance(history, torch.Tensor):
142
+ return history[order.to(torch.long)]
143
+ if hasattr(history, "__getitem__") and hasattr(history, "dtype"): # numpy & friends
144
+ return history[list(picks)]
145
+ return [history[i] for i in picks]
146
+
147
+
148
+ def _as_int_list(values, length: int, default) -> list:
149
+ """vLLM hands per-row sampling params over as python lists; normalise them."""
150
+ if values is None:
151
+ return [default] * length
152
+ # vLLM hands prompt_lens/start_pos over as numpy arrays and sampling params
153
+ # as python lists, and the plugin builds some fields from torch tensors.
154
+ if hasattr(values, "tolist") and not isinstance(values, (list, tuple)):
155
+ values = values.tolist()
156
+ if not isinstance(values, (list, tuple)):
157
+ values = [values] * length
158
+ out = list(values)[:length]
159
+ out.extend([default] * (length - len(out)))
160
+ return out
161
+
162
+
163
+ class Qwen3CoderForCausalLM:
164
+ """The class vLLM instantiates. Registered as ``TTQwen3MoeForCausalLM``.
165
+
166
+ ``Qwen3MoeForCausalLM`` is the architecture string in this checkpoint's
167
+ ``config.json``; the TT plugin registers every model under a ``TT`` prefix.
168
+ """
169
+
170
+ #: Read off the *class* by ``vllm_tt_plugin.platform.check_and_update_config``
171
+ #: before anything is instantiated.
172
+ #:
173
+ #: * ``supports_sample_on_device`` -- the full model's traced split sampling
174
+ #: is the measured token-out path; without this flag the readiness runner's
175
+ #: ``sample_on_device_mode: all`` is a hard config error.
176
+ #: * ``supports_async_decode`` -- ``decode_forward(read_from_device=False)``
177
+ #: returns device handles, ``read_decode_output(async_read=True)`` does the
178
+ #: deferred read and records an event, ``process_decode_output_host`` does
179
+ #: host formatting only. This also gates ``--async-scheduling``, which is
180
+ #: safe here because ``_merge_scheduler_view`` prefers the device's token
181
+ #: and position over an async-ahead scheduler's.
182
+ #: * ``supports_prefix_caching`` -- **False**. Phase 3 REVERTED it: with caching
183
+ #: on, cold-vs-warm greedy output matched on only 1 of 10 prompts, while the
184
+ #: same test with --no-enable-prefix-caching matched 10/10. See
185
+ #: doc/prefix_caching/probes/phase3_cold_warm_rate.json and
186
+ #: phase3_control_no_prefix_caching.json. The adapter wiring below is correct
187
+ #: and stays; the flag must not go back to True until that gap is closed.
188
+ #: (Historical note, kept because the wiring depends on it:)
189
+ #: vLLM then sends a non-zero ``start_pos`` (= ``num_computed_tokens``) with
190
+ #: the FULL prompt and the FULL ``prompt_lens``; the model slices the suffix
191
+ #: itself, matching tt_transformers' ``tokens[i, num_cached:seq_len]``.
192
+ #: Two vLLM invariants make our generator-side guards exact rather than
193
+ #: defensive, both READ OFF vllm rather than assumed:
194
+ #: - ``max_cache_hit_length = request.num_tokens - 1`` (kv_cache_manager.py)
195
+ #: so a full hit still recomputes the last token: ``start < prompt_len``.
196
+ #: - cache hits are whole blocks and ``allocate_slots`` requires
197
+ #: block-aligned ``num_computed_tokens``: ``start % 32 == 0``.
198
+ #: Chunked prefill is force-disabled by the plugin platform, so prefix
199
+ #: caching is the only source of a non-zero ``start_pos``.
200
+ #: Kill switch: ``QWEN3_PREFIX_CACHING=0`` restores the old refusal.
201
+ model_capabilities = {
202
+ # QUALITY-GATE EDIT (doc/prefix_caching/QUALITY_BAR.md): flipped True to
203
+ # run the caching-ON arm. Revert this single line to False if the gate
204
+ # fails. See doc/prefix_caching/quality_gate/.
205
+ "supports_prefix_caching": True,
206
+ "supports_async_decode": True,
207
+ "supports_sample_on_device": True,
208
+ }
209
+
210
+ # -- construction ---------------------------------------------------------
211
+
212
+ @classmethod
213
+ def initialize_vllm_model(
214
+ cls,
215
+ hf_config,
216
+ mesh_device,
217
+ max_batch_size,
218
+ max_seq_len: int | None = None,
219
+ n_layers: int | None = None,
220
+ tt_data_parallel: int = 1,
221
+ optimizations: str | None = None,
222
+ **kwargs: Any,
223
+ ) -> "Qwen3CoderForCausalLM":
224
+ if int(tt_data_parallel) != 1:
225
+ raise ValueError(
226
+ f"tt_data_parallel={tt_data_parallel} is unsupported: this port occupies the whole "
227
+ "1x4 mesh with tensor parallelism, so there is no submesh left to replicate onto."
228
+ )
229
+ supported = _supported_context()
230
+ max_seq_len = supported if max_seq_len is None else int(max_seq_len)
231
+ if max_seq_len > supported:
232
+ raise ValueError(
233
+ f"--max-model-len {max_seq_len} exceeds the context recorded in {CONTEXT_CONTRACT} " f"({supported})."
234
+ )
235
+ max_batch_size = int(max_batch_size)
236
+ if not 1 <= max_batch_size <= SAMPLING_SLOTS:
237
+ raise ValueError(
238
+ f"max_num_seqs={max_batch_size} is outside [1,{SAMPLING_SLOTS}]. "
239
+ "nlp_create_qkv_heads_decode and ttnn.sampling both address 32 fixed user slots."
240
+ )
241
+
242
+ # Reduced serving target for the bring-up inner loop only: the same
243
+ # adapter, generator, registration, cache/page-table shapes, terminal
244
+ # norm/LM head, sampler and trace behaviour, with fewer copies of the one
245
+ # layer kind this model has. Never used for accuracy or performance
246
+ # evidence -- the final run leaves it unset and gets all 48 layers.
247
+ reduced = os.getenv("QWEN3_VLLM_NUM_LAYERS")
248
+ if reduced:
249
+ n_layers = int(reduced)
250
+ logger.warning(
251
+ "QWEN3_VLLM_NUM_LAYERS={} -- REDUCED serving target, bring-up inner loop only. "
252
+ "Do not report accuracy or performance from this server.",
253
+ n_layers,
254
+ )
255
+
256
+ precision = os.getenv("QWEN3_PRECISION_CONFIG") or str(SELECTED_PRECISION_CONFIG)
257
+ generator = build_generator(
258
+ MODEL_DIR,
259
+ mesh_device,
260
+ max_batch_size=max_batch_size,
261
+ max_context_len=max_seq_len,
262
+ # The traced decode loop advances ``rotary_position`` on device and
263
+ # nothing on device clamps it, so the cos/sin tables must already
264
+ # cover every position this server may be asked to serve. Sizing
265
+ # them here means no serving step can ever grow them -- growing
266
+ # reallocates, and a captured trace holds the old identities.
267
+ rope_cache_len=max_seq_len,
268
+ precision=precision,
269
+ **({} if n_layers is None else {"override_num_layers": int(n_layers)}),
270
+ )
271
+ # Logged after the generator exists so ``active_row_gating`` is read off
272
+ # the model that was actually built rather than re-parsed from the
273
+ # environment here. Every leg of an A/B then carries its own
274
+ # configuration in its own server log, instead of the two legs being
275
+ # distinguishable only by the work-log prose that says which was which.
276
+ logger.info(
277
+ "Qwen3-Coder-30B-A3B vLLM init: max_num_seqs={} max_model_len={} precision={} "
278
+ "active_row_gating={} optimizations={}",
279
+ max_batch_size,
280
+ max_seq_len,
281
+ precision,
282
+ generator.model.active_row_gating,
283
+ optimizations,
284
+ )
285
+ return cls(generator, max_model_len=max_seq_len, max_num_seqs=max_batch_size)
286
+
287
+ def __init__(self, generator: Qwen3CoderGenerator, *, max_model_len: int, max_num_seqs: int):
288
+ self.generator = generator
289
+ self.model = generator.model
290
+ self.mesh_device = generator.mesh_device
291
+ self.max_model_len = int(max_model_len)
292
+ self.max_num_seqs = int(max_num_seqs)
293
+
294
+ #: vLLM-owned cache; set by ``allocate_kv_cache`` and never re-created.
295
+ self.kv_cache: list | None = None
296
+ #: True until a decode step has installed host state into the trace.
297
+ #: Every prefill sets it, because prefill admits a new request into a
298
+ #: slot whose device token/position belong to whoever held it before.
299
+ self._needs_decode_install = True
300
+ #: Set by ``warmup_model_prefill`` so the plugin's two-phase warmup does
301
+ #: not repeat the prefill sweep (the plugin resets this itself).
302
+ self.already_warmed_up_prefill = False
303
+ #: Runtime-fallback bookkeeping, reported by ``serving_audit``.
304
+ self._audit = {
305
+ "device_sampled_decode_steps": 0,
306
+ "host_sampled_decode_steps": 0,
307
+ "device_sampled_prefills": 0,
308
+ "host_sampled_prefills": 0,
309
+ "decode_trace_installs": 0,
310
+ "top_k_clamped_requests": 0,
311
+ "penalised_decode_steps": 0,
312
+ "ignored_seed_requests": 0,
313
+ #: Steps on which vLLM actually took the async split -- i.e. called
314
+ #: ``read_decode_output(async_read=True)`` rather than reading the
315
+ #: device handle synchronously inside ``execute_model``. This is what
316
+ #: makes ``supports_async_decode`` a measurement instead of a claim;
317
+ #: see the one-time log line in ``read_decode_output``.
318
+ "async_decode_reads": 0,
319
+ "sync_decode_reads": 0,
320
+ }
321
+ self._warned: set[str] = set()
322
+
323
+ #: Decode graph widths this server may capture, ascending. Read from
324
+ #: ``QWEN3_DECODE_WIDTHS`` (comma-separated, e.g. ``1,8,32``); anything
325
+ #: above ``max_num_seqs`` is dropped and ``max_num_seqs`` is always
326
+ #: present, so the default -- unset -- is exactly the shipped single
327
+ #: fixed-width graph and nothing below changes behaviour.
328
+ #:
329
+ #: Why this exists: expert, router and paged-SDPA cost is paid per row
330
+ #: **configured**, not per row live (``doc/optimized_vllm/README.md``'s
331
+ #: control curve: 227.9 ms fixed + 1.28 ms x live_rows at 32 slots). The
332
+ #: only lever that removes the fixed term is a graph with fewer rows.
333
+ self._decode_widths = self._configured_widths()
334
+ #: Graph row -> vLLM slot for the live trace, or ``None`` when the graph
335
+ #: is full width and the mapping is the identity. Rewritten only on an
336
+ #: install, which is the only step on which the batch layout may change
337
+ #: (``model_runner.py`` sets ``reset_batch`` from a sticky
338
+ #: ``_decode_layout_changed_since_last_decode``).
339
+ self._compaction: torch.Tensor | None = None
340
+ #: One entry per decode forward that is still awaiting its host read: the
341
+ #: graph-row -> vLLM-slot mapping **that forward was issued with**.
342
+ #:
343
+ #: Why a queue and not just ``_compaction``. The un-permutation happens in
344
+ #: ``process_decode_output_host``, which under ``--async-scheduling`` does
345
+ #: not run in the same step as the forward that produced the tokens: the
346
+ #: forward returns a device handle and the host read happens later. A
347
+ #: mapping stored on the adapter can therefore be **rewritten by a later
348
+ #: install before the earlier step's tokens are scattered**, which would
349
+ #: put every token on the wrong slot -- silently, as a correctness bug.
350
+ #:
351
+ #: The plugin happens to order this safely today: only a layout change can
352
+ #: rewrite the mapping, and ``model_runner.py`` drains pending async decodes
353
+ #: whenever the layout changed. But that is an invariant of a *different*
354
+ #: repository, which this one must not modify and cannot pin with a test.
355
+ #: Pairing each output with the mapping its own forward used replaces
356
+ #: that dependency with a weaker and detectable one. It no longer relies
357
+ #: on drain-on-layout-change; it does still assume the plugin finalizes
358
+ #: decode steps in issue order and exactly once, which is an invariant of
359
+ #: the same foreign repository. The difference is that a violation now
360
+ #: raises (tag mismatch, underflow, or depth cap) instead of silently
361
+ #: scattering a step's tokens through another step's permutation.
362
+ self._pending_orders: deque = deque()
363
+ #: Monotonic id of the next decode forward to be issued, and of the next
364
+ #: output expected. They are compared on every pop: a queue that has
365
+ #: skipped or reordered an entry shows up as a tag mismatch rather than
366
+ #: as tokens quietly landing on the wrong requests.
367
+ self._next_issue_tag = 0
368
+ self._next_output_tag = 0
369
+ #: Hard ceiling on outstanding decode forwards. Async scheduling runs at
370
+ #: most a step or two ahead; anything approaching this is a leak, not
371
+ #: depth.
372
+ self._pending_orders_cap = 64
373
+ self._audit["narrow_decode_installs"] = 0
374
+ self._audit["decode_graph_width"] = self.max_num_seqs
375
+ #: Times the output path found no queued mapping. Must stay 0; a nonzero
376
+ #: value means forwards and host reads are not paired one-to-one. The
377
+ #: path raises rather than guessing -- applying the adapter's current
378
+ #: mapping here would be the exact mis-scatter the queue exists to
379
+ #: prevent -- so this counter records a raise, not a silent fallback.
380
+ self._audit["compaction_fifo_underflows"] = 0
381
+ self._audit["compaction_fifo_max_depth"] = 0
382
+
383
+ @property
384
+ def _compaction_enabled(self) -> bool:
385
+ """Whether the row-mapping queue is in use at all.
386
+
387
+ Derived from ``_decode_widths`` rather than cached, because probes and
388
+ tests rebind that list after construction to switch the ladder on and
389
+ off; a cached flag would go stale and silently disable the pairing.
390
+
391
+ With the ladder disabled there is no permutation to pair, so nothing is
392
+ pushed or popped and the shipped path gains neither the bookkeeping nor
393
+ its failure modes -- in particular it cannot raise the errors below.
394
+ """
395
+ return len(self._decode_widths) > 1
396
+
397
+ def _reset_pending_orders(self) -> None:
398
+ """Drop queued mappings whose outputs can no longer be read.
399
+
400
+ Called wherever the decode traces are released or replaced. A queued
401
+ entry refers to a forward whose sampled tokens live in the trace's
402
+ persistent output tensor; once that trace is gone the handle cannot be
403
+ read at all, so the entry is dead and keeping it would desync every
404
+ later pop. The tags are realigned rather than zeroed so the invariant
405
+ ("the n-th output pairs with the n-th forward") survives the reset.
406
+ """
407
+ self._pending_orders.clear()
408
+ self._next_output_tag = self._next_issue_tag
409
+
410
+ #: The ladder used when ``QWEN3_DECODE_WIDTHS`` is unset. Powers of two to
411
+ #: ``max_num_seqs``: each step runs in the narrowest graph that holds the
412
+ #: live rows, so the cost a user pays tracks occupancy instead of the slot
413
+ #: count the server was configured with.
414
+ #:
415
+ #: On by default because the alternative default is *known wrong*: a
416
+ #: ``max_num_seqs=32`` server decodes a single user at 4.3464 t/s/u against
417
+ #: 49.3636 with the ladder, and a deployment that simply does not set an
418
+ #: environment variable gets the slow one. That is the same failure shape as
419
+ #: a missing ``sample_on_device_mode`` -- a config key whose absence looks
420
+ #: like broken hardware rather than a default.
421
+ #:
422
+ #: ``QWEN3_DECODE_WIDTHS=32`` (or any single width equal to ``max_num_seqs``)
423
+ #: restores the previous fixed-width behaviour exactly.
424
+ DEFAULT_DECODE_WIDTHS = (1, 2, 4, 8, 16, 32)
425
+
426
+ def _configured_widths(self) -> list[int]:
427
+ """Decode graph widths this server may capture, ascending.
428
+
429
+ ``max_num_seqs`` is always present -- a graph can never be wider than
430
+ the slots the caller sends, and the full-width graph must exist as the
431
+ fallback -- and anything above it is dropped.
432
+ """
433
+ raw = os.getenv("QWEN3_DECODE_WIDTHS", "").strip()
434
+ source = raw.split(",") if raw else [str(w) for w in self.DEFAULT_DECODE_WIDTHS]
435
+ widths = {self.max_num_seqs}
436
+ for piece in source:
437
+ piece = piece.strip()
438
+ if not piece:
439
+ continue
440
+ value = int(piece)
441
+ if 1 <= value <= self.max_num_seqs:
442
+ widths.add(value)
443
+ return sorted(widths)
444
+
445
+ def _choose_width(self, live_rows: int) -> int:
446
+ """Narrowest configured graph that can hold ``live_rows`` requests."""
447
+ for width in self._decode_widths:
448
+ if width >= max(1, live_rows):
449
+ return width
450
+ return self.max_num_seqs
451
+
452
+ @staticmethod
453
+ def _compaction_order(host_positions: torch.Tensor, width: int, rows: int) -> torch.Tensor:
454
+ """Graph row -> vLLM slot, live slots first, then spare slots.
455
+
456
+ The live slots go to rows ``0..live-1`` in their original order; the
457
+ remaining graph rows are filled from the *inactive* slots so that every
458
+ graph row still names a distinct vLLM slot and therefore still carries a
459
+ real (zero-filled) page-table row. Those rows install ``current_pos =
460
+ -1``, which is the inactive sentinel the traced graph already relies on.
461
+
462
+ This is a permutation of **only** the three per-row inputs -- position,
463
+ rotary index and page-table row -- plus the token. No KV page moves: the
464
+ cache is reached exclusively through page-table entries, so a request's
465
+ pages are wherever its page-table row says they are, in whatever graph
466
+ row that row is installed.
467
+ """
468
+ live = torch.nonzero(host_positions >= 0, as_tuple=False).reshape(-1)
469
+ spare = torch.nonzero(host_positions < 0, as_tuple=False).reshape(-1)
470
+ order = torch.cat((live, spare))[:width]
471
+ if order.numel() < width: # fewer vLLM slots than graph rows: cannot happen
472
+ raise RuntimeError(f"cannot fill a {width}-row graph from {rows} slots")
473
+ return order.to(torch.int64)
474
+
475
+ def _warn_once(self, key: str, message: str) -> None:
476
+ if key not in self._warned:
477
+ self._warned.add(key)
478
+ logger.warning(message)
479
+
480
+ # -- scheduler sizing -----------------------------------------------------
481
+
482
+ @classmethod
483
+ def get_max_tokens_all_users(
484
+ cls,
485
+ model_name: str = "",
486
+ num_devices: int = 1,
487
+ tt_data_parallel: int = 1,
488
+ max_model_len: int | None = None,
489
+ max_num_seqs: int | None = None,
490
+ **kwargs: Any,
491
+ ) -> int:
492
+ """Total KV tokens vLLM may allocate blocks for, across all users.
493
+
494
+ The whole advertised context for one user. ``config/context_contract.json``
495
+ records 262144 as the supported context and the paged decode probe that
496
+ reached position 262143; this is what makes vLLM size enough blocks for
497
+ a single request to actually use it. At this port's 4 dies the KV cost is
498
+ 512 B per token per layer per die over 48 layers -- 24 KiB per token per
499
+ die, so 262144 tokens is 6.29 GiB of the 34.18 GiB each die reports.
500
+
501
+ The worker adds ``block_size * max_num_seqs`` of its own headroom on top
502
+ and converts to blocks, so nothing here needs to model that.
503
+ """
504
+ supported = _supported_context()
505
+ return supported if max_model_len is None else min(int(max_model_len), supported)
506
+
507
+ # -- vLLM-owned KV cache --------------------------------------------------
508
+
509
+ def allocate_kv_cache(self, kv_cache_shape, dtype, num_layers: int):
510
+ """Allocate the attention KV cache **for vLLM**, at vLLM's geometry.
511
+
512
+ ``kv_cache_shape`` is ``(num_blocks, num_kv_heads_per_device, block_size,
513
+ head_dim)``; the plugin has already divided the head count by the mesh
514
+ size, which for this port's 4 dies and 4 KV heads gives the 1 local head
515
+ per die the model expects. The block size is vLLM's, so it is installed
516
+ into the generator here -- before warmup, before any forward, before any
517
+ trace -- rather than assumed.
518
+
519
+ ``dtype`` is vLLM's torch dtype and is deliberately **not** used: the KV
520
+ dtype is part of the selected precision policy
521
+ (``kv_cache_dtype`` in ``selected_precision_config.json``) and serving
522
+ must not silently run a different one than the sweep measured.
523
+ """
524
+ num_blocks, kv_heads, block_size, head_dim = (int(v) for v in kv_cache_shape)
525
+ if num_layers != self.model.num_layers:
526
+ if not os.getenv("QWEN3_VLLM_NUM_LAYERS"):
527
+ raise ValueError(f"vLLM asked for {num_layers} attention layers, model has {self.model.num_layers}")
528
+ # Reduced bring-up target: vLLM sized blocks for the real depth, the
529
+ # model only has a few layers. Allocating the model's depth is the
530
+ # right thing -- it is strictly less memory and the page geometry,
531
+ # which is what this loop is testing, is unchanged.
532
+ logger.warning(
533
+ "Reduced target: allocating {} layer caches, vLLM planned for {}",
534
+ self.model.num_layers,
535
+ num_layers,
536
+ )
537
+ if kv_heads != self.model.config.local_attention.num_key_value_heads:
538
+ raise ValueError(
539
+ f"vLLM computed {kv_heads} local KV heads, this port shards to "
540
+ f"{self.model.config.local_attention.num_key_value_heads} per die"
541
+ )
542
+ if head_dim != self.model.head_dim:
543
+ raise ValueError(f"vLLM head_dim {head_dim} != model head_dim {self.model.head_dim}")
544
+
545
+ pages_per_user = min(math.ceil(self.max_model_len / block_size), num_blocks)
546
+ self.generator.configure_paging(
547
+ page_block_size=block_size, pages_per_user=pages_per_user, num_blocks=num_blocks
548
+ )
549
+ logger.info(
550
+ "vLLM-owned KV cache: {} blocks x {} tokens = {} tokens, {} local KV heads, dtype {} "
551
+ "(from the selected precision config; vLLM asked for {})",
552
+ num_blocks,
553
+ block_size,
554
+ num_blocks * block_size,
555
+ kv_heads,
556
+ self.model.precision.kv_cache_dtype,
557
+ dtype,
558
+ )
559
+ self.kv_cache = self.model.allocate_kv_cache(num_blocks=num_blocks)
560
+ return self.kv_cache
561
+
562
+ # -- warmup ---------------------------------------------------------------
563
+
564
+ def warmup_model_prefill(self, *, kv_cache, can_sample_on_device: bool, enable_trace: bool, **kwargs: Any) -> None:
565
+ """Compile the prefill programs on the shapes serving will actually use.
566
+
567
+ A prefill program is compiled per sequence length, so what "the shapes
568
+ serving will use" means depends entirely on whether bucketing is on:
569
+
570
+ * **Bucketed** (the default, :func:`prefill_bucket_ladder`) the shape
571
+ space is the ladder -- 19 rungs at ``pow2_half`` -- and it is finite,
572
+ so it can be enumerated and warmed. That is the whole reason
573
+ bucketing exists.
574
+ * **``QWEN3_PREFILL_BUCKETS=exact``** every distinct prompt length is
575
+ its own program and no warmup can cover them. This falls back to the
576
+ two lengths it always warmed, one of them (129) deliberately aligned
577
+ to nothing: not a multiple of the page block, of a tile, or of any
578
+ power of two, because the serving path must accept such a length.
579
+
580
+ Both arms also warm the **cached-suffix** branch, which is a different
581
+ program set again: at ``start_pos > 0`` attention switches to
582
+ ``chunked_scaled_dot_product_attention`` at
583
+ ``q_chunk_size = k_chunk_size = sdpa_chunk_size(start_pos)``, and that
584
+ is reachable at only three values -- ``min(256, start & -start)`` over
585
+ block-aligned starts is ``{64, 128, 256}`` at a 64-token block. Three
586
+ warm passes therefore cover **every** prefix-cache hit; leaving them
587
+ cold cost a measured 4.7 s on the first request that hit the cache.
588
+
589
+ Prefill is eager (``enable_trace`` is accepted and ignored -- this port
590
+ has no prefill trace), so the plugin's second warmup phase is a no-op.
591
+
592
+ **Depth is gated on the kernel cache, not on the boot.** ``tt-model``
593
+ bind-mounts a per-model host directory as ``TT_METAL_CACHE`` and keeps
594
+ it across container removal on purpose (``tt_kernel/container.py``:
595
+ "the ~10-minute cost the mounted TT_METAL_CACHE exists to avoid paying
596
+ twice"). Compiles therefore persist, but *running* the ladder does not
597
+ get cheaper -- a rung is a real prefill, and the default ladder up to
598
+ the cap is ~26 k tokens, ~26 s at 0.96 ms/token. Paying that once per
599
+ host is a bargain and once per restart is not, so ``auto`` warms only
600
+ when its marker is absent from the cache directory. The marker lives
601
+ *inside* that directory, so clearing the cache clears the marker too
602
+ and the next boot re-warms.
603
+
604
+ ``QWEN3_PREFILL_WARMUP=auto|full|off`` (``full`` ignores the marker,
605
+ ``off`` skips entirely) and ``QWEN3_PREFILL_WARMUP_MAX`` (rungs above
606
+ it are left to compile lazily; ``0`` means no cap) are the knobs.
607
+ """
608
+ # Also a boot-progress landmark. `tt-model serve` (the kernel package
609
+ # manager, which packages this model) tracks boot phases by matching
610
+ # server-log lines, and its "warming up the model" phase starts on
611
+ # ``Warming up prefill`` / ``Starting decode warmup``
612
+ # (tt_kernel/boot_progress.py::VLLM_PHASES). Logging only completion
613
+ # lines left that phase dark for the whole warmup.
614
+ if self.already_warmed_up_prefill or enable_trace:
615
+ return
616
+ self.already_warmed_up_prefill = True
617
+
618
+ mode = os.getenv("QWEN3_PREFILL_WARMUP", "auto").strip().lower()
619
+ if mode not in ("auto", "full", "off"):
620
+ raise ValueError(f"QWEN3_PREFILL_WARMUP must be auto, full or off; got {mode!r}")
621
+ if mode == "off":
622
+ logger.info("Prefill warmup skipped (QWEN3_PREFILL_WARMUP=off)")
623
+ return
624
+
625
+ lengths = self._warmup_prefill_lengths()
626
+ suffixes = self._warmup_suffix_starts()
627
+ marker = self._prefill_warmup_marker(lengths, suffixes)
628
+ if mode == "auto" and marker is not None and marker.exists():
629
+ # The programs are already in the persistent kernel cache; a first
630
+ # request loads them rather than compiling them.
631
+ logger.info(
632
+ "Prefill warmup skipped -- kernel cache already primed for this ladder ({})",
633
+ marker,
634
+ )
635
+ return
636
+
637
+ logger.info(
638
+ "Warming up prefill ({} lengths {}; {} cached-suffix shapes, chunk sizes {} x suffix rungs {})",
639
+ len(lengths),
640
+ list(lengths),
641
+ len(suffixes),
642
+ sorted({chunk for _, chunk, _ in suffixes}),
643
+ sorted({suffix for _, _, suffix in suffixes}),
644
+ )
645
+ sampling = self._neutral_sampling_params(1) if can_sample_on_device else None
646
+ for length in lengths:
647
+ self.prefill_forward(
648
+ tokens=torch.zeros((1, length), dtype=torch.int32),
649
+ page_table=self._warmup_page_table(1, length),
650
+ kv_cache=kv_cache,
651
+ enable_trace=False,
652
+ prompt_lens=[length],
653
+ start_pos=[0],
654
+ sampling_params=sampling,
655
+ )
656
+ for start, _chunk, suffix in suffixes:
657
+ # ``suffix`` new tokens on top of a ``start``-token "prefix". The
658
+ # content is irrelevant -- only the shapes are being compiled -- but
659
+ # the start must be block-aligned and strictly inside the prompt,
660
+ # which ``_warmup_suffix_starts`` guarantees.
661
+ total = start + suffix
662
+ self.prefill_forward(
663
+ tokens=torch.zeros((1, total), dtype=torch.int32),
664
+ page_table=self._warmup_page_table(1, total),
665
+ kv_cache=kv_cache,
666
+ enable_trace=False,
667
+ prompt_lens=[total],
668
+ start_pos=[start],
669
+ sampling_params=sampling,
670
+ )
671
+ if marker is not None:
672
+ try:
673
+ marker.parent.mkdir(parents=True, exist_ok=True)
674
+ marker.write_text("qwen3-coder-30b-a3b prefill warmup complete\n")
675
+ except OSError as exc:
676
+ # A read-only or missing cache directory means the next boot
677
+ # warms again -- slower, never wrong. Not worth failing a boot.
678
+ logger.warning("Could not write the prefill warmup marker {}: {}", marker, exc)
679
+ logger.info("Prefill warmup done ({} lengths + {} cached-suffix shapes)", len(lengths), len(suffixes))
680
+
681
+ def _warmup_prefill_lengths(self) -> tuple[int, ...]:
682
+ """Prefill lengths to warm: the bucket ladder, capped.
683
+
684
+ Capped twice over -- by ``QWEN3_PREFILL_WARMUP_MAX`` because a rung is a
685
+ real prefill and the tail of the ladder is enormous (the top rung is the
686
+ whole 256 k context, ~4 minutes on its own), and by the KV cache,
687
+ because a single warm row cannot occupy more blocks than one user is
688
+ allotted or than the cache holds.
689
+ """
690
+ ladder = self.generator.prefill_buckets
691
+ if not ladder:
692
+ # ``exact``: nothing to enumerate. 129 is aligned to nothing on
693
+ # purpose -- see the class docstring.
694
+ return (129, 128)
695
+ raw = os.getenv("QWEN3_PREFILL_WARMUP_MAX", "8192").strip()
696
+ cap = int(raw) if raw else 8192
697
+ block = self.generator.page_block_size
698
+ affordable = min(self.generator.pages_per_user, self.generator.num_blocks) * block
699
+ if cap > 0:
700
+ # The cap is a promise about PROMPTS, so it has to include the rung
701
+ # that *covers* it, not merely the rungs below it. Measured the
702
+ # other way round: at a cap of 8192 an 8192-token prompt rounded to
703
+ # the 8320 rung, which had been left out, and paid 15.088 s against
704
+ # 7.610 s once compiled -- the cap silently failing at exactly the
705
+ # length it named.
706
+ affordable = min(affordable, self.generator.prefill_padded_len(cap))
707
+ return tuple(v for v in ladder if v <= affordable)
708
+
709
+ def _warmup_suffix_starts(self) -> tuple[tuple[int, int, int], ...]:
710
+ """``(start_pos, chunk_size, suffix_len)`` for every cached-suffix shape warmed.
711
+
712
+ The chunked-SDPA program is keyed on *both* axes, so warming one is not
713
+ enough:
714
+
715
+ * ``chunk_size`` is ``sdpa_chunk_size(start) == min(256, start & -start)``,
716
+ and a split prefill's ``start`` must be a multiple of the KV block
717
+ size -- so the reachable set is small and closed (``{64, 128, 256}``
718
+ at a 64-token block). This enumerates block-aligned starts and keeps
719
+ the cheapest start reaching each distinct chunk size, rather than
720
+ hardcoding numbers that were only true at one block size.
721
+ * ``suffix_len`` is the *bucketed* length of the new tokens, so it is a
722
+ rung of the prefill ladder -- which is exactly why bucketing had to
723
+ come first. Unbucketed, this axis is unbounded and the cross product
724
+ is unwarmable.
725
+
726
+ Capped by ``QWEN3_PREFILL_WARMUP_SUFFIX_MAX`` (default 1024) because the
727
+ cross product is chunk sizes x rungs and only the low rungs are common:
728
+ a prefix-cache hit's suffix is the *new* tail of a conversation, not a
729
+ whole prompt. A hit with a longer suffix compiles lazily, once, and the
730
+ persistent kernel cache keeps it.
731
+ """
732
+ block = self.generator.page_block_size
733
+ affordable = min(self.generator.pages_per_user, self.generator.num_blocks) * block
734
+ raw = os.getenv("QWEN3_PREFILL_WARMUP_SUFFIX_MAX", "1024").strip()
735
+ suffix_cap = int(raw) if raw else 1024
736
+ # 512 is past the point of diminishing returns: chunk sizes saturate at
737
+ # 256, so no larger start reaches a chunk size a smaller one did not.
738
+ starts: dict[int, int] = {}
739
+ start = block
740
+ while start <= 512:
741
+ starts.setdefault(sdpa_chunk_size(start), start)
742
+ start += block
743
+ rungs = [v for v in (self.generator.prefill_buckets or (block,)) if v <= suffix_cap]
744
+ out = []
745
+ for chunk, start in sorted(starts.items()):
746
+ for rung in rungs:
747
+ if start + rung <= affordable:
748
+ out.append((start, chunk, rung))
749
+ return tuple(out)
750
+
751
+ def _prefill_warmup_marker(self, lengths, suffixes) -> Optional[Path]:
752
+ """Where ``auto`` records that this ladder has been compiled once.
753
+
754
+ Inside the kernel cache directory, and named for everything that would
755
+ invalidate it: the ladder, the block size, and the mesh shape. If no
756
+ cache directory is configured there is nothing persistent to key off,
757
+ so return ``None`` and let ``auto`` behave like ``full``.
758
+ """
759
+ root = os.getenv("TT_METAL_CACHE", "").strip()
760
+ if not root:
761
+ return None
762
+ try:
763
+ shape = "x".join(str(int(v)) for v in tuple(self.generator.mesh_device.shape))
764
+ except Exception: # pragma: no cover -- naming detail, never worth a failed boot
765
+ shape = "unknown"
766
+ stamp = "-".join(
767
+ (
768
+ str(len(lengths)),
769
+ str(lengths[-1] if lengths else 0),
770
+ str(len(suffixes)),
771
+ str(self.generator.page_block_size),
772
+ shape,
773
+ )
774
+ )
775
+ return Path(root) / f".qwen3_coder_30b_a3b.prefill_warm.v1.{stamp}"
776
+
777
+ def warmup_model_decode(
778
+ self,
779
+ *,
780
+ kv_cache,
781
+ max_batch_size: int,
782
+ num_blocks: int,
783
+ can_sample_on_device: bool,
784
+ enable_trace: bool,
785
+ **kwargs: Any,
786
+ ) -> None:
787
+ """Compile and then capture a decode graph for **every** ladder width.
788
+
789
+ Capturing here rather than on the first real token matters twice over:
790
+ it keeps a multi-second trace capture out of the benchmark's
791
+ inter-token latency, and it is the phase the plugin gives us for exactly
792
+ that (phase 1 ``enable_trace=False`` compiles, phase 2 captures).
793
+
794
+ **Every rung, not just the serving batch.** A width's *first* capture
795
+ also pays an eager warm pass to compile its programs -- 1.9-4.8 s
796
+ against the 713 ms a capture costs once they are cached
797
+ (``doc/batch_scaling/README.md``, "Cost of switching") -- and
798
+ ``_decode_compiled_keys`` only spares a width it has already seen. Left
799
+ lazy, that multi-second compile lands on whichever request first needs
800
+ the rung, which for a single user against a 32-slot server is the very
801
+ first one. Warming the ladder moves all of it into the phase built for
802
+ it, for 24.8 MB of per-width persistent inputs across all six -- 0.11 %
803
+ of the 21.9 GB/die a maximal prefill leaves free, and already measured
804
+ with every width captured at the shipped 48 MB trace region.
805
+
806
+ **Descending, so width 1 is the one left resident.** Only one trace pair
807
+ is resident at a time -- ``_refresh_trace_state`` releases on a
808
+ ``graph_width``, ``active_batch``, ``kv_cache`` or page-table-shape
809
+ change -- so the last rung warmed is the one a first request can replay
810
+ for free. Ending at the narrowest makes the single-user case, the one
811
+ the ladder exists for, cost no capture at all. Each rung is warmed at
812
+ ``active_batch == graph_width`` and at vLLM's own table width, because
813
+ both are part of what a release is keyed on: warming a width at some
814
+ other batch would compile the programs but still discard the trace.
815
+
816
+ **Both sampling strategies, greedy last.** ``_decode_graph_key`` covers
817
+ ``_sampling_stochastic`` as well as the width, and the argmax and
818
+ split-sampling graphs are different program sets --
819
+ ``set_sampling_params`` releases the traces when the flag flips. A
820
+ served model does not choose what its callers send, so both are warmed:
821
+ measured on 4 dies, an unwarmed strategy cost **1.9416 s** of TTFT on a
822
+ first request against **0.2989 s** warmed, and the warmed cost holds
823
+ even for a strategy the server's own default never uses (a
824
+ Qwen-recommended ``temperature=0.7, top_p=0.8, top_k=20`` request
825
+ against a greedy-default server: 0.2579 s).
826
+
827
+ Greedy is warmed last, so it is the resident one. Order is nearly free
828
+ -- flipping strategy at width 1 measured 0.2434 s TTFT against 0.2430 s
829
+ steady, i.e. inside noise -- but greedy is both the faster path
830
+ (20.35 ms/token, 49.1 t/s/u, against 25.88 ms and 38.6) and what
831
+ ``tt-model.yaml`` makes the default request, via
832
+ ``--override-generation-config '{"temperature": 0}'`` on top of
833
+ ``--generation-config vllm``. Without that override an unparameterised
834
+ request would arrive at vLLM's own ``temperature=1.0, top_p=1.0``, and
835
+ ``stochastic = any(k > 1) or any(p > 0.0)`` would put the default path
836
+ on the slower graph. ``QWEN3_WARMUP_SAMPLING`` trims this to ``greedy``
837
+ or ``stochastic`` alone when a deployment knows it only serves one.
838
+
839
+ Penalties are deliberately *not* warmed. ``_penalty_mode`` is in the key
840
+ too, but the same ``--generation-config vllm`` is what keeps this
841
+ checkpoint's ``repetition_penalty=1.05`` from reaching every request, so
842
+ mode 0 is the steady state and warming a penalised graph would spend
843
+ 1.9-4.8 s per width on a path a default deployment never takes. A server
844
+ run *without* that flag inverts this, and would want mode 1 warmed
845
+ instead. ``rope_cache_len`` is stable (``decode_horizon`` is the served
846
+ context, sized at construction) and ``active_row_gating`` is fixed for
847
+ the process, so neither adds a dimension here.
848
+ """
849
+ if not can_sample_on_device:
850
+ # Host-sampled decode is eager by construction; nothing to capture.
851
+ return
852
+ batch = min(int(max_batch_size), self.max_num_seqs)
853
+ # Descending, and clipped to what the caller can actually send: a graph
854
+ # wider than ``batch`` is unreachable this run, and warming it would
855
+ # compile programs no request can replay. ``batch`` itself is always
856
+ # warmed -- it is the fallback rung ``_configured_widths`` guarantees --
857
+ # so with the ladder off this is exactly the previous single warmup.
858
+ widths = sorted({min(w, batch) for w in self._decode_widths}, reverse=True)
859
+ strategies = self._warmup_sampling_strategies()
860
+ # Boot-progress landmark as well as a plan: see warmup_model_prefill.
861
+ logger.info(
862
+ "Starting decode warmup ({} graphs: widths {} descending x sampling {})",
863
+ len(widths) * len(strategies),
864
+ widths,
865
+ ["stochastic" if st else "greedy" for st in strategies],
866
+ )
867
+ for stochastic in strategies:
868
+ for graph_width in widths:
869
+ page_table = self._warmup_page_table(graph_width, self.generator.page_block_size, width=int(num_blocks))
870
+ self.decode_forward(
871
+ tokens=torch.zeros((graph_width, 1), dtype=torch.int32),
872
+ page_table=page_table,
873
+ kv_cache=kv_cache,
874
+ start_pos=torch.zeros(graph_width, dtype=torch.int64),
875
+ enable_trace=enable_trace,
876
+ read_from_device=True,
877
+ sampling_params=self._warmup_sampling_params(graph_width, stochastic=stochastic),
878
+ reset_batch=True,
879
+ )
880
+ # Each warmup step wrote a token at position 0 of every row it used and
881
+ # advanced the device positions; the first real request must not inherit
882
+ # that, and must re-decide its own width against a real batch layout.
883
+ self._needs_decode_install = True
884
+ self._reset_pending_orders()
885
+ self._audit["decode_graph_width"] = widths[-1]
886
+ logger.info(
887
+ "Decode warmup done ({} widths {} descending x sampling {}, resident {} {}, enable_trace={})",
888
+ len(widths) * len(strategies),
889
+ widths,
890
+ ["stochastic" if st else "greedy" for st in strategies],
891
+ "stochastic" if strategies[-1] else "greedy",
892
+ widths[-1],
893
+ enable_trace,
894
+ )
895
+
896
+ def _warmup_sampling_strategies(self) -> list[bool]:
897
+ """Sampling strategies to warm, in warm order -- the last stays resident.
898
+
899
+ ``False`` is argmax/greedy, ``True`` is split sampling. Both by default,
900
+ because a served model does not get to choose what its callers send and
901
+ an uncompiled strategy costs seconds on a real request -- measured on 4
902
+ dies, an unwarmed strategy put 1.9416 s on a first request's TTFT
903
+ against 0.2989 s warmed.
904
+
905
+ **Greedy is warmed last, so it is the one left resident.** Order is a
906
+ near-free choice rather than a load-bearing one: switching strategy at
907
+ width 1 releases and re-captures, and that cost was unmeasurable
908
+ (0.2434 s TTFT switching against 0.2430 s steady). Greedy gets the slot
909
+ because it is both the faster path (20.34 ms/token against 25.88 ms,
910
+ 49.1 t/s/u against 38.6) and what ``tt-model.yaml`` makes the default
911
+ request via ``--override-generation-config '{"temperature": 0}'``.
912
+
913
+ ``QWEN3_WARMUP_SAMPLING=greedy|stochastic`` warms one only, for a
914
+ deployment that knows which it serves; anything else is rejected rather
915
+ than silently ignored, since a typo would quietly cost that same
916
+ multi-second compile on a real request.
917
+ """
918
+ choice = os.getenv("QWEN3_WARMUP_SAMPLING", "both").strip().lower()
919
+ if choice in ("", "both"):
920
+ return [True, False]
921
+ if choice == "greedy":
922
+ return [False]
923
+ if choice == "stochastic":
924
+ return [True]
925
+ raise ValueError(f"QWEN3_WARMUP_SAMPLING must be both, greedy or stochastic; got {choice!r}")
926
+
927
+ def _warmup_sampling_params(self, rows: int, *, stochastic: bool):
928
+ """Neutral params for one warm pass, on the requested sampling strategy.
929
+
930
+ The stochastic arm mirrors what vLLM sends for an unparameterised
931
+ request under ``--generation-config vllm`` -- ``temperature=1.0``,
932
+ ``top_p=1.0`` -- with ``top_k`` already at the device limit rather than
933
+ vLLM's "disabled" 0, which ``_apply_sampling_params`` would clamp to the
934
+ same 32 while incrementing ``top_k_clamped_requests`` and warning. The
935
+ graph depends only on the boolean, so the values beyond that are
936
+ irrelevant to what gets compiled.
937
+ """
938
+ if not stochastic:
939
+ return self._neutral_sampling_params(rows)
940
+ from vllm_tt_plugin.model_input import TTSamplingParams
941
+
942
+ return TTSamplingParams(
943
+ temperature=[1.0] * rows,
944
+ top_k=[SAMPLING_SLOTS] * rows,
945
+ top_p=[1.0] * rows,
946
+ seed=[None] * rows,
947
+ )
948
+
949
+ def _warmup_page_table(self, batch: int, token_count: int, *, width: int | None = None) -> torch.Tensor:
950
+ """A disjoint block assignment for warmup only, at vLLM's table width."""
951
+ width = self.generator.pages_per_user if width is None else int(width)
952
+ blocks = max(1, math.ceil(token_count / self.generator.page_block_size))
953
+ table = torch.zeros((batch, width), dtype=torch.int32)
954
+ for row in range(batch):
955
+ span = min(blocks, width)
956
+ table[row, :span] = torch.arange(row * span, row * span + span, dtype=torch.int32)
957
+ return table
958
+
959
+ def _neutral_sampling_params(self, rows: int):
960
+ from vllm_tt_plugin.model_input import TTSamplingParams
961
+
962
+ return TTSamplingParams(
963
+ temperature=[0.0] * rows,
964
+ top_k=[1] * rows,
965
+ top_p=[1.0] * rows,
966
+ # The plugin translates its own "no seed" sentinel to ``None`` before
967
+ # the model sees it; the dataclass default of ``0`` is a real seed.
968
+ seed=[None] * rows,
969
+ )
970
+
971
+ # -- sampling translation -------------------------------------------------
972
+
973
+ def _apply_sampling_params(self, sampling_params, rows: int, *, order=None, graph_rows: int | None = None) -> None:
974
+ """vLLM's per-row sampling request -> the generator's ``(k, p, temp)``.
975
+
976
+ Nothing here samples. It only sets the three persistent device parameter
977
+ tensors that ``Qwen3CoderModel.sample_split`` reads, and only when they
978
+ actually changed -- ``set_sampling_params`` no-ops on an identical
979
+ snapshot, so a steady greedy benchmark costs zero host copies per token.
980
+ """
981
+ temps = _as_int_list(getattr(sampling_params, "temperature", None), rows, 0.0)
982
+ top_ks = _as_int_list(getattr(sampling_params, "top_k", None), rows, 1)
983
+ top_ps = _as_int_list(getattr(sampling_params, "top_p", None), rows, 1.0)
984
+ self._audit_unsupported(sampling_params, rows)
985
+ # ``ttnn.sampling``'s per-slot parameters address the *graph*'s rows, so
986
+ # a compacted batch must present them in the same order the rows are in.
987
+ if order is not None:
988
+ picks = [int(v) for v in order.tolist()]
989
+ temps = [temps[i] for i in picks]
990
+ top_ks = [top_ks[i] for i in picks]
991
+ top_ps = [top_ps[i] for i in picks]
992
+ rows = rows if graph_rows is None else int(graph_rows)
993
+
994
+ k_out: list[int] = []
995
+ p_out: list[float] = []
996
+ t_out: list[float] = []
997
+ for row in range(rows):
998
+ temperature = float(temps[row])
999
+ top_k = int(top_ks[row])
1000
+ top_p = float(top_ps[row])
1001
+ if temperature <= 0.0:
1002
+ # Greedy. The generator maps temperature 0 to k=1, p=0 itself and
1003
+ # then routes to the argmax strategy.
1004
+ k_out.append(1)
1005
+ p_out.append(0.0)
1006
+ t_out.append(0.0)
1007
+ continue
1008
+ if top_k <= 0 or top_k > SAMPLING_SLOTS:
1009
+ # vLLM spells "no top-k" as <=0 and allows any k up to the
1010
+ # vocabulary; ``Sampling1DConfig(max_top_k=32)`` is a device
1011
+ # limit, so both collapse to the widest supported candidate set.
1012
+ if top_k > SAMPLING_SLOTS or top_k <= 0:
1013
+ self._audit["top_k_clamped_requests"] += 1
1014
+ self._warn_once(
1015
+ "top_k",
1016
+ f"top_k={top_k} clamped to {SAMPLING_SLOTS}: the on-device sampler's "
1017
+ "max_top_k is 32 candidates per die-gathered slot.",
1018
+ )
1019
+ top_k = SAMPLING_SLOTS
1020
+ k_out.append(top_k)
1021
+ p_out.append(min(max(top_p, 0.0), 1.0))
1022
+ t_out.append(temperature)
1023
+ self.generator.set_sampling_params(top_k=k_out, top_p=p_out, temperature=t_out, active_batch=rows)
1024
+
1025
+ def _apply_penalties(
1026
+ self, sampling_params, rows: int, prompt_tokens, output_tokens, *, order=None, graph_rows: int | None = None
1027
+ ) -> None:
1028
+ """vLLM's three penalties -> the generator's staged on-device penalty stage.
1029
+
1030
+ The plugin packs ``presence_penalty`` / ``frequency_penalty`` /
1031
+ ``repetition_penalty`` into ``TTSamplingParams`` and sends the token
1032
+ history alongside them (``model_runner.py`` populates ``prompt_tokens``
1033
+ and ``output_tokens`` "if penalties are needed (decode only)"), because
1034
+ ``platform.py`` deliberately does **not** route penalised requests to host
1035
+ sampling. This is the model side of that contract; the stage itself is
1036
+ ``_WatcherCleanSampling1D._apply_penalties``.
1037
+
1038
+ Neutral on every row is the fast path: ``set_penalty_params`` returns
1039
+ False, the ops are not in the captured trace at all, and nothing is
1040
+ uploaded.
1041
+ """
1042
+ presence = _as_int_list(getattr(sampling_params, "presence_penalty", None), rows, 0.0)
1043
+ frequency = _as_int_list(getattr(sampling_params, "frequency_penalty", None), rows, 0.0)
1044
+ repetition = _as_int_list(getattr(sampling_params, "repetition_penalty", None), rows, 1.0)
1045
+ if order is not None:
1046
+ picks = [int(v) for v in order.tolist()]
1047
+ presence = [presence[i] for i in picks]
1048
+ frequency = [frequency[i] for i in picks]
1049
+ repetition = [repetition[i] for i in picks]
1050
+ # The staged penalty rows are per *graph* row too, and the history
1051
+ # they are keyed on has to travel with them.
1052
+ prompt_tokens = _reorder_history(prompt_tokens, order, picks)
1053
+ output_tokens = _reorder_history(output_tokens, order, picks)
1054
+ live, graph_changed = self.generator.set_penalty_params(
1055
+ presence=presence,
1056
+ frequency=frequency,
1057
+ repetition=repetition,
1058
+ prompt_tokens=prompt_tokens,
1059
+ output_tokens=output_tokens,
1060
+ active_batch=rows if graph_rows is None else int(graph_rows),
1061
+ )
1062
+ if graph_changed:
1063
+ # The mode flip released the decode traces; the next step must
1064
+ # reinstall host state rather than replay a freed trace.
1065
+ self._needs_decode_install = True
1066
+ # Any queued mapping refers to a forward whose output tensor the
1067
+ # released trace owned, so those outputs can no longer be read.
1068
+ self._reset_pending_orders()
1069
+ if live:
1070
+ self._audit["penalised_decode_steps"] += 1
1071
+
1072
+ def _audit_unsupported(self, sampling_params, rows: int) -> None:
1073
+ """Record -- loudly, once -- the request features this sampler drops."""
1074
+ seeds = _as_int_list(getattr(sampling_params, "seed", None), rows, None)
1075
+ if any(s is not None for s in seeds):
1076
+ self._audit["ignored_seed_requests"] += 1
1077
+ self._warn_once(
1078
+ "seed",
1079
+ "A per-request seed was supplied but this port's sampler draws from its own device "
1080
+ "RNG buffer; sampled output is not reproducible from the request seed. See "
1081
+ "doc/vllm_integration/README.md, Limitations.",
1082
+ )
1083
+
1084
+ # -- prefill --------------------------------------------------------------
1085
+
1086
+ def prefill_forward(
1087
+ self,
1088
+ *,
1089
+ tokens: torch.Tensor,
1090
+ page_table: torch.Tensor,
1091
+ kv_cache,
1092
+ enable_trace: bool = False,
1093
+ prompt_lens=None,
1094
+ start_pos=None,
1095
+ sampling_params=None,
1096
+ empty_slots=None,
1097
+ page_tables_per_layer=None,
1098
+ **kwargs: Any,
1099
+ ):
1100
+ """One serving prefill step, straight into ``generator.prefill_forward``.
1101
+
1102
+ ``tokens`` is ``[num_reqs, max(prompt_lens)]`` with each row's real
1103
+ length in ``prompt_lens`` and *garbage past it* -- vLLM slices a shared
1104
+ buffer. The generator prefills each row at exactly its own logical
1105
+ length (``tokens[user, :prompt_len]``), so a prompt length that is not a
1106
+ multiple of the page block, the tile height or any chunk size needs no
1107
+ special case: nothing rounds it up on the way in and the selected row is
1108
+ ``prompt_len - 1``.
1109
+ """
1110
+ if page_tables_per_layer is not None:
1111
+ raise ValueError("this port has one uniform full-attention KV-cache group; per-layer tables are not used")
1112
+ active = int(tokens.shape[0])
1113
+ starts = _as_int_list(start_pos, active, 0) if start_pos is not None else [0] * active
1114
+ if any(int(p) != 0 for p in starts) and not _PREFIX_CACHING_ENABLED:
1115
+ raise ValueError(
1116
+ "non-zero prefill start_pos means prefix caching, but it has been "
1117
+ "disabled via QWEN3_PREFIX_CACHING=0 while model_capabilities still "
1118
+ "advertises supports_prefix_caching=True. Those two must agree: either "
1119
+ "unset the kill switch or set supports_prefix_caching=False."
1120
+ )
1121
+ lengths = [int(n) for n in _as_int_list(prompt_lens, active, int(tokens.shape[1]))]
1122
+ device_sampling = sampling_params is not None
1123
+ if device_sampling:
1124
+ self._apply_sampling_params(sampling_params, active)
1125
+ self._audit["device_sampled_prefills"] += 1
1126
+ else:
1127
+ self._audit["host_sampled_prefills"] += 1
1128
+
1129
+ out = self.generator.prefill_forward(
1130
+ tokens.to(torch.int64),
1131
+ page_table=self._page_table_for_generator(page_table, active),
1132
+ kv_cache=self._require_cache(kv_cache),
1133
+ prompt_lens=lengths,
1134
+ sampling_mode="device" if device_sampling else "host",
1135
+ # A new request is admitted while other slots are mid-decode; see
1136
+ # the argument in ``Qwen3CoderGenerator.prefill_forward``, and the
1137
+ # measurement behind this default in
1138
+ # ``doc/vllm_integration/work_log.md``.
1139
+ preserve_decode_traces=PRESERVE_DECODE_TRACES,
1140
+ start_pos=starts if _PREFIX_CACHING_ENABLED else None,
1141
+ )
1142
+ # Whoever held these slots before is gone; the next decode must reinstall
1143
+ # host state rather than replay over the device's stale token/position.
1144
+ self._needs_decode_install = True
1145
+ self._reset_pending_orders()
1146
+ if device_sampling:
1147
+ return self.generator.read_sampled_tokens(out, active).reshape(active, 1)
1148
+ # Host-sampling compatibility mode: vLLM wants ``[B, S, vocab]`` and
1149
+ # reads ``[:, -1, :]``. The generator already returns one row per user.
1150
+ return out.reshape(active, 1, -1)
1151
+
1152
+ # -- decode ---------------------------------------------------------------
1153
+
1154
+ def decode_forward(
1155
+ self,
1156
+ *,
1157
+ tokens: torch.Tensor,
1158
+ page_table: torch.Tensor,
1159
+ kv_cache,
1160
+ start_pos,
1161
+ enable_trace: bool = True,
1162
+ read_from_device: bool = True,
1163
+ sampling_params=None,
1164
+ reset_batch: bool | None = None,
1165
+ slot_remap=None,
1166
+ prompt_tokens=None,
1167
+ output_tokens=None,
1168
+ page_tables_per_layer=None,
1169
+ **kwargs: Any,
1170
+ ):
1171
+ """One serving decode step.
1172
+
1173
+ The batch is always the full ``max_num_seqs`` rows -- vLLM pads it so the
1174
+ trace shape is constant -- with inactive slots carrying position ``-1``,
1175
+ which is the same inactive-row convention the generator's low-level API
1176
+ already had.
1177
+
1178
+ Steady state is the whole point: ``tokens``, ``start_pos`` and
1179
+ ``page_table`` are all passed as ``None``/unchanged, the two traces
1180
+ replay non-blocking, the sampled token is fed back on device and both
1181
+ position tensors advance on device. Host work per token is two
1182
+ ``ttnn.execute_trace`` calls and one page-table equality check.
1183
+ """
1184
+ if page_tables_per_layer is not None:
1185
+ raise ValueError("this port has one uniform full-attention KV-cache group; per-layer tables are not used")
1186
+ caches = self._require_cache(kv_cache)
1187
+ rows = int(tokens.shape[0])
1188
+ host_tokens = tokens.reshape(-1).to(torch.int64)
1189
+ host_positions = torch.as_tensor(start_pos).reshape(-1).to(torch.int64)
1190
+
1191
+ if sampling_params is None:
1192
+ # Explicit host-sampling compatibility mode. vLLM routes a request
1193
+ # here on its own (logprobs on a 4-die mesh, min_p, bad_words,
1194
+ # logit_bias, structured output); it is never the measured path.
1195
+ # The eager decode allocates, so ``decode_forward`` releases the
1196
+ # captured decode traces -- hence ``_needs_decode_install`` below,
1197
+ # which makes the next device-sampled step re-capture through
1198
+ # ``_refresh_trace_state`` instead of replaying a released trace.
1199
+ self._audit["host_sampled_decode_steps"] += 1
1200
+ # Make the demotion loud. vLLM decides this per request and logs
1201
+ # nothing, so without this line a served request silently drops from
1202
+ # ~49 t/s/u to ~3.6 and the only visible symptom is that the model
1203
+ # "got slow" -- which is exactly how this port's batch-scaling defect
1204
+ # was first reported. The server-level ``sample_on_device_mode: all``
1205
+ # is still correct and still says nothing about it.
1206
+ self._warn_once(
1207
+ "host_sampled_decode",
1208
+ "This request was routed to HOST sampling by vLLM, so decode runs "
1209
+ "eager with no captured trace and no width compaction: measured "
1210
+ "3.595 t/s/u against 49.345 on the traced path, a ~14x slowdown "
1211
+ "for the affected requests. On this 4-die mesh the usual cause is "
1212
+ "`logprobs` -- ANY value including 0, because "
1213
+ "`model_runner.check_perform_device_sampling` tests "
1214
+ "`max_num_logprobs is not None` and then rejects a mesh that is "
1215
+ "not 8 or 32 dies. Other triggers are min_p, bad_words, "
1216
+ "logit_bias and structured output. Drop the offending parameter "
1217
+ "to stay on the traced path; see doc/batch_scaling/README.md, "
1218
+ "'logprobs silently cost 14x on this mesh'.",
1219
+ )
1220
+ self._needs_decode_install = True
1221
+ self._reset_pending_orders()
1222
+ logits = self.generator.decode_forward(
1223
+ host_tokens,
1224
+ torch.clamp(host_positions, min=0),
1225
+ page_table=self._page_table_for_generator(page_table, rows),
1226
+ kv_cache=caches,
1227
+ sampling_mode="host",
1228
+ enable_trace=False,
1229
+ active_batch=rows,
1230
+ validate_page_coverage=False,
1231
+ )
1232
+ return logits.reshape(rows, 1, -1)
1233
+
1234
+ self._audit["device_sampled_decode_steps"] += 1
1235
+ install = bool(reset_batch) or self._needs_decode_install
1236
+
1237
+ # The graph width, and with it the graph-row -> vLLM-slot mapping, may
1238
+ # only change on an install: that is the one step the plugin guarantees
1239
+ # is not steady-decode eligible, so nothing is in flight against the old
1240
+ # trace. On every other step the previous mapping still describes the
1241
+ # live trace and is reused unchanged.
1242
+ previous_order = self._compaction
1243
+ if install and len(self._decode_widths) > 1:
1244
+ # ``rows`` is the padded decode batch, normally ``max_num_seqs``; a
1245
+ # graph can never be wider than the slots the caller actually sent.
1246
+ chosen = min(self._choose_width(int((host_positions >= 0).sum())), rows)
1247
+ order = self._compaction_order(host_positions, chosen, rows)
1248
+ identity = chosen == rows and bool(torch.equal(order, torch.arange(rows)))
1249
+ self._compaction = None if identity else order
1250
+ if chosen < rows:
1251
+ self._audit["narrow_decode_installs"] += 1
1252
+ self._audit["decode_graph_width"] = chosen
1253
+ width = rows if self._compaction is None else int(self._compaction.numel())
1254
+ order = self._compaction
1255
+ # With no extra widths configured this stays ``None`` and the generator
1256
+ # keeps its own default -- the full configured slot count -- so the
1257
+ # shipped path is untouched down to which graph gets captured.
1258
+ requested_width = None if self._compaction is None and len(self._decode_widths) == 1 else width
1259
+
1260
+ self._apply_sampling_params(sampling_params, rows, order=order, graph_rows=width)
1261
+ # Before the trace is touched: a penalty-mode change releases the decode
1262
+ # traces (the ops either are or are not in the captured graph), and the
1263
+ # buffers it may allocate cannot be allocated during a capture.
1264
+ self._apply_penalties(sampling_params, rows, prompt_tokens, output_tokens, order=order, graph_rows=width)
1265
+ # A penalty-mode change releases the traces, so it forces an install even
1266
+ # when the scheduler layout did not move. Re-read the flag rather than
1267
+ # trusting the value taken before the call; the width decision above is
1268
+ # unaffected, because the batch layout is what picks the width and that
1269
+ # has not changed.
1270
+ install = install or self._needs_decode_install
1271
+
1272
+ if install:
1273
+ merged_tokens, merged_positions = self._merge_scheduler_view(
1274
+ host_tokens, host_positions, page_table, slot_remap, rows, previous_order
1275
+ )
1276
+ if order is not None:
1277
+ merged_tokens = merged_tokens[order]
1278
+ merged_positions = merged_positions[order]
1279
+ sampled = self.generator.decode_forward(
1280
+ merged_tokens,
1281
+ merged_positions,
1282
+ page_table=self._compact_page_table(page_table, rows, order),
1283
+ kv_cache=caches,
1284
+ sampling_mode="device",
1285
+ enable_trace=True,
1286
+ active_batch=width,
1287
+ graph_width=requested_width,
1288
+ # Sized once at construction to the served context, so this only
1289
+ # asserts the horizon rather than growing anything.
1290
+ decode_horizon=self.max_model_len,
1291
+ # vLLM's block tables are its own: rows of an unused slot are
1292
+ # zero-filled rather than -1, so the standalone disjointness
1293
+ # check does not describe them.
1294
+ validate_page_coverage=False,
1295
+ )
1296
+ self._needs_decode_install = False
1297
+ self._audit["decode_trace_installs"] += 1
1298
+ else:
1299
+ sampled = self.generator.decode_forward(
1300
+ None,
1301
+ None,
1302
+ page_table=self._compact_page_table(page_table, rows, order),
1303
+ kv_cache=caches,
1304
+ sampling_mode="device",
1305
+ enable_trace=True,
1306
+ active_batch=width,
1307
+ graph_width=requested_width,
1308
+ )
1309
+
1310
+ # Pair this forward's tokens with the mapping it was issued with, before
1311
+ # anything can read them back. ``order`` is ``None`` at full width, which
1312
+ # is a meaningful entry: it says "this step needs no un-permutation".
1313
+ if self._compaction_enabled:
1314
+ self._pending_orders.append((self._next_issue_tag, order))
1315
+ self._next_issue_tag += 1
1316
+ depth = len(self._pending_orders)
1317
+ self._audit["compaction_fifo_max_depth"] = max(self._audit["compaction_fifo_max_depth"], depth)
1318
+ if depth > self._pending_orders_cap:
1319
+ # Outputs are being issued and never read: the queue is leaking.
1320
+ # Fail here rather than let it grow unbounded and mis-pair later.
1321
+ raise RuntimeError(
1322
+ f"decode row-mapping queue reached {depth} entries (cap {self._pending_orders_cap}). "
1323
+ "Decode forwards are being issued without their outputs being read, so the "
1324
+ "mapping queue no longer tracks in-flight steps. See "
1325
+ "Qwen3CoderForCausalLM._pending_orders."
1326
+ )
1327
+ if read_from_device:
1328
+ return self.process_decode_output_host(sampled, is_tokens=True)
1329
+ return sampled
1330
+
1331
+ def _merge_scheduler_view(
1332
+ self,
1333
+ host_tokens: torch.Tensor,
1334
+ host_positions: torch.Tensor,
1335
+ page_table: torch.Tensor,
1336
+ slot_remap,
1337
+ rows: int,
1338
+ previous_order: torch.Tensor | None = None,
1339
+ ) -> tuple[torch.Tensor, torch.Tensor]:
1340
+ """Reinstall host state, but keep the device's where the device is right.
1341
+
1342
+ Called only on a layout change, never per token. For a slot that is
1343
+ simply continuing, the device already advanced ``current_pos`` past the
1344
+ token it just sampled, so ``device_pos`` equals the scheduler's position
1345
+ (synchronous scheduling) or is one ahead of it (``--async-scheduling``,
1346
+ where vLLM submits step *N+1* before it has applied token *N*). Taking
1347
+ the device's pair in both cases is what makes async scheduling safe:
1348
+ the host's token would be stale and its position would re-decode a
1349
+ position that is already in the cache.
1350
+
1351
+ A slot that changed hands must take the host's pair. Position continuity
1352
+ alone cannot tell the two apart -- a recycled slot can coincidentally
1353
+ land on a matching position -- so this also requires the slot's page-table
1354
+ row to be byte-identical to the one the live trace was captured against.
1355
+ A newly admitted request is given fresh physical blocks, so its row moves.
1356
+ """
1357
+ state = self.generator.decode_device_state()
1358
+ if state is None or state["page_table"] is None:
1359
+ return host_tokens, host_positions
1360
+
1361
+ # The live trace's rows are *graph* rows. ``previous_order`` says which
1362
+ # vLLM slot each one held; scatter them back into slot order before any
1363
+ # comparison with the scheduler's view, and leave slots the narrow graph
1364
+ # did not cover at the inactive sentinel so they are never "continuing".
1365
+ if previous_order is None:
1366
+ device_tokens = state["tokens"][:rows].clone()
1367
+ device_positions = state["positions"][:rows].clone()
1368
+ snapshot = state["page_table"][:rows]
1369
+ else:
1370
+ covered = previous_order[: state["width"]]
1371
+ device_tokens = torch.zeros(rows, dtype=torch.int64)
1372
+ device_positions = torch.full((rows,), -1, dtype=torch.int64)
1373
+ device_tokens[covered] = state["tokens"][: covered.numel()]
1374
+ device_positions[covered] = state["positions"][: covered.numel()]
1375
+ snapshot = torch.zeros((rows, state["page_table"].shape[1]), dtype=state["page_table"].dtype)
1376
+ snapshot[covered] = state["page_table"][: covered.numel()]
1377
+ incoming = torch.as_tensor(page_table).to(torch.int32)[:rows]
1378
+ width = min(snapshot.shape[1], incoming.shape[1])
1379
+
1380
+ if slot_remap is not None:
1381
+ remap = torch.as_tensor(slot_remap).reshape(-1)[:rows].to(torch.int64)
1382
+ device_tokens = device_tokens[remap]
1383
+ device_positions = device_positions[remap]
1384
+ snapshot = snapshot[remap]
1385
+
1386
+ pages_unchanged = torch.all(snapshot[:, :width] == incoming[:, :width], dim=1)
1387
+ continuing = (
1388
+ ((device_positions == host_positions) | (device_positions == host_positions + 1))
1389
+ & (device_positions >= 0)
1390
+ & (host_positions >= 0)
1391
+ & pages_unchanged
1392
+ )
1393
+ merged_tokens = torch.where(continuing, device_tokens, host_tokens)
1394
+ # ``host_positions`` is the scheduler's view, and the plugin pads rows it
1395
+ # is not serving with ``-1`` (``model_runner.py`` pads decode positions
1396
+ # with ``-1`` "to indicate no position"). That ``-1`` is exactly the
1397
+ # inactive sentinel the traced graph relies on: ``ttnn.plus_one(...,
1398
+ # skip_negative_entries=True)`` leaves it alone across replays, and
1399
+ # ``_decode_active_mask`` derives the expert-gating mask from
1400
+ # ``current_pos >= 0``.
1401
+ #
1402
+ # Clamping it to 0 here would install an inactive row as position 0, the
1403
+ # mask would read it as live, and inactive-row expert gating would
1404
+ # silently become a no-op for that slot until the next prefill released
1405
+ # the traces. Single-request runs never expose it -- every row is either
1406
+ # continuing or genuinely live -- but a server churning 4 of 32 slots
1407
+ # would see the gating win appear and disappear with request turnover.
1408
+ # So preserve the sentinel and only clamp what is not already a sentinel.
1409
+ host_positions_kept = torch.where(host_positions < 0, torch.full_like(host_positions, -1), host_positions)
1410
+ merged_positions = torch.where(continuing, device_positions, host_positions_kept)
1411
+ return merged_tokens, merged_positions
1412
+
1413
+ # -- async split ----------------------------------------------------------
1414
+
1415
+ def read_decode_output(self, tt_out, async_read: bool = False):
1416
+ """Deferred, minimal host read of the sampled-token tensor.
1417
+
1418
+ The payload is one ``[1,1,1,32]`` uint32 tensor -- 128 bytes -- because
1419
+ the token was sampled on device. There is no logits readback here and
1420
+ there is nothing else to move.
1421
+ """
1422
+ if isinstance(tt_out, torch.Tensor):
1423
+ # Host-sampling compatibility mode already returned host logits.
1424
+ return tt_out, []
1425
+ if not async_read:
1426
+ return tt_out.cpu()
1427
+ self._audit["async_decode_reads"] += 1
1428
+ self._warn_once(
1429
+ "async_split",
1430
+ "vLLM took the async decode split: read_decode_output(async_read=True) on a device "
1431
+ "handle returned by decode_forward(read_from_device=False). supports_async_decode is "
1432
+ "being exercised, not merely declared.",
1433
+ )
1434
+ host = tt_out.cpu(blocking=False)
1435
+ return host, [ttnn.record_event(self.mesh_device, 0)]
1436
+
1437
+ def process_decode_output_host(self, tt_out, is_tokens: bool = False):
1438
+ """Host formatting only; submits no device work.
1439
+
1440
+ Accepts the device handle (the plugin's synchronous path skips
1441
+ ``read_decode_output`` entirely), the host ttnn tensor from the async
1442
+ path, or an already-host torch tensor.
1443
+ """
1444
+ if isinstance(tt_out, torch.Tensor):
1445
+ return tt_out
1446
+ if not is_tokens:
1447
+ raise ValueError("host-sampled decode already returns torch logits; nothing to format")
1448
+ if ttnn.is_tensor_storage_on_device(tt_out):
1449
+ # Only a *device*-resident handle means the plugin skipped
1450
+ # ``read_decode_output`` -- its synchronous path -- so the readback
1451
+ # happens now, inside ``execute_model``, rather than after the async
1452
+ # boundary. Counted so the async/sync split is evidence, not prose.
1453
+ #
1454
+ # The discriminator has to be device residency, not
1455
+ # ``not isinstance(tt_out, torch.Tensor)``: the torch case already
1456
+ # returned above, so that test was dead and fired on every step, and
1457
+ # the async path's ``read_decode_output`` hands us
1458
+ # ``tt_out.cpu(blocking=False)`` -- a ttnn *host* tensor, not a
1459
+ # ``torch.Tensor`` -- so async reads were being counted as sync.
1460
+ self._audit["sync_decode_reads"] += 1
1461
+ tokens = self.generator.read_sampled_tokens(tt_out, self.max_num_seqs)
1462
+ # Take the mapping belonging to *this* output, not whatever the adapter
1463
+ # currently holds -- see ``_pending_orders``. FIFO is the right pairing
1464
+ # because decode forwards are finalized in issue order.
1465
+ order = None
1466
+ if self._compaction_enabled:
1467
+ if not self._pending_orders:
1468
+ # There is no safe answer here. Falling back to the adapter's
1469
+ # current mapping is precisely the bug the queue exists to
1470
+ # prevent, and it would be applied in the one state where the
1471
+ # pairing is known to be broken -- every token would go to the
1472
+ # wrong request, silently. A crash is strictly better.
1473
+ self._audit["compaction_fifo_underflows"] += 1
1474
+ raise RuntimeError(
1475
+ "decode output arrived with no queued row mapping. Forwards and host reads are "
1476
+ "no longer paired one-to-one, so the sampled tokens cannot be attributed to "
1477
+ "requests. Refusing to scatter them through a mapping that is not theirs. See "
1478
+ "Qwen3CoderForCausalLM._pending_orders."
1479
+ )
1480
+ tag, order = self._pending_orders.popleft()
1481
+ if tag != self._next_output_tag:
1482
+ raise RuntimeError(
1483
+ f"decode row-mapping queue is out of step: popped tag {tag}, expected "
1484
+ f"{self._next_output_tag}. An output has been skipped or read twice, so this "
1485
+ "mapping does not belong to these tokens. See "
1486
+ "Qwen3CoderForCausalLM._pending_orders."
1487
+ )
1488
+ self._next_output_tag += 1
1489
+ if order is not None:
1490
+ # Graph row *i* sampled for vLLM slot ``order[i]``. Scatter back;
1491
+ # slots the narrow graph did not cover hold no live request and vLLM
1492
+ # discards whatever is there.
1493
+ restored = torch.zeros(self.max_num_seqs, dtype=tokens.dtype)
1494
+ restored[order] = tokens[: order.numel()]
1495
+ tokens = restored
1496
+ return tokens.reshape(-1, 1)
1497
+
1498
+ # -- helpers --------------------------------------------------------------
1499
+
1500
+ def _require_cache(self, kv_cache):
1501
+ """The cache vLLM allocated, and only that one.
1502
+
1503
+ The generator would happily allocate its own on a ``None``; in serving
1504
+ that would be a silent second cache that vLLM's block manager knows
1505
+ nothing about, so it is an error instead.
1506
+ """
1507
+ cache = self.kv_cache if kv_cache is None else kv_cache
1508
+ if cache is None:
1509
+ raise RuntimeError("vLLM has not called allocate_kv_cache; there is no serving cache to use")
1510
+ if isinstance(cache, (list, tuple)) and cache and isinstance(cache[0], (list, tuple)):
1511
+ raise ValueError("this port is single-submesh; a per-submesh cache list is not expected")
1512
+ return cache
1513
+
1514
+ def _compact_page_table(self, page_table, rows: int, order) -> torch.Tensor:
1515
+ """vLLM's block table, reordered into graph-row order.
1516
+
1517
+ The page table is the *only* thing that ties a request to its KV pages,
1518
+ so permuting its rows is what moves a request between graph rows -- and
1519
+ it is why nothing in the cache has to move.
1520
+ """
1521
+ table = self._page_table_for_generator(page_table, rows)
1522
+ return table if order is None else table[order].contiguous()
1523
+
1524
+ def _page_table_for_generator(self, page_table, rows: int) -> torch.Tensor:
1525
+ """vLLM's block table at the generator's table width.
1526
+
1527
+ vLLM sizes its table to ``max_num_blocks_per_req``; ``configure_paging``
1528
+ already made that the generator's width, so this is normally a no-op.
1529
+ Where it is not, pad with **0** rather than the generator's standalone
1530
+ ``-1``: the paged decode SDPA kernel rounds its read up to a tile/eight
1531
+ page boundary and dereferences every page in the rounded window before
1532
+ causal masking, so a tail page must map somewhere valid. vLLM pads its
1533
+ own unused entries with 0 for the same reason.
1534
+ """
1535
+ table = torch.as_tensor(page_table).to(torch.int32)
1536
+ if table.ndim != 2:
1537
+ raise ValueError(f"page_table must be rank two, got {tuple(table.shape)}")
1538
+ target = self.generator.pages_per_user
1539
+ if table.shape[1] < target:
1540
+ table = torch.nn.functional.pad(table, (0, target - table.shape[1]), value=0)
1541
+ elif table.shape[1] > target:
1542
+ table = table[:, :target]
1543
+ if table.shape[0] < rows:
1544
+ table = torch.nn.functional.pad(table, (0, 0, 0, rows - table.shape[0]), value=0)
1545
+ return table.contiguous()
1546
+
1547
+ # -- audit ----------------------------------------------------------------
1548
+
1549
+ def serving_audit(self) -> dict:
1550
+ """What the serving path actually did, for the stage's fallback audit."""
1551
+ audit = dict(self._audit)
1552
+ audit["trace_stats"] = dict(self.generator.trace_stats)
1553
+ audit["precision_config"] = str(SELECTED_PRECISION_CONFIG)
1554
+ audit["max_model_len"] = self.max_model_len
1555
+ audit["max_num_seqs"] = self.max_num_seqs
1556
+ audit["page_block_size"] = self.generator.page_block_size
1557
+ audit["pages_per_user"] = self.generator.pages_per_user
1558
+ audit["kv_cache_blocks"] = self.generator.num_blocks
1559
+ audit["model_runtime_fallbacks"] = self.model.runtime_fallback_audit(self.max_num_seqs)
1560
+ return audit
1561
+
1562
+
1563
+ #: The architecture string in this checkpoint's ``config.json`` is
1564
+ #: ``Qwen3MoeForCausalLM``; the TT plugin registers models ``TT``-prefixed.
1565
+ HF_ARCHITECTURE = "Qwen3MoeForCausalLM"
1566
+ VLLM_ARCHITECTURE = "TT" + HF_ARCHITECTURE
1567
+
1568
+ __all__ = ["Qwen3CoderForCausalLM", "HF_ARCHITECTURE", "VLLM_ARCHITECTURE", "HF_MODEL_ID"]
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/model.py ADDED
@@ -0,0 +1,1723 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Full 48-layer Qwen3-Coder-30B-A3B-Instruct on the 4-die P300_X2 mesh.
5
+
6
+ Stage 05. This module is the *wrapper* around the stage-04 optimized multichip
7
+ decoder layer and it deliberately changes nothing about that layer's strategy:
8
+
9
+ * attention TP=4 (8 Q heads, 1 K head, 1 V head per die), experts EP=4 (32 of
10
+ 128 per die), router and both residual RMSNorms and the residual replicated;
11
+ * two all-reduces per layer, ``FABRIC_1D_RING``, 2 links prefill / 1 decode;
12
+ * expert weights ``bfloat4_b`` at LoFi with ``in0_block_w`` 16/12, attention
13
+ projections ``bfloat8_b`` DRAM-sharded, paged KV cache ``bfloat16``;
14
+ * router top-k in fp32 logit space;
15
+ * **the inter-layer residual layout contract**: every layer takes and returns a
16
+ replicated ``[1, 1, B, 2048]`` bfloat16 ``TILE`` ``DRAM_MEMORY_CONFIG``
17
+ tensor, and there is no collective, gather, reshard or layout conversion
18
+ between layers. ``prefill_hidden`` and ``decode_hidden`` below are literally a
19
+ ``for`` loop over 48 layers with the residual threaded straight through.
20
+
21
+ What the wrapper adds, and where each new boundary lives:
22
+
23
+ ``embed_tokens``
24
+ **Replicated**, bf16, so the embedding output *is* the residual contract
25
+ with no collective at all. A hidden-sharded embedding would be 4x smaller
26
+ per die but would owe an all-gather on every prefill chunk and every decode
27
+ token; at 0.622 GB/die against 22.35 GB of measured headroom
28
+ (``config/context_contract.json``) the replicated table is free and the
29
+ collective is not. This is also the shape the stage-03 footprint probe
30
+ allocated, so the published capacity numbers describe what actually runs.
31
+
32
+ ``model.norm`` (final RMSNorm)
33
+ Replicated, and shares the layer code: decode uses
34
+ ``multichip_decoder.decode_residual_norm`` (width-sharded over 8 L1 cores,
35
+ the same kernel and compute config as the two residual norms), prefill uses
36
+ the interleaved ``ttnn.rms_norm``.
37
+
38
+ ``lm_head``
39
+ **Column-parallel over the vocabulary**: die *d* owns columns
40
+ ``37984*d .. 37984*d+37983`` of ``[2048, 151936]``. 151936 = 4 * 37984 and
41
+ 37984 = 32 * 1187, so the split is exact and needs no vocabulary padding.
42
+ **Logits never reach the host on the token-out path, and neither strategy
43
+ all-gathers them.** Both reduce first and gather the survivors: greedy takes
44
+ a per-die argmax and all-gathers four candidate values and indices
45
+ (``_WatcherCleanSampling1D._sample_argmax``), top-k/top-p takes a per-die
46
+ top-32 and all-gathers 32 values and indices. Which of the two is faster here
47
+ was measured, not assumed -- see ``sample_greedy_argmax``.
48
+
49
+ ``rotary`` (decode only)
50
+ ``ttnn.experimental.rotary_embedding_hf(is_decode_mode=True)`` reading a
51
+ per-user cos/sin pair **gathered on device** by ``ttnn.embedding`` from a
52
+ position tensor the trace advances with ``ttnn.plus_one``. The layer's
53
+ shipped spelling, ``ttnn.experimental.rotary_embedding``, takes the position
54
+ as a **Python int** compile-time argument and therefore cannot be replayed:
55
+ a captured trace would rotate every subsequent token at the position it was
56
+ captured at. Note this is the *HF* rotary, same ``rotate_half`` channel
57
+ convention -- so unlike stage 04's rejected ``rotary_embedding_llama`` lever
58
+ (README limitation 4) it needs no weight permutation, changes no KV-cache
59
+ channel convention and leaves prefill untouched.
60
+ """
61
+
62
+ from __future__ import annotations
63
+
64
+ import contextlib
65
+ import gc
66
+ import json
67
+ import math
68
+ import os
69
+ from collections.abc import Sequence
70
+ from pathlib import Path
71
+
72
+ import torch
73
+ from safetensors import safe_open
74
+ from transformers import AutoConfig
75
+
76
+ import ttnn
77
+ from models.common.modules.sampling.sampling_1d import Sampling1D, Sampling1DConfig
78
+
79
+ from .functional_decoder import DecoderLayerConfig, KVCache
80
+ from .multichip_decoder import (
81
+ MESH_SHAPE,
82
+ NUM_DEVICES,
83
+ TOPOLOGY,
84
+ MeshContext,
85
+ MeshDecoderConfig,
86
+ MultichipWeights,
87
+ _head_shard,
88
+ _norm_compute_config,
89
+ build_local_sparsity,
90
+ decode_residual_norm,
91
+ decoder_layer_decode_multichip,
92
+ decoder_layer_prefill_multichip,
93
+ fallback_audit,
94
+ mesh_context,
95
+ upload_multichip_weights,
96
+ )
97
+ from .precision import DEFAULT_PRECISION, PrecisionConfig, dtype_to_name
98
+
99
+ HF_MODEL_ID = "Qwen/Qwen3-Coder-30B-A3B-Instruct"
100
+ HF_REVISION = "b2cff646eb4bb1d68355c01b18ae02e7cf42d120"
101
+
102
+ HIDDEN_SIZE = 2048
103
+ VOCAB_SIZE = 151936
104
+ NUM_LAYERS = 48
105
+ HEAD_DIM = 128
106
+ MAX_CONTEXT = 262144
107
+ DEFAULT_PAGE_BLOCK_SIZE = 32
108
+ DEFAULT_MAX_BATCH_SIZE = 1
109
+ #: ``ttnn.sampling`` and ``nlp_create_qkv_heads_decode`` both address 32 fixed
110
+ #: user slots; decode is always one 32-row tile regardless of the active batch.
111
+ SAMPLING_SLOTS = 32
112
+ #: Trace region per device. Two traces (model decode + sampling) over 48 layers.
113
+ DEFAULT_TRACE_REGION_SIZE = 300_000_000
114
+ #: RoPE table rows materialised at construction; grown on demand to the request
115
+ #: horizon by ``ensure_rope_capacity`` so a short request never pays for 262144
116
+ #: rows (which would be 64 MB/die of cos plus 64 MB of sin).
117
+ DEFAULT_ROPE_CACHE_LEN = 8192
118
+
119
+ #: ``lm_head`` weight dtype. bfloat8_b halves the 155 MB/die bf16 read that a
120
+ #: decode step would otherwise make against a 2048x37984 weight.
121
+ #:
122
+ #: Since stage 07 this is an **alias** for ``DEFAULT_PRECISION.lm_head_dtype``,
123
+ #: not the source of truth: a model built at a non-default ``PrecisionConfig``
124
+ #: does not read it. See ``tt/precision.py``.
125
+ LM_HEAD_WEIGHT_DTYPE = DEFAULT_PRECISION.lm_head_dtype
126
+ #: The embedding table stays bf16: it is a gather, not a matmul, and bfloat8_b
127
+ #: would quantise every token's hidden state at the very top of the stack.
128
+ #: Alias for ``DEFAULT_PRECISION.embedding_dtype``, as above.
129
+ EMBED_WEIGHT_DTYPE = DEFAULT_PRECISION.embedding_dtype
130
+
131
+
132
+ #: ``_WatcherCleanSampling1D._sample_argmax``'s "not a winner" sentinel. Any value
133
+ #: strictly greater than the vocabulary works; 2**20 is exact in int32 and leaves
134
+ #: ``idx - BIG`` far from overflow.
135
+ _DIST_ARGMAX_BIG = 1 << 20
136
+
137
+
138
+ class _WatcherCleanSampling1D(Sampling1D):
139
+ """``Sampling1D`` with the force-argmax gather spelled the way this layer spells it.
140
+
141
+ Two overrides, for two different reasons.
142
+
143
+ ------------------------------------------------------------------------
144
+ ``_sample_argmax`` -- reduce first, gather second
145
+ ------------------------------------------------------------------------
146
+
147
+ ``Sampling1D._sample_argmax`` all-gathers the whole column-parallel logit
148
+ shard (37984 bf16 columns per die) up to the full 151936 on **every** die,
149
+ untilizes 151936 columns and runs one ``ttnn.argmax`` over them.
150
+ ``doc/full_model/tt_perf_report_full_model_decode.txt`` shows those two ops
151
+ at ``AllGatherAsync 889 us`` and ``ArgMax 859 us``. **Neither number is a
152
+ share of a token-out step, and the two must not be summed against one.**
153
+ That report is stage 05's **2-layer** window, which charges the terminal
154
+ path against two layers instead of 48 and so over-weights it by
155
+ construction -- the two rows are 27.5% and 26.5% of *that* window. And the
156
+ column is per-op device-kernel time summed over the op's own cores (2 for
157
+ the gather, 110 for the argmax), which is a different accounting from the
158
+ wall clock of a decode step. An earlier revision of this docstring set
159
+ ``889 + 859`` against "the 1.87 ms of non-layer work in a 22.079 ms
160
+ token-out step"; the near-agreement was a coincidence between two
161
+ incommensurable measurements and the claim is withdrawn.
162
+
163
+ The full 48-layer profile is the accounting that means something. On the
164
+ shipped tree ``doc/optimized_full_model/probes/profile_summary_decode.json``
165
+ puts the *whole* terminal block -- final norm, LM head, this sampler and the
166
+ token feedback -- at **366.5 us of an 18889.5 us decode iteration, 1.94%**,
167
+ of which this sampler is **126.2 us**. The baseline path was replaced before
168
+ that profile was taken and so has no 48-layer op row of its own; its
169
+ in-model price is a token-out delta and is quoted as one under
170
+ ``sample_greedy_argmax``.
171
+
172
+ The override computes the same token by reducing on each die first and
173
+ all-gathering only the four survivors::
174
+
175
+ rm = untilize(local_shard) # bf16 [1,1,32,37984]
176
+ rm = rm[:, :, :B, :] # bf16 [1,1,B,37984]
177
+ local_idx = argmax(rm, -1, keepdim) # uint32 [1,1,B,1]
178
+ local_max = gather(rm, -1, local_idx) # bf16 [1,1,B,1]
179
+ global_idx = local_idx + rank*37984 # int32, sharded constant
180
+ vals4 = all_gather(local_max) # bf16 [1,1,B,4]
181
+ idx4 = all_gather(global_idx) # int32 [1,1,B,4]
182
+ gmax = max(vals4, -1, keepdim)
183
+ mask = (vals4 == gmax) # int32 0/1
184
+ token = min(BIG + mask*(idx4 - BIG), -1)
185
+ token = pad(token, to=32, value=0) # uint32 [1,1,32]
186
+
187
+ ``doc/optimized_full_model/probes/distributed_argmax_probe.py`` measures the
188
+ two against each other at the shipped shape, trace-captured, median of 100:
189
+ **1.1432 ms baseline against 0.6275 ms, 1.82x**. Five things in that spelling
190
+ are load-bearing and were each established on the device, not assumed:
191
+
192
+ * **The local maximum must come from ``ttnn.gather``, not ``ttnn.max``.**
193
+ ``ttnn.max`` over the 37984-wide shard costs 0.494 ms -- more than the
194
+ ``ttnn.argmax`` over the same tensor (0.371 ms). ``ttnn.gather`` at the
195
+ index the argmax already produced costs 0.059 ms. That single substitution
196
+ is the difference between 1.05x and 1.82x.
197
+ * **Only the live user rows are reduced.** The logit tile is logically 32
198
+ rows because ``ttnn.sampling`` addresses 32 slots, but at batch ``B`` the
199
+ other ``32-B`` are zero-logit padding and reduce to token 0 by
200
+ construction. ``ttnn.argmax``'s kernel compares scalar-wise on a
201
+ data-movement RISC, so the cost is linear in rows: the whole reduction is
202
+ **631.6 us over 32 rows and 250.8 us over 1**, and the ``ttnn.pad(value=0)``
203
+ that restores the 32 slots writes back exactly the values the 32-row
204
+ reduction produced. ``argmax_outer_dim_probe.py`` checks that on the
205
+ device rather than asserting it.
206
+ * **Untilize before the argmax.** ``ttnn.argmax``'s multicore path needs
207
+ ROW_MAJOR; the TILE path is single-core and the whole leg becomes 23.25 ms.
208
+ The untilize itself is 0.075 ms.
209
+ * **Indices are INT32 end to end.** FLOAT32 elementwise rounds an index
210
+ through bf16 (36885 -> 36864), and ``ttnn.where`` on int32 operands returns
211
+ bit garbage -- hence the arithmetic select ``BIG + mask*(idx-BIG)`` rather
212
+ than a ``where``. ``ttnn.gather`` in turn demands a UINT32 index, which is
213
+ exactly what ``ttnn.argmax`` emits, so no cast happens on that edge.
214
+ * **The cross-die reduction is a ``min`` over masked indices, never a sum.**
215
+ On an exact tie both lanes survive the mask; ``sum(mask*idx)`` would add
216
+ the two indices together, ``min`` keeps the lower one. Because the dies own
217
+ contiguous ascending vocabulary ranges and ``ttnn.argmax`` returns the
218
+ first occurrence within a die, that is precisely ``torch.argmax``'s
219
+ first-maximal rule. The probe checks it with crafted cross-die, within-die
220
+ and triple ties, and checks the first-occurrence property of ``ttnn.argmax``
221
+ itself.
222
+
223
+ **Output contract.** The base writes the token into the caller's
224
+ ``tt_out_tok`` via ``ttnn.argmax(output_tensor=...)`` and returns it. The
225
+ traced decode loop feeds that same buffer back as the next token's input, so
226
+ returning a *new* tensor would silently break token feedback
227
+ (``models/common/sampling/generator.py::_validate_trace_inputs`` checks
228
+ identity, and this model's trace binds ``token`` as both sampler output and
229
+ model input). The override therefore ends in ``ttnn.copy`` into the caller's
230
+ buffer and returns that exact object -- same dtype (uint32), layout
231
+ (ROW_MAJOR), shape and buffer address.
232
+
233
+ **Fallback.** The fast path is only taken when the reduction it performs is
234
+ provably the same function as the base's. It falls back to
235
+ ``super()._sample_argmax`` whenever ``valid_vocab_size < vocab_size`` (a
236
+ padded vocabulary needs the invalid tail masked *before* the local argmax,
237
+ which ``_mask_invalid_vocab_logits`` /
238
+ ``_can_slice_valid_vocab_for_argmax`` do around the base's gather and this
239
+ path does not reproduce), whenever any invalid-vocab mask buffer is present,
240
+ on a single device, or when the logits do not arrive as an exact even shard.
241
+ For this model ``valid_vocab_size == vocab_size == 151936 == 4*37984``, so
242
+ the fast path is what runs -- but a token id >= the real vocabulary stays
243
+ impossible either way, because a padded vocabulary never reaches it.
244
+
245
+ ------------------------------------------------------------------------
246
+ ``_argmax_all_gather`` -- no ``Topology::Linear`` + ``num_workers_per_link=1``
247
+ ------------------------------------------------------------------------
248
+
249
+ Still overridden, still needed: the split top-k/top-p path is live for any
250
+ request with ``top_k > 1`` or ``top_p > 0`` (``sample_split``), and
251
+ ``_sample_argmax``'s fallback branch above uses it too.
252
+
253
+ ``ttnn.experimental.all_gather_async`` trips a
254
+ BRISC ``ASSERT`` in ``minimal_default_writer.cpp`` when it is given
255
+ ``topology=Topology::Linear`` **together with** ``num_workers_per_link=1``.
256
+ Neither alone does it; the pair does, at any width. The full A/B matrix is
257
+ ``doc/full_model/watcher_ab.log`` and the model-free reproducer is
258
+ ``doc/full_model/probes/ccl_watcher_ab.py --leg linear_workers1``.
259
+
260
+ ``Sampling1D._argmax_all_gather`` walks straight into that pair on any mesh
261
+ smaller than T3K. Its first branch -- Ring, no barrier -- is guarded by
262
+ ``default_topology(mesh) == Topology.Ring``, which is **False** on this 1x4
263
+ Blackhole mesh, so the branch is unreachable here. The fallback then runs
264
+ ``_get_argmax_all_gather_config``, which forces ``Topology.Linear`` for any
265
+ mesh under 8 devices, and the call below it hardcodes
266
+ ``num_workers_per_link=1``. Linear + 1 worker: exactly the tripping pair.
267
+
268
+ The decoder layer's own two all-reduces have been watcher-clean for four
269
+ stages, and the reason is visible in the same matrix: the layer never passes
270
+ ``num_workers_per_link`` at all, so the op picks its default. This override
271
+ does the same thing -- same op, same ``dim``, same semaphores, same
272
+ ``Topology.Ring`` the layer uses, and **no tuning knobs pinned**. The
273
+ matrix's ``sampler_shape_default_knobs`` leg is this exact call at this exact
274
+ shape, and it is clean.
275
+
276
+ This is a local workaround for an upstream bug, not a fix for it. Both
277
+ reports (the op, and ``sampling_1d.py``'s unreachable Ring branch) still
278
+ stand and should still be filed; this subclass just means stage 05 does not
279
+ ship an unchecked-but-violated device invariant while they are open. When
280
+ the op is fixed, delete this class and pass ``Sampling1D`` directly.
281
+
282
+ Subclassing is the seam because ``Sampling1D.from_config`` builds through
283
+ ``object.__new__(cls)`` and ``_bind_strategy`` binds
284
+ ``self._pre_argmax_gather = self._argmax_all_gather`` by attribute lookup on
285
+ the instance -- so the override is what gets bound. **No shared code is
286
+ edited.**
287
+ """
288
+
289
+ def _argmax_all_gather(self, logits):
290
+ cfg = self.config
291
+ return ttnn.experimental.all_gather_async(
292
+ logits,
293
+ persistent_output_buffer=None,
294
+ dim=3,
295
+ multi_device_global_semaphore=cfg.tt_ccl.get_and_cycle_ag_semaphore_handles(),
296
+ barrier_semaphore=cfg.tt_ccl.get_and_cycle_barrier_semaphore_handle(),
297
+ num_links=cfg.num_argmax_gather_links,
298
+ memory_config=logits.memory_config(),
299
+ topology=cfg.ag_topology,
300
+ # Deliberately NOT passing chunks_per_sync / num_workers_per_link /
301
+ # num_buffers_per_channel. Pinning num_workers_per_link=1 is the half
302
+ # of the tripping pair we control. See the class docstring.
303
+ )
304
+
305
+ # -- distributed argmax ---------------------------------------------------
306
+
307
+ def _distributed_argmax_local_vocab(self):
308
+ """Per-die vocabulary width if the distributed argmax applies, else ``None``.
309
+
310
+ Every condition here is a condition under which the reduction below is
311
+ *provably* the same function as ``Sampling1D._sample_argmax``'s. Anything
312
+ else falls back to the base implementation rather than being approximated.
313
+ """
314
+ cfg = self.config
315
+ if getattr(self, "_invalid_vocab_mask", None) is not None:
316
+ return None
317
+ if getattr(self, "_invalid_vocab_tail_mask", None) is not None:
318
+ return None
319
+ valid = cfg.valid_vocab_size if cfg.valid_vocab_size is not None else cfg.vocab_size
320
+ if valid != cfg.vocab_size:
321
+ # A padded vocabulary needs the invalid tail masked before the *local*
322
+ # argmax, which this path does not do. The base masks/slices around
323
+ # its full gather and stays correct; use it.
324
+ return None
325
+ num_devices = cfg.mesh_device.get_num_devices()
326
+ if num_devices < 2 or cfg.vocab_size % num_devices != 0:
327
+ return None
328
+ local = cfg.vocab_size // num_devices
329
+ if local % ttnn.TILE_SIZE != 0:
330
+ return None
331
+ return local
332
+
333
+ #: Live user rows in the sampler's 32-slot logit tile. ``None`` means "all 32"
334
+ #: and reproduces the pre-row-slicing behaviour exactly. The model sets it to
335
+ #: its own ``max_batch_size`` (1 by default) -- see ``_sample_argmax``, and
336
+ #: ``doc/optimized_full_model/probes/argmax_outer_dim_probe.py`` for why it is
337
+ #: worth 2.5x on the whole sampler.
338
+ _dist_active_rows = None
339
+
340
+ def _distributed_argmax_active_rows(self, slots: int) -> int:
341
+ rows = self._dist_active_rows
342
+ if rows is None:
343
+ return int(slots)
344
+ return max(1, min(int(rows), int(slots)))
345
+
346
+ # -- sampling penalties ---------------------------------------------------
347
+ #
348
+ # ``Sampling1D`` has no penalty stage at all, and the vLLM TT plugin does not
349
+ # route penalised requests to host sampling (``platform.py`` sends ``min_p``,
350
+ # ``bad_words``, ``logit_bias``, ``allowed_token_ids``, ``min_tokens``,
351
+ # ``prompt_logprobs`` and structured output to the host sampler -- penalties
352
+ # are deliberately *not* in that list). It packs all three into
353
+ # ``TTSamplingParams`` and hands the model the token history it needs
354
+ # (``model_runner.py``: ``prompt_tokens`` / ``output_tokens`` are populated
355
+ # "if penalties are needed (decode only)"), expecting the model's on-device
356
+ # sampler to apply them. This is that stage.
357
+ #
358
+ # ------------------------------------------------------------------------
359
+ # The shard-boundary problem, and why this spelling cannot get it wrong
360
+ # ------------------------------------------------------------------------
361
+ #
362
+ # Logits are column-parallel: die ``d`` holds vocabulary ids
363
+ # ``d*37984 .. d*37984+37983`` of the 151936, contiguous and ascending -- the
364
+ # same decomposition ``load_device_buffers`` above builds ``_dist_die_offset``
365
+ # from, and ``_dist_local_vocab`` is reused here rather than recomputed. A
366
+ # penalty is keyed by a **global** token id, so for id ``t`` only die
367
+ # ``t // 37984`` may touch column ``t % 37984``; penalising local index
368
+ # ``t % 37984`` on the *other three* dies would silently penalise three
369
+ # unrelated tokens and produce plausible-looking wrong output rather than an
370
+ # error.
371
+ #
372
+ # This stage never does that arithmetic in a kernel. The penalty operands are
373
+ # built on the host as **full-vocabulary** ``[1, 1, 32, 151936]`` tensors --
374
+ # indexed by global id, which is the only frame in which a penalty is
375
+ # defined -- and handed to the device through
376
+ # ``ttnn.ShardTensorToMesh(dim=-1)``, the *same* mapper and the same even
377
+ # 4-way split the logits themselves were produced under by the
378
+ # column-parallel LM head. Column ``t`` of the host tensor therefore lands on
379
+ # exactly the die and exactly the local column that holds logit ``t``, by
380
+ # construction rather than by a computed index. Every op below is
381
+ # elementwise between two tensors with identical per-die shapes, so no op
382
+ # ever needs to know a global id.
383
+ #
384
+ # The identity is *checked* rather than assumed:
385
+ # ``probes/penalty_shard_boundary_probe.py`` penalises one token in die 0's
386
+ # range and one in die 3's, and asserts both moved and that the same local
387
+ # index on the other dies did not.
388
+ #
389
+ # ------------------------------------------------------------------------
390
+ # The arithmetic
391
+ # ------------------------------------------------------------------------
392
+ #
393
+ # vLLM's ``model_executor/layers/utils.py::apply_penalties`` is the contract,
394
+ # and its order is load-bearing -- repetition first, on the raw logit:
395
+ #
396
+ # repetition p (over prompt+output): x = x/p if x > 0 else x*p
397
+ # frequency f (over output): x -= f * count(t in output)
398
+ # presence q (over output): x -= q * (count(t in output) > 0)
399
+ #
400
+ # The repetition rule is sign-dependent, so it is *not* expressible as an
401
+ # additive delta. It is spelled as a per-column multiplicative factor whose
402
+ # two branches are both uploaded:
403
+ #
404
+ # pos = gtz(x) # 1.0 where x > 0, else 0.0
405
+ # factor = rep_neg + pos * rep_dif # rep_neg = p, rep_dif = 1/p - p
406
+ # x = x * factor
407
+ # x = x - add_delta # f*count + q*presence, host-summed
408
+ #
409
+ # For a column no row penalises, the host writes ``rep_neg = 1.0``,
410
+ # ``rep_dif = 0.0``, ``add_delta = 0.0``: ``x * 1.0 - 0.0`` is **bit-exact**
411
+ # in bf16, so an unpenalised token is not merely close to unchanged, it is
412
+ # unchanged. That is what makes the cross-die non-perturbation claim a
413
+ # property of the arithmetic and not of a tolerance.
414
+ #
415
+ # Per-row isolation is likewise structural: the operands are ``[1,1,32,V]``
416
+ # and every op is elementwise, so row *i*'s columns are only ever combined
417
+ # with row *i*'s logits. Padding slots get the neutral row and are untouched.
418
+ #
419
+ # Baking the per-row scalars (p, 1/p, f, q) into the full-width tensors on
420
+ # the host, rather than broadcasting a ``[1,1,32,1]`` scalar column on
421
+ # device, costs one more upload but removes every H-broadcast from the traced
422
+ # graph -- and the host is rebuilding these rows anyway, because vLLM re-sends
423
+ # the whole token history each step.
424
+ #
425
+ # ------------------------------------------------------------------------
426
+ # Fast path
427
+ # ------------------------------------------------------------------------
428
+ #
429
+ # ``_penalty_mode`` is a *graph* property, not a value: 0 means the ops below
430
+ # are not in the captured trace at all, so an unpenalised request pays
431
+ # nothing -- no op, no buffer, no upload. Bit 0 is the repetition stage and
432
+ # bit 1 the additive stage, and they are independent, so a repetition-only
433
+ # request never pays for the additive tensor. The generator releases and
434
+ # re-captures the decode traces when the mode changes, exactly as it already
435
+ # does when ``_sampling_stochastic`` flips between the argmax and split
436
+ # strategies.
437
+
438
+ #: Bitmask: 1 = repetition stage in the graph, 2 = frequency/presence stage.
439
+ _penalty_mode = 0
440
+ _penalty_rep_neg = None
441
+ _penalty_add = None
442
+
443
+ def penalty_buffer_shape(self) -> tuple[int, int]:
444
+ """``(slots, vocab_size)`` the host-side penalty operands must have."""
445
+ cfg = self.config
446
+ return int(cfg.max_batch_size), int(cfg.vocab_size)
447
+
448
+ def penalty_shard_geometry(self) -> tuple[int, int]:
449
+ """``(num_devices, local_vocab)`` -- the split the operands must be staged in.
450
+
451
+ The **same** decomposition ``load_device_buffers`` builds
452
+ ``_dist_die_offset`` from, read off the same config rather than
453
+ recomputed, so the staging path and the distributed argmax cannot drift
454
+ apart.
455
+ """
456
+ cfg = self.config
457
+ devices = cfg.mesh_device.get_num_devices()
458
+ vocab = int(cfg.vocab_size)
459
+ if vocab % devices:
460
+ raise RuntimeError(f"penalties need an even column-parallel split; {vocab} % {devices} != 0")
461
+ return devices, vocab // devices
462
+
463
+ def allocate_penalty_buffers(self, mode: int) -> None:
464
+ """Allocate/free the per-stage operands for ``mode``.
465
+
466
+ Called by the generator **outside** any trace capture -- ``ttnn.from_torch``
467
+ inside ``begin_trace_capture`` raises and leaves the capture open (stage-04
468
+ ``work_log.md`` §6), which is the same reason ``load_device_buffers``
469
+ builds ``_dist_die_offset`` eagerly.
470
+ """
471
+ mode = int(mode)
472
+ if mode == self._penalty_mode:
473
+ return
474
+ cfg = self.config
475
+ slots, vocab = self.penalty_buffer_shape()
476
+ num_devices = cfg.mesh_device.get_num_devices()
477
+ if mode and vocab % num_devices != 0:
478
+ raise RuntimeError(f"penalties need an even column-parallel vocabulary split; {vocab} % {num_devices} != 0")
479
+
480
+ def _alloc(fill: float):
481
+ return ttnn.from_torch(
482
+ torch.full((1, 1, slots, vocab), fill, dtype=torch.bfloat16),
483
+ dtype=ttnn.bfloat16,
484
+ layout=ttnn.TILE_LAYOUT,
485
+ device=cfg.mesh_device,
486
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
487
+ mesh_mapper=ttnn.ShardTensorToMesh(cfg.mesh_device, dim=-1),
488
+ )
489
+
490
+ for want, names, fills in (
491
+ (mode & 1, ("_penalty_rep_neg",), (1.0,)),
492
+ (mode & 2, ("_penalty_add",), (0.0,)),
493
+ ):
494
+ for name, fill in zip(names, fills):
495
+ current = getattr(self, name, None)
496
+ if want and current is None:
497
+ setattr(self, name, _alloc(fill))
498
+ elif not want and current is not None:
499
+ ttnn.deallocate(current, True)
500
+ setattr(self, name, None)
501
+ self._penalty_mode = mode
502
+
503
+ def penalty_device_buffers(self) -> dict:
504
+ """The live operands, keyed by name; the generator uploads into these."""
505
+ return {"rep_neg": self._penalty_rep_neg, "add": self._penalty_add}
506
+
507
+ def _apply_penalties(self, logits):
508
+ """``logits`` -> penalised logits, or ``logits`` itself when the mode is 0.
509
+
510
+ Returns ``(tensor, is_new)``; the caller deallocates when ``is_new``.
511
+ """
512
+ mode = self._penalty_mode
513
+ if not mode:
514
+ return logits, False
515
+ if int(logits.shape[-1]) != self.penalty_buffer_shape()[1] // self.config.mesh_device.get_num_devices():
516
+ # Already gathered, or some shape this stage was not built for. The
517
+ # penalty operands are per-die shards; refusing is the only safe
518
+ # answer, because applying them at the wrong width would penalise
519
+ # the wrong tokens.
520
+ raise RuntimeError(
521
+ f"penalty operands are per-die shards of width "
522
+ f"{self.penalty_buffer_shape()[1] // self.config.mesh_device.get_num_devices()}, "
523
+ f"got logits of width {int(logits.shape[-1])}"
524
+ )
525
+ out = logits
526
+ if mode & 1:
527
+ # ``rep_dif`` (= 1/p - p) is derived **on device** rather than
528
+ # uploaded. It used to be a second full-width operand, and staging one
529
+ # of those costs 2.049 ms of host time per decode step -- more than
530
+ # every device op in this stage put together. ``ttnn.reciprocal`` of
531
+ # the operand gives the same thing for free, because the operand is
532
+ # ``p`` at penalised columns and exactly ``1.0`` everywhere else.
533
+ #
534
+ # This is only allowed to be here because ``reciprocal(1.0)`` is
535
+ # **exactly** 1.0 on this device -- checked, not assumed
536
+ # (``penalty_shard_boundary_probe.py``'s reference and
537
+ # bit-identity legs both fail if it is not). That is what keeps the
538
+ # unpenalised column at ``x * 1.0 - 0.0``, i.e. bit-exact, which is
539
+ # the whole cross-die non-perturbation argument. At *penalised*
540
+ # columns the LLK reciprocal differs from a host-computed ``1/p`` by
541
+ # up to about one bf16 ulp (p=1.05: 0.95703 against 0.95313), which is
542
+ # inside the accuracy the bf16 operand already has.
543
+ inv = ttnn.reciprocal(self._penalty_rep_neg)
544
+ rep_dif = ttnn.subtract(inv, self._penalty_rep_neg)
545
+ ttnn.deallocate(inv)
546
+ # gtz, not a where: ttnn.where on this path is the op the argmax
547
+ # override already avoids, and gtz is a single unary.
548
+ pos = ttnn.gtz(out)
549
+ scaled = ttnn.multiply(pos, rep_dif)
550
+ ttnn.deallocate(pos)
551
+ ttnn.deallocate(rep_dif)
552
+ factor = ttnn.add(scaled, self._penalty_rep_neg)
553
+ ttnn.deallocate(scaled)
554
+ out = ttnn.multiply(out, factor)
555
+ ttnn.deallocate(factor)
556
+ if mode & 2:
557
+ penalised = ttnn.subtract(out, self._penalty_add)
558
+ if out is not logits:
559
+ ttnn.deallocate(out)
560
+ out = penalised
561
+ return out, True
562
+
563
+ def decode_forward(self, logits, **kwargs):
564
+ """Penalty stage, then ``Sampling1D``'s own routing -- unchanged.
565
+
566
+ Overriding here rather than in each strategy means both the argmax path
567
+ and the top-k/top-p split path get penalties from one place, applied
568
+ **before** any selection, which is the only order that is correct.
569
+ """
570
+ penalised, is_new = self._apply_penalties(logits)
571
+ try:
572
+ return super().decode_forward(penalised, **kwargs)
573
+ finally:
574
+ if is_new:
575
+ ttnn.deallocate(penalised)
576
+
577
+ def load_device_buffers(self):
578
+ """Base buffers, plus the per-die vocabulary offset the reduction adds.
579
+
580
+ Built here rather than lazily in ``_sample_argmax`` because the first
581
+ ``_sample_argmax`` may already be inside ``begin_trace_capture``, and
582
+ ``ttnn.from_torch`` inside a capture raises and leaves the capture open.
583
+ """
584
+ already_loaded = self._device_buffers_loaded
585
+ super().load_device_buffers()
586
+ if already_loaded and getattr(self, "_dist_die_offset", None) is not None:
587
+ return
588
+ local_vocab = self._distributed_argmax_local_vocab()
589
+ if local_vocab is None:
590
+ self._dist_die_offset = None
591
+ return
592
+ cfg = self.config
593
+ num_devices = cfg.mesh_device.get_num_devices()
594
+ # ``_dist_active_rows`` rows, not ``cfg.max_batch_size``: the reduction
595
+ # below runs over the live user rows only and the shapes must match
596
+ # exactly, because an H-broadcast here would silently re-expand the
597
+ # result back to 32 rows. See ``_sample_argmax``.
598
+ rows = self._distributed_argmax_active_rows(cfg.max_batch_size)
599
+ offsets = (
600
+ (
601
+ torch.arange(num_devices, dtype=torch.int64)
602
+ .reshape(1, 1, 1, num_devices)
603
+ .expand(1, 1, rows, num_devices)
604
+ * local_vocab
605
+ )
606
+ .contiguous()
607
+ .to(torch.int32)
608
+ )
609
+ # Sharded on the last dim: die d holds the single column ``d*local_vocab``.
610
+ self._dist_die_offset = ttnn.from_torch(
611
+ offsets,
612
+ dtype=ttnn.int32,
613
+ layout=ttnn.TILE_LAYOUT,
614
+ device=cfg.mesh_device,
615
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
616
+ mesh_mapper=ttnn.ShardTensorToMesh(cfg.mesh_device, dim=-1),
617
+ )
618
+ self._dist_local_vocab = local_vocab
619
+
620
+ def _sample_argmax(self, logits, tt_out_tok):
621
+ """Distributed argmax: reduce per die, all-gather 4 candidates, reduce again.
622
+
623
+ Honours ``Sampling1D._sample_argmax``'s contract exactly -- writes the
624
+ caller's ``tt_out_tok`` in place and returns ``(tt_out_tok, None)``; the
625
+ argmax path never emits logprobs. See the class docstring for why each
626
+ step is spelled the way it is and for the measured 1.82x.
627
+ """
628
+ self.load_device_buffers()
629
+ die_offset = getattr(self, "_dist_die_offset", None)
630
+ if die_offset is None or int(logits.shape[-1]) != self._dist_local_vocab:
631
+ # Not an even per-die shard (already gathered, padded vocab, 1x1
632
+ # mesh, ...): the base path is the one that is still correct.
633
+ return super()._sample_argmax(logits, tt_out_tok)
634
+
635
+ # -- per-die reduction, over this die's own columns only ---------------
636
+ # ROW_MAJOR because ttnn.argmax's TILE path is single-core (23 ms).
637
+ rm = ttnn.untilize(logits, use_multicore=True)
638
+ # **Reduce the live user rows, not the padding.** ``decode_terminal``
639
+ # hands the sampler a logically-32-row tile because ``ttnn.sampling``
640
+ # addresses 32 fixed slots, but at batch B only the first B rows carry a
641
+ # user: the rest are the zero rows ``ttnn.pad(..., value=0.0)`` put on the
642
+ # pre-head hidden, and ``lm_head`` has no bias, so their logits are exactly
643
+ # zero. ``ttnn.argmax``'s multicore kernel does the comparison as a scalar
644
+ # C++ ``>`` loop on the RISCV_1 *data-movement* core -- 32 x 37984 values
645
+ # over 110 cores is ~11k compares each, and that, not the 32-round
646
+ # semaphore barrier, is why the op costs 366 us and sits 75x off
647
+ # bandwidth. Dropping the padding rows drops the work proportionally.
648
+ #
649
+ # Measured standalone at the shipped shape
650
+ # (``doc/optimized_full_model/probes/argmax_outer_dim_probe.py``,
651
+ # trace-captured, median of 60; the harness floor is ~58 us):
652
+ #
653
+ # argmax over 32 rows 371.1 us
654
+ # argmax over 32 rows, keepdim=False 309.3 (one barrier, not 32)
655
+ # ROW_MAJOR slice to 1 row + argmax 58.0 (i.e. at the floor)
656
+ # whole reduction, 32 rows 631.6
657
+ # whole reduction, 1 row 250.8 **2.52x**
658
+ #
659
+ # ``keepdim=False`` is a real but small effect and is *not* taken: it buys
660
+ # 62 us on its own and nothing at all once the rows are sliced (251.0 vs
661
+ # 250.8), while costing a ``[1,1,B] -> [1,1,B,1]`` reshape.
662
+ #
663
+ # The substitution is exact, not an approximation. The probe's
664
+ # ``padding_rows_produce_token_zero`` leg checks on the device that a
665
+ # zero logit row reduces to token **0** on the shipped 32-row path -- all
666
+ # four dies tie at 0.0, so the masked ``min`` keeps global index 0 -- which
667
+ # is precisely the value the ``ttnn.pad`` below writes back.
668
+ slots = int(rm.shape[-2])
669
+ active = self._distributed_argmax_active_rows(slots)
670
+ if active < slots:
671
+ live = ttnn.slice(rm, [0, 0, 0, 0], [1, 1, active, self._dist_local_vocab])
672
+ ttnn.deallocate(rm)
673
+ rm = live
674
+ local_idx = ttnn.argmax(rm, dim=-1, keepdim=True) # uint32 RM [1,1,B,1]
675
+ # ttnn.gather, NOT ttnn.max: 0.059 ms against 0.494 ms. This is the win.
676
+ local_max = ttnn.to_layout(ttnn.gather(rm, dim=-1, index=local_idx), ttnn.TILE_LAYOUT)
677
+ ttnn.deallocate(rm)
678
+ # INT32, not FLOAT32: fp32 elementwise rounds the index through bf16.
679
+ local_idx_i32 = ttnn.to_layout(ttnn.typecast(local_idx, ttnn.int32), ttnn.TILE_LAYOUT)
680
+ ttnn.deallocate(local_idx)
681
+ global_idx = ttnn.add(local_idx_i32, die_offset)
682
+ ttnn.deallocate(local_idx_i32)
683
+
684
+ # -- gather 4 candidates, not the whole vocabulary ---------------------
685
+ vals4 = self._argmax_all_gather(local_max) # bf16 [1,1,B,4]
686
+ idx4 = self._argmax_all_gather(global_idx) # int32 [1,1,B,4]
687
+ ttnn.deallocate(local_max)
688
+ ttnn.deallocate(global_idx)
689
+
690
+ # -- cross-die reduction ----------------------------------------------
691
+ gmax = ttnn.max(vals4, dim=-1, keepdim=True)
692
+ mask = ttnn.typecast(ttnn.eq(vals4, gmax), ttnn.int32) # 0/1
693
+ # NOT sum(mask*idx): on a tie that adds the tied indices together.
694
+ # BIG + mask*(idx-BIG) sends losers to BIG and leaves every tied winner at
695
+ # its own global index, so min() keeps the lowest -- the first-maximal one,
696
+ # because die ranges ascend.
697
+ sel = ttnn.add(ttnn.multiply(mask, ttnn.subtract(idx4, _DIST_ARGMAX_BIG)), _DIST_ARGMAX_BIG)
698
+ token = ttnn.min(sel, dim=-1, keepdim=False) # int32 TILE
699
+ for scratch in (vals4, idx4, gmax, mask, sel):
700
+ ttnn.deallocate(scratch)
701
+
702
+ # -- match ttnn.argmax's output contract: UINT32 / ROW_MAJOR ------------
703
+ token = ttnn.typecast(ttnn.to_layout(token, ttnn.ROW_MAJOR_LAYOUT), ttnn.uint32)
704
+ if active < slots:
705
+ # Restore the 32-slot vector ``tt_out_tok`` is. 0 is not a convenient
706
+ # filler, it is the token the shipped 32-row reduction *already*
707
+ # produces for a padding row (see the comment above the slice), so the
708
+ # buffer's contents are unchanged slot for slot. It also keeps every
709
+ # slot a valid id: ``embed_decode`` runs ``ttnn.embedding`` over all 32
710
+ # before slicing to ``batch``, and an out-of-vocabulary id there would
711
+ # be an out-of-bounds table read.
712
+ padded = ttnn.pad(token, [(0, 0), (0, 0), (0, slots - active)], value=0)
713
+ ttnn.deallocate(token)
714
+ token = padded
715
+ if tt_out_tok is None:
716
+ return token, None
717
+ # Write **into the caller's buffer**. The traced decode loop feeds this
718
+ # exact tensor back as the next token, so the object and its address must
719
+ # survive; returning a new tensor breaks token feedback silently.
720
+ ttnn.copy(ttnn.reshape(token, tt_out_tok.shape), tt_out_tok)
721
+ ttnn.deallocate(token)
722
+ return tt_out_tok, None
723
+
724
+
725
+ def _resolve_precision(precision) -> PrecisionConfig:
726
+ """Accept a ``PrecisionConfig``, a dict, a path to JSON, or ``None``.
727
+
728
+ ``None`` is ``DEFAULT_PRECISION``, so every existing caller keeps the
729
+ shipped policy. The dict and path forms exist so a sweep runner -- and,
730
+ later, the vLLM construction path -- can pass the *artifact*
731
+ (``selected_precision_config.json``) rather than importing the dataclass,
732
+ which is what makes the artifact something the model consumes rather than
733
+ something written next to it.
734
+ """
735
+ if precision is None:
736
+ return DEFAULT_PRECISION
737
+ if isinstance(precision, PrecisionConfig):
738
+ return precision
739
+ if isinstance(precision, dict):
740
+ return PrecisionConfig.from_dict(precision)
741
+ if isinstance(precision, (str, Path)):
742
+ return PrecisionConfig.read_json(precision)
743
+ # A ``PrecisionConfig`` from a *duplicate copy* of ``tt.precision``, which is
744
+ # a real hazard in this tree and not a hypothetical one: ``tt/generator.py``
745
+ # imports ``tt.model`` by absolute path while tests and probes import it
746
+ # relatively, and under pytest's ``--import-mode=importlib`` (this repo's
747
+ # ``addopts``) with no ``models/__init__.py`` the two spellings produce two
748
+ # distinct module objects and therefore two distinct classes. ``isinstance``
749
+ # is then False for an object that is, by every meaning that matters, the
750
+ # right one. Rebuild it through the serialised form rather than refusing it.
751
+ if type(precision).__name__ == "PrecisionConfig" and hasattr(precision, "to_dict"):
752
+ return PrecisionConfig.from_dict(precision.to_dict())
753
+ raise TypeError(f"precision must be a PrecisionConfig, dict, path or None; got {type(precision).__name__}")
754
+
755
+
756
+ def _lm_head_compute_config(device, precision: PrecisionConfig = DEFAULT_PRECISION):
757
+ return ttnn.init_device_compute_kernel_config(
758
+ device.arch(),
759
+ math_fidelity=precision.lm_head_fidelity,
760
+ math_approx_mode=False,
761
+ fp32_dest_acc_en=False,
762
+ packer_l1_acc=True,
763
+ )
764
+
765
+
766
+ class ShardedCheckpoint:
767
+ """Read named tensors out of a sharded safetensors checkpoint on demand.
768
+
769
+ The full checkpoint is 30.5B parameters, ~61 GB in bf16. Materialising it as
770
+ one ``state_dict`` to build a model that uploads it layer by layer would
771
+ need that whole 61 GB of host RAM at once; this reads only the tensors asked
772
+ for, from only the shards that hold them, and holds nothing.
773
+ """
774
+
775
+ def __init__(self, path: str | Path):
776
+ self.path = Path(path)
777
+ index_path = self.path / "model.safetensors.index.json"
778
+ if not index_path.is_file():
779
+ raise FileNotFoundError(f"checkpoint index is missing: {index_path}")
780
+ self.weight_map: dict[str, str] = json.loads(index_path.read_text())["weight_map"]
781
+
782
+ def get(self, name: str) -> torch.Tensor:
783
+ shard = self.weight_map.get(name)
784
+ if shard is None:
785
+ raise KeyError(name)
786
+ with safe_open(self.path / shard, framework="pt") as f:
787
+ return f.get_tensor(name)
788
+
789
+ def layer(self, layer_idx: int) -> dict[str, torch.Tensor]:
790
+ """Every ``model.layers.<i>.*`` tensor, keyed layer-relative."""
791
+ prefix = f"model.layers.{layer_idx}."
792
+ by_shard: dict[str, list[str]] = {}
793
+ for name, shard in self.weight_map.items():
794
+ if name.startswith(prefix):
795
+ by_shard.setdefault(shard, []).append(name)
796
+ if not by_shard:
797
+ raise KeyError(f"no tensors for layer {layer_idx}")
798
+ out: dict[str, torch.Tensor] = {}
799
+ for shard, names in by_shard.items():
800
+ with safe_open(self.path / shard, framework="pt") as f:
801
+ for name in names:
802
+ out[name[len(prefix) :]] = f.get_tensor(name)
803
+ return out
804
+
805
+
806
+ def _validate_mesh(mesh_device) -> None:
807
+ shape = tuple(int(v) for v in mesh_device.shape)
808
+ if shape != MESH_SHAPE:
809
+ raise ValueError(f"Qwen3CoderModel requires mesh {MESH_SHAPE}, got {shape}")
810
+ if mesh_device.get_num_devices() != NUM_DEVICES:
811
+ raise ValueError(f"Qwen3CoderModel requires exactly {NUM_DEVICES} devices")
812
+
813
+
814
+ def _rope_parameters(hf_config) -> dict:
815
+ """``rope_parameters`` on current transformers, ``rope_theta`` on older ones.
816
+
817
+ ``Qwen3MoeConfig`` no longer exposes a top-level ``rope_theta`` attribute --
818
+ reading it raises ``AttributeError`` rather than returning ``None`` -- so
819
+ the dict is the only spelling that works on both.
820
+ """
821
+ params = getattr(hf_config, "rope_parameters", None)
822
+ if params:
823
+ return dict(params)
824
+ return {"rope_theta": hf_config.rope_theta, "rope_type": "default"}
825
+
826
+
827
+ def _rope_type(hf_config) -> str:
828
+ return str(_rope_parameters(hf_config).get("rope_type", "default"))
829
+
830
+
831
+ def _rope_theta(hf_config) -> float:
832
+ return float(_rope_parameters(hf_config)["rope_theta"])
833
+
834
+
835
+ def _rope_tables(hf_config, capacity: int) -> tuple[torch.Tensor, torch.Tensor]:
836
+ """The HF ``(cos, sin)`` tables for positions ``0..capacity-1``.
837
+
838
+ Built here rather than through ``Qwen3MoeRotaryEmbedding`` so that no
839
+ transformers model object is constructed at load time; the formula is the
840
+ default rope (``rope_scaling`` is null in this checkpoint, which
841
+ ``from_checkpoint`` asserts).
842
+ """
843
+ head_dim = int(getattr(hf_config, "head_dim", HEAD_DIM))
844
+ theta = _rope_theta(hf_config)
845
+ inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float64) / head_dim))
846
+ angles = torch.outer(torch.arange(capacity, dtype=torch.float64), inv_freq)
847
+ angles = torch.cat([angles, angles], dim=-1)
848
+ return angles.cos().float(), angles.sin().float()
849
+
850
+
851
+ class Qwen3CoderModel:
852
+ """The 48-layer causal LM over the stage-04 multichip decoder layer."""
853
+
854
+ def __init__(
855
+ self,
856
+ *,
857
+ mesh_device,
858
+ hf_config,
859
+ checkpoint: ShardedCheckpoint,
860
+ max_batch_size: int = DEFAULT_MAX_BATCH_SIZE,
861
+ max_cache_len: int = MAX_CONTEXT,
862
+ num_layers: int = NUM_LAYERS,
863
+ page_block_size: int = DEFAULT_PAGE_BLOCK_SIZE,
864
+ rope_cache_len: int = DEFAULT_ROPE_CACHE_LEN,
865
+ precision: "PrecisionConfig | dict | str | Path | None" = None,
866
+ ) -> None:
867
+ _validate_mesh(mesh_device)
868
+ if not 1 <= int(max_batch_size) <= 32:
869
+ # nlp_create_qkv_heads_decode_device_operation.cpp:51 asserts
870
+ # num_users <= 32; a TTNN op limit, unchanged by TP.
871
+ raise ValueError(f"max_batch_size must be in [1,32], got {max_batch_size}")
872
+ if not 1 <= int(num_layers) <= int(hf_config.num_hidden_layers):
873
+ raise ValueError(f"num_layers must be in [1,{hf_config.num_hidden_layers}]")
874
+ if not 1 <= int(max_cache_len) <= int(hf_config.max_position_embeddings):
875
+ raise ValueError(f"max_cache_len must be in [1,{hf_config.max_position_embeddings}]")
876
+ if int(hf_config.hidden_size) != HIDDEN_SIZE or int(hf_config.vocab_size) != VOCAB_SIZE:
877
+ raise ValueError("HF config does not match the Qwen3-Coder-30B-A3B full-model contract")
878
+ if bool(hf_config.tie_word_embeddings):
879
+ raise ValueError("this checkpoint has an untied lm_head; tied weights would be a different contract")
880
+ if _rope_type(hf_config) != "default":
881
+ raise ValueError(f"rope_type {_rope_type(hf_config)!r} is not supported by this port's rotary tables")
882
+
883
+ # The precision policy, resolved once and then read by every builder and
884
+ # every forward below. ``None`` -> ``DEFAULT_PRECISION``, the shipped
885
+ # stage-06 policy, so a caller that says nothing gets exactly the model
886
+ # it got before this parameter existed. A ``dict`` or a path is accepted
887
+ # too, so a sweep runner can hand over a ``selected_precision_config.json``
888
+ # without importing the dataclass.
889
+ self.precision = _resolve_precision(precision)
890
+
891
+ self.mesh_device = mesh_device
892
+ self.hf_config = hf_config
893
+ self.max_batch_size = int(max_batch_size)
894
+ self.max_cache_len = int(max_cache_len)
895
+ self.num_layers = int(num_layers)
896
+ self.page_block_size = int(page_block_size)
897
+ self.hidden_size = HIDDEN_SIZE
898
+ self.vocab_size = VOCAB_SIZE
899
+ self.head_dim = int(getattr(hf_config, "head_dim", HEAD_DIM))
900
+ self.rms_norm_eps = float(hf_config.rms_norm_eps)
901
+ # Exact: 151936 = 4 * 37984 and 37984 = 32 * 1187.
902
+ assert self.vocab_size % (32 * NUM_DEVICES) == 0, self.vocab_size
903
+ self.local_vocab_size = self.vocab_size // NUM_DEVICES
904
+
905
+ #: Skip the expert work of decode rows that hold no live request. On by
906
+ #: default and a no-op at ``max_batch_size == 1``; see
907
+ #: ``_decode_active_mask``. ``QWEN3_DECODE_ACTIVE_ROW_GATING=0`` restores
908
+ #: the stage-08 graph exactly, which is what
909
+ #: ``doc/optimized_vllm/probes/inactive_row_gating_probe.py`` A/Bs
910
+ #: against for the token-equality leg.
911
+ self.active_row_gating = os.getenv("QWEN3_DECODE_ACTIVE_ROW_GATING", "1") not in ("0", "", "false", "no")
912
+
913
+ #: Width of the decode graph currently being built or captured.
914
+ #: Equal to ``max_batch_size`` everywhere except inside
915
+ #: ``decode_width_scope``, which the generator opens to capture a
916
+ #: **narrower** decode graph than the configured slot count -- see
917
+ #: ``doc/batch_scaling/README.md``. Every decode-path use of the row
918
+ #: count reads this, not ``max_batch_size``: the embedding slice, the
919
+ #: rotary gather and shard, and the active-row mask. Prefill and the
920
+ #: sampler are untouched -- ``decode_terminal`` pads to the 32 fixed
921
+ #: ``SAMPLING_SLOTS`` regardless, so the sampler never sees the width.
922
+ self.decode_width = self.max_batch_size
923
+
924
+ self.ctx: MeshContext = mesh_context(mesh_device)
925
+ self.config = MeshDecoderConfig.from_hf(hf_config)
926
+ self.global_config: DecoderLayerConfig = self.config.global_config
927
+
928
+ self.embed_tokens = self._build_embedding(checkpoint)
929
+ self.layers: list[MultichipWeights] = self._build_layers(checkpoint)
930
+ self.final_norm, self.final_norm_rm = self._build_final_norm(checkpoint)
931
+ self.lm_head = self._build_lm_head(checkpoint)
932
+
933
+ self.sparsity = build_local_sparsity(mesh_device, self.config.local_moe)
934
+ self.lm_head_compute_config = _lm_head_compute_config(mesh_device, self.precision)
935
+ self.norm_compute_config = _norm_compute_config(mesh_device, self.precision)
936
+ # Set by ``local_logits`` / the sampler-input path once a forward has
937
+ # run, so ``runtime_fallback_audit`` can report the dtypes the terminal
938
+ # path *produced* rather than the ones the config asked for.
939
+ self._observed_logits_dtype = None
940
+ self._observed_sampling_dtype = None
941
+
942
+ self.rope_cache_len = 0
943
+ self.cos_table = None
944
+ self.sin_table = None
945
+ self.ensure_rope_capacity(min(int(rope_cache_len), self.max_cache_len))
946
+
947
+ # ``_WatcherCleanSampling1D`` rather than ``Sampling1D``: same module,
948
+ # same strategies, the force-argmax gather spelled without the pinned
949
+ # ``num_workers_per_link`` that trips the watcher on this mesh. See the
950
+ # class docstring above and ``doc/full_model/watcher_ab.log``.
951
+ self.sampler = _WatcherCleanSampling1D.from_config(
952
+ Sampling1DConfig(
953
+ vocab_size=self.vocab_size,
954
+ valid_vocab_size=self.vocab_size,
955
+ mesh_device=mesh_device,
956
+ tt_ccl=self.ctx.ccl,
957
+ max_batch_size=32,
958
+ max_top_k=32,
959
+ num_gather_links=1,
960
+ sampling_memory_config=ttnn.DRAM_MEMORY_CONFIG,
961
+ allow_force_argmax=True,
962
+ num_argmax_gather_links=1,
963
+ ag_topology=TOPOLOGY,
964
+ # **False, and that is a measurement.** ``Sampling1D``'s comment
965
+ # calls the power-of-two pad a "big device-perf win for
966
+ # non-power-of-2 vocab on the multi-device path". For a per-die
967
+ # shard of 37984 it is the opposite: the pad is to 65536, a 1.73x
968
+ # blow-up of the tensor ``ttnn.topk`` then scans, and
969
+ # ``probes/sampler_probe.py`` measures the whole split path at
970
+ # **11.006 ms padded against 6.151 ms unpadded**, 1.79x, at the
971
+ # shipped logits shape with the sampled token unchanged.
972
+ pad_to_power_of_2=False,
973
+ )
974
+ )
975
+ # ``max_batch_size=32`` above is the *slot* count ``ttnn.sampling`` and
976
+ # ``decode_terminal`` address; this is how many of those slots carry a
977
+ # user. The distributed argmax reduces only those rows -- the rest are the
978
+ # zero-logit padding ``decode_terminal`` adds -- which is worth 2.52x on
979
+ # the whole sampler at batch 1. Set before ``load_device_buffers`` because
980
+ # the per-die offset constant is built to this row count.
981
+ self.sampler._dist_active_rows = self.max_batch_size
982
+ self.sampler.load_device_buffers()
983
+ self.kv_cache: list[KVCache] | None = None
984
+
985
+ # -- construction ---------------------------------------------------------
986
+
987
+ @classmethod
988
+ def from_checkpoint(
989
+ cls,
990
+ checkpoint_path: str | Path,
991
+ *,
992
+ mesh_device,
993
+ max_batch_size: int = DEFAULT_MAX_BATCH_SIZE,
994
+ max_cache_len: int = MAX_CONTEXT,
995
+ num_layers: int = NUM_LAYERS,
996
+ page_block_size: int = DEFAULT_PAGE_BLOCK_SIZE,
997
+ rope_cache_len: int = DEFAULT_ROPE_CACHE_LEN,
998
+ precision: "PrecisionConfig | dict | str | Path | None" = None,
999
+ ) -> "Qwen3CoderModel":
1000
+ checkpoint_path = Path(checkpoint_path)
1001
+ hf_config = AutoConfig.from_pretrained(checkpoint_path)
1002
+ checkpoint = ShardedCheckpoint(checkpoint_path)
1003
+ model = cls(
1004
+ mesh_device=mesh_device,
1005
+ hf_config=hf_config,
1006
+ checkpoint=checkpoint,
1007
+ max_batch_size=max_batch_size,
1008
+ max_cache_len=max_cache_len,
1009
+ num_layers=num_layers,
1010
+ page_block_size=page_block_size,
1011
+ rope_cache_len=rope_cache_len,
1012
+ precision=precision,
1013
+ )
1014
+ gc.collect()
1015
+ return model
1016
+
1017
+ def _build_embedding(self, checkpoint: ShardedCheckpoint) -> ttnn.Tensor:
1018
+ host = checkpoint.get("model.embed_tokens.weight").float()
1019
+ tensor = ttnn.from_torch(
1020
+ host,
1021
+ dtype=self.precision.embedding_dtype,
1022
+ layout=ttnn.ROW_MAJOR_LAYOUT,
1023
+ device=self.mesh_device,
1024
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1025
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
1026
+ )
1027
+ del host
1028
+ gc.collect()
1029
+ return tensor
1030
+
1031
+ def _build_layers(self, checkpoint: ShardedCheckpoint) -> list[MultichipWeights]:
1032
+ from .weight_mapping import convert_layer_weights
1033
+
1034
+ layers = []
1035
+ for layer_idx in range(self.num_layers):
1036
+ sd = checkpoint.layer(layer_idx)
1037
+ torch_weights = convert_layer_weights(sd, self.hf_config)
1038
+ del sd
1039
+ layers.append(
1040
+ upload_multichip_weights(torch_weights, self.mesh_device, self.config, precision=self.precision)
1041
+ )
1042
+ del torch_weights
1043
+ gc.collect()
1044
+ return layers
1045
+
1046
+ def _build_final_norm(self, checkpoint: ShardedCheckpoint):
1047
+ host = checkpoint.get("model.norm.weight").float().reshape(-1)
1048
+ tiled = ttnn.from_torch(
1049
+ host.reshape(1, 1, 1, -1),
1050
+ dtype=self.precision.norm_weight_dtype,
1051
+ layout=ttnn.TILE_LAYOUT,
1052
+ device=self.mesh_device,
1053
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1054
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
1055
+ )
1056
+ # The layout the sharded rms_norm program factory reads; see
1057
+ # ``multichip_decoder.upload_multichip_weights.norm_row_major``.
1058
+ row_major = ttnn.from_torch(
1059
+ host.reshape(1, 1, host.numel() // 32, 32).contiguous(),
1060
+ dtype=self.precision.norm_weight_dtype,
1061
+ layout=ttnn.ROW_MAJOR_LAYOUT,
1062
+ device=self.mesh_device,
1063
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1064
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
1065
+ )
1066
+ return tiled, row_major
1067
+
1068
+ def _build_lm_head(self, checkpoint: ShardedCheckpoint) -> ttnn.Tensor:
1069
+ host = checkpoint.get("lm_head.weight").float().transpose(-2, -1).contiguous()
1070
+ assert tuple(host.shape) == (self.hidden_size, self.vocab_size), tuple(host.shape)
1071
+ tensor = ttnn.from_torch(
1072
+ host.reshape(1, 1, self.hidden_size, self.vocab_size),
1073
+ dtype=self.precision.lm_head_dtype,
1074
+ layout=ttnn.TILE_LAYOUT,
1075
+ device=self.mesh_device,
1076
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1077
+ mesh_mapper=ttnn.ShardTensorToMesh(self.mesh_device, dim=-1),
1078
+ )
1079
+ del host
1080
+ gc.collect()
1081
+ return tensor
1082
+
1083
+ # -- rotary ---------------------------------------------------------------
1084
+
1085
+ def ensure_rope_capacity(self, required_len: int) -> bool:
1086
+ """Grow the device cos/sin tables to cover ``required_len`` positions."""
1087
+ required_len = int(required_len)
1088
+ if required_len <= self.rope_cache_len:
1089
+ return False
1090
+ if required_len > self.max_cache_len:
1091
+ raise ValueError(f"rotary capacity {required_len} exceeds context {self.max_cache_len}")
1092
+ capacity = min(self.max_cache_len, max(32, 1 << (required_len - 1).bit_length()))
1093
+ cos, sin = _rope_tables(self.hf_config, capacity)
1094
+ new = []
1095
+ for host in (cos, sin):
1096
+ new.append(
1097
+ ttnn.from_torch(
1098
+ host.reshape(1, 1, capacity, self.head_dim),
1099
+ dtype=ttnn.bfloat16,
1100
+ layout=ttnn.TILE_LAYOUT,
1101
+ device=self.mesh_device,
1102
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1103
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
1104
+ )
1105
+ )
1106
+ old = (self.cos_table, self.sin_table)
1107
+ self.cos_table, self.sin_table = new
1108
+ for tensor in old:
1109
+ if tensor is not None:
1110
+ ttnn.deallocate(tensor, True)
1111
+ self.rope_cache_len = capacity
1112
+ return True
1113
+
1114
+ def rope_decode_tables(self, rotary_position: ttnn.Tensor):
1115
+ """Per-user ``(cos, sin)`` for one decode step, gathered **on device**.
1116
+
1117
+ ``rotary_position`` is a ``[1, batch]`` uint32 device tensor. The gather
1118
+ is ``ttnn.embedding`` against the replicated cos/sin tables, so the
1119
+ position never leaves the device and the whole thing is capturable; the
1120
+ trace advances ``rotary_position`` itself with ``ttnn.plus_one``.
1121
+
1122
+ Returns the height-sharded ``[1, batch, 1, head_dim]`` pair that
1123
+ ``rotary_embedding_hf(is_decode_mode=True)`` requires -- one core per
1124
+ user, the same ``_head_shard`` layout ``nlp_create_qkv_heads_decode``
1125
+ emits for Q and K.
1126
+ """
1127
+ batch = self.decode_width
1128
+ shard = _head_shard(32, self.head_dim, batch)
1129
+ out = []
1130
+ for table in (self.cos_table, self.sin_table):
1131
+ # [1, batch] -> [1, batch, head_dim] -> [1, 1, batch, head_dim]
1132
+ # -> [1, batch, 1, head_dim], the layout rotary_embedding_hf's decode
1133
+ # factory reads. Same sequence as ``RotarySetup1D.decode_forward``.
1134
+ gathered = ttnn.unsqueeze_to_4D(ttnn.embedding(rotary_position, table, layout=ttnn.TILE_LAYOUT))
1135
+ transposed = ttnn.transpose(gathered, 1, 2)
1136
+ if int(transposed.shape[1]) != batch:
1137
+ trimmed = ttnn.slice(transposed, [0, 0, 0, 0], [1, batch, 1, self.head_dim])
1138
+ ttnn.deallocate(transposed, True)
1139
+ transposed = trimmed
1140
+ out.append(ttnn.interleaved_to_sharded(transposed, shard))
1141
+ ttnn.deallocate(transposed, True)
1142
+ return out[0], out[1]
1143
+
1144
+ def _rope_decode(self, tensor: ttnn.Tensor, cos_sharded, sin_sharded, _token_index):
1145
+ """The ``rope=`` seam handed to ``decoder_layer_decode_multichip``.
1146
+
1147
+ ``_token_index`` is accepted and ignored: the position lives in the
1148
+ cos/sin pair, which is what makes this spelling replayable where the
1149
+ layer's default one is not.
1150
+ """
1151
+ shard = _head_shard(32, self.head_dim, self.decode_width)
1152
+ staged = ttnn.to_memory_config(tensor, shard)
1153
+ rotated = ttnn.experimental.rotary_embedding_hf(staged, cos_sharded, sin_sharded, is_decode_mode=True)
1154
+ ttnn.deallocate(staged, True)
1155
+ out = ttnn.to_memory_config(rotated, ttnn.DRAM_MEMORY_CONFIG)
1156
+ ttnn.deallocate(rotated, True)
1157
+ return out
1158
+
1159
+ # -- KV cache -------------------------------------------------------------
1160
+
1161
+ def allocate_kv_cache(
1162
+ self,
1163
+ *,
1164
+ max_cache_len: int | None = None,
1165
+ num_blocks: int | None = None,
1166
+ page_table: ttnn.Tensor | None = None,
1167
+ ) -> list[KVCache]:
1168
+ """One paged ``KVCache`` per layer, 1 local KV head per die.
1169
+
1170
+ 512 B per token per layer per die -- a quarter of the single-die 2048 --
1171
+ which is what makes the advertised 262144 context fit; see
1172
+ ``config/context_contract.json``.
1173
+ """
1174
+ cache_len = self.max_cache_len if max_cache_len is None else int(max_cache_len)
1175
+ blocks_per_seq = math.ceil(cache_len / self.page_block_size)
1176
+ total_blocks = self.max_batch_size * blocks_per_seq if num_blocks is None else int(num_blocks)
1177
+ local = self.config.local_attention
1178
+ caches = []
1179
+ for _ in range(self.num_layers):
1180
+ k, v = (
1181
+ ttnn.from_torch(
1182
+ torch.zeros(total_blocks, local.num_key_value_heads, self.page_block_size, local.head_dim),
1183
+ dtype=self.precision.kv_cache_dtype,
1184
+ layout=ttnn.TILE_LAYOUT,
1185
+ device=self.mesh_device,
1186
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1187
+ mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
1188
+ )
1189
+ for _ in range(2)
1190
+ )
1191
+ caches.append(KVCache(k=k, v=v, page_table=page_table, block_size=self.page_block_size))
1192
+ return caches
1193
+
1194
+ def ensure_internal_kv_cache(self, page_table: ttnn.Tensor | None = None) -> list[KVCache]:
1195
+ if self.kv_cache is None:
1196
+ self.kv_cache = self.allocate_kv_cache(page_table=page_table)
1197
+ return self.kv_cache
1198
+
1199
+ @staticmethod
1200
+ def bind_page_table(kv_cache: Sequence[KVCache], page_table: ttnn.Tensor | None) -> list[KVCache]:
1201
+ """Point every layer's cache at ``page_table`` in place.
1202
+
1203
+ The page table is a *persistent device tensor* owned by the caller (the
1204
+ generator, or vLLM later). Rebinding mutates the ``KVCache`` records
1205
+ rather than reallocating, so the tensor identity a captured trace
1206
+ recorded is preserved and an unchanged page table costs nothing.
1207
+ """
1208
+ for cache in kv_cache:
1209
+ cache.page_table = page_table
1210
+ return list(kv_cache)
1211
+
1212
+ def reset_kv_cache(self, kv_cache: Sequence[KVCache] | None = None) -> None:
1213
+ selected = self.ensure_internal_kv_cache() if kv_cache is None else kv_cache
1214
+ for cache in selected:
1215
+ ttnn.fill(cache.k, 0.0, memory_config=cache.k.memory_config(), output_tensor=cache.k)
1216
+ ttnn.fill(cache.v, 0.0, memory_config=cache.v.memory_config(), output_tensor=cache.v)
1217
+
1218
+ # -- forward: prefill -----------------------------------------------------
1219
+
1220
+ def embed_prefill(self, tokens: ttnn.Tensor) -> ttnn.Tensor:
1221
+ """``[1, S]`` uint32 -> replicated ``[1, 1, S, 2048]``, no collective."""
1222
+ hidden = ttnn.embedding(
1223
+ tokens,
1224
+ self.embed_tokens,
1225
+ layout=ttnn.TILE_LAYOUT,
1226
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1227
+ dtype=self.precision.activation_dtype,
1228
+ )
1229
+ hidden = ttnn.unsqueeze_to_4D(hidden)
1230
+ return ttnn.reshape(hidden, (1, 1, int(hidden.shape[-2]), self.hidden_size))
1231
+
1232
+ def prefill_hidden(
1233
+ self,
1234
+ tokens: ttnn.Tensor,
1235
+ *,
1236
+ kv_cache: Sequence[KVCache] | None = None,
1237
+ user_id: int = 0,
1238
+ start_pos: int = 0,
1239
+ chunk_page_table=None,
1240
+ fill_page_table=None,
1241
+ fill_len: int | None = None,
1242
+ ) -> ttnn.Tensor:
1243
+ """Run the whole stack over one user's prompt. ``S`` is arbitrary.
1244
+
1245
+ Nothing here constrains ``S``: the collectives scatter on dim 3 (hidden,
1246
+ 2048, fixed), ``attention_prefill`` slices RoPE's tile padding back, and
1247
+ ``moe_prefill_optimized`` pads to its chunk internally and slices back.
1248
+ """
1249
+ caches = self.ensure_internal_kv_cache() if kv_cache is None else kv_cache
1250
+ if len(caches) != self.num_layers:
1251
+ raise ValueError(f"kv_cache has {len(caches)} layers, expected {self.num_layers}")
1252
+ hidden = self.embed_prefill(tokens)
1253
+ seq_len = int(hidden.shape[-2])
1254
+ # A split prefill's suffix occupies absolute positions
1255
+ # [start_pos, start_pos + seq_len), so the tables must cover the END of
1256
+ # the range, not its length.
1257
+ self.ensure_rope_capacity(start_pos + seq_len)
1258
+ # Exactly ``seq_len`` rows, including non-tile-aligned lengths -- the
1259
+ # same shape the single-layer prefill gates pass at S = 33/100/257.
1260
+ # RoPE is applied at ABSOLUTE positions: the suffix of a split prefill
1261
+ # must rotate at [start_pos, start_pos + seq_len), not from 0, or its keys
1262
+ # disagree with the ones already in the cache. Identical to the shipped
1263
+ # slice when start_pos == 0.
1264
+ # When the requested window is the WHOLE table, ``ttnn.slice`` hands
1265
+ # back a view of its input rather than a copy -- as a different Python
1266
+ # object, so an ``is`` guard does not catch it. The deallocate at the
1267
+ # end of this method would then free ``self.cos_table``'s own DRAM, and
1268
+ # the next prefill dies on "Input Tensor is not allocated" rather than
1269
+ # anywhere near here. Identical hazard, and identical fix, to the
1270
+ # one-token case in ``select_prefill_rows``.
1271
+ #
1272
+ # It is reachable whenever ``start_pos == 0`` and the length equals the
1273
+ # rope capacity, which ``ensure_rope_capacity`` rounds to a power of
1274
+ # two -- so a 8192- or 16384-token prompt, or any prefill bucket rung
1275
+ # that is a power of two. Found by the bucket ladder's 8192 rung.
1276
+ whole_table = start_pos == 0 and seq_len >= self.rope_cache_len
1277
+ if whole_table:
1278
+ cos = ttnn.clone(self.cos_table, memory_config=ttnn.DRAM_MEMORY_CONFIG)
1279
+ sin = ttnn.clone(self.sin_table, memory_config=ttnn.DRAM_MEMORY_CONFIG)
1280
+ else:
1281
+ cos = ttnn.slice(self.cos_table, [0, 0, start_pos, 0], [1, 1, start_pos + seq_len, self.head_dim])
1282
+ sin = ttnn.slice(self.sin_table, [0, 0, start_pos, 0], [1, 1, start_pos + seq_len, self.head_dim])
1283
+ for layer_idx in range(self.num_layers):
1284
+ hidden = decoder_layer_prefill_multichip(
1285
+ hidden,
1286
+ self.layers[layer_idx],
1287
+ self.config,
1288
+ self.ctx,
1289
+ cos,
1290
+ sin,
1291
+ self.sparsity,
1292
+ kv_cache=caches[layer_idx],
1293
+ user_id=user_id,
1294
+ precision=self.precision,
1295
+ start_pos=start_pos,
1296
+ chunk_page_table=chunk_page_table,
1297
+ fill_page_table=fill_page_table,
1298
+ fill_len=fill_len,
1299
+ )
1300
+ ttnn.deallocate(cos, True)
1301
+ ttnn.deallocate(sin, True)
1302
+ return hidden
1303
+
1304
+ def prefill_norm(self, hidden: ttnn.Tensor) -> ttnn.Tensor:
1305
+ return ttnn.rms_norm(
1306
+ hidden,
1307
+ weight=self.final_norm,
1308
+ epsilon=self.rms_norm_eps,
1309
+ compute_kernel_config=self.norm_compute_config,
1310
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1311
+ )
1312
+
1313
+ def select_prefill_rows(self, hidden: ttnn.Tensor, rows: Sequence[int]) -> ttnn.Tensor:
1314
+ """Keep only ``rows`` of a ``[1, 1, S, H]`` prefill result."""
1315
+ seq_len = int(hidden.shape[-2])
1316
+ pieces = []
1317
+ for row in rows:
1318
+ if not 0 <= int(row) < seq_len:
1319
+ raise ValueError(f"prefill row {row} is outside [0,{seq_len})")
1320
+ if seq_len == 1:
1321
+ # At a **one-token prompt** the requested slice covers the whole
1322
+ # tensor, and ``ttnn.slice`` then hands back a view of its input
1323
+ # rather than a copy -- as a *different* Python object, so an
1324
+ # ``is`` guard does not catch it. The caller deallocates
1325
+ # ``hidden`` immediately afterwards, leaving the retained row
1326
+ # pointing at freed DRAM; that does not raise, it **segfaults**
1327
+ # in whatever reads it next (the final norm here). Copy instead.
1328
+ # `probes/prompt_len_1_repro.py` is the four-line reproduction.
1329
+ pieces.append(ttnn.clone(hidden, memory_config=ttnn.DRAM_MEMORY_CONFIG))
1330
+ continue
1331
+ pieces.append(
1332
+ ttnn.slice(
1333
+ hidden,
1334
+ [0, 0, int(row), 0],
1335
+ [1, 1, int(row) + 1, self.hidden_size],
1336
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1337
+ )
1338
+ )
1339
+ if len(pieces) == 1:
1340
+ return pieces[0]
1341
+ out = ttnn.concat(pieces, dim=2, memory_config=ttnn.DRAM_MEMORY_CONFIG)
1342
+ for piece in pieces:
1343
+ ttnn.deallocate(piece, True)
1344
+ return out
1345
+
1346
+ def local_logits(self, normed: ttnn.Tensor) -> ttnn.Tensor:
1347
+ """``[1, 1, rows, 2048]`` -> this die's ``[1, 1, rows, 37984]`` logits."""
1348
+ out = ttnn.linear(
1349
+ normed,
1350
+ self.lm_head,
1351
+ compute_kernel_config=self.lm_head_compute_config,
1352
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1353
+ dtype=self.precision.logits_dtype,
1354
+ )
1355
+ # Observed, not asserted: the dtype the produced tensor actually carries.
1356
+ # ``runtime_fallback_audit`` reports it so ``logits_dtype`` is verified
1357
+ # off a real tensor rather than echoed back out of the config -- see the
1358
+ # ``*_observed`` entries there.
1359
+ self._observed_logits_dtype = out.dtype
1360
+ return out
1361
+
1362
+ def gather_logits_to_torch(self, local_logits: ttnn.Tensor, *, valid_rows: int | None = None) -> torch.Tensor:
1363
+ """Host-side full-vocabulary logits. **Not** on the token-out path.
1364
+
1365
+ Used by ``return_all_logits`` prefill checks and the host-sampling
1366
+ compatibility mode only; the measured decode path never calls this.
1367
+ """
1368
+ gathered = ttnn.all_gather(
1369
+ local_logits,
1370
+ dim=3,
1371
+ num_links=1,
1372
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1373
+ topology=TOPOLOGY,
1374
+ )
1375
+ host = ttnn.to_torch(ttnn.get_device_tensors(gathered)[0]).float()
1376
+ ttnn.deallocate(gathered, True)
1377
+ if valid_rows is not None:
1378
+ host = host[..., : int(valid_rows), :]
1379
+ return host[..., : self.vocab_size]
1380
+
1381
+ # -- forward: decode ------------------------------------------------------
1382
+
1383
+ def embed_decode(self, tokens: ttnn.Tensor) -> ttnn.Tensor:
1384
+ """``[1, 1, 1, 32]`` uint32 -> replicated ``[1, 1, batch, 2048]``."""
1385
+ hidden = ttnn.embedding(
1386
+ tokens,
1387
+ self.embed_tokens,
1388
+ layout=ttnn.TILE_LAYOUT,
1389
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1390
+ dtype=self.precision.activation_dtype,
1391
+ )
1392
+ hidden = ttnn.unsqueeze_to_4D(hidden)
1393
+ flat = ttnn.reshape(hidden, (1, 1, int(hidden.shape[-2]), self.hidden_size))
1394
+ if int(flat.shape[-2]) == self.decode_width:
1395
+ return flat
1396
+ sliced = ttnn.slice(
1397
+ flat, [0, 0, 0, 0], [1, 1, self.decode_width, self.hidden_size], memory_config=ttnn.DRAM_MEMORY_CONFIG
1398
+ )
1399
+ ttnn.deallocate(flat, True)
1400
+ return sliced
1401
+
1402
+ @contextlib.contextmanager
1403
+ def decode_width_scope(self, width: int):
1404
+ """Build the decode graph ``width`` rows wide instead of ``max_batch_size``.
1405
+
1406
+ The narrow graph is legal because **nothing in decode binds a user to a
1407
+ slot index except the three per-row inputs** -- ``current_pos``, the
1408
+ rotary position and the page-table row. The KV cache is fully paged:
1409
+ ``paged_update_cache`` and ``paged_scaled_dot_product_attention_decode``
1410
+ both reach the cache only through ``page_table_tensor`` rows and
1411
+ ``cur_pos_tensor`` entries, and neither takes a ``batch_offset``. So a
1412
+ request can be decoded in *any* row provided its page-table row, its
1413
+ position and its token travel with it; no cache page moves.
1414
+
1415
+ What the width actually changes is the amount of work: the expert
1416
+ ``ttnn.sparse_matmul`` visits ``width x local_experts`` slots per layer
1417
+ with ``nnz=None``, the expert tail is dense over the same product, the
1418
+ router runs one ``topk`` per row, and paged SDPA reads one window per
1419
+ row. That cost is paid per row *configured*, not per row live -- see
1420
+ ``doc/optimized_vllm/README.md``'s control curve -- so narrowing the
1421
+ graph is the only lever that removes it.
1422
+
1423
+ Sampling is deliberately outside the scope: ``decode_terminal`` pads to
1424
+ the 32 fixed ``SAMPLING_SLOTS`` whatever the width is, so ``tt_out_tok``
1425
+ keeps its ``[1,1,1,32]`` shape and the sampler's per-slot parameters
1426
+ keep their meaning. Row *i* of a narrow graph is sampling slot *i*.
1427
+ """
1428
+ width = int(width)
1429
+ if not 1 <= width <= self.max_batch_size:
1430
+ raise ValueError(f"decode width must be in [1,{self.max_batch_size}], got {width}")
1431
+ previous = self.decode_width
1432
+ self.decode_width = width
1433
+ try:
1434
+ yield width
1435
+ finally:
1436
+ self.decode_width = previous
1437
+
1438
+ def _decode_active_mask(self, current_pos: ttnn.Tensor):
1439
+ """``[1, 1, batch, 1]`` of 1.0 for live slots and 0.0 for inactive ones.
1440
+
1441
+ A serving decode batch is always the configured ``max_num_seqs`` rows --
1442
+ vLLM pads it so the trace shape is constant -- with inactive slots
1443
+ carrying ``current_pos = -1``. Those rows still embed a token, still run
1444
+ attention, and, critically, still route to a full top-8 of experts, so
1445
+ their ``(row, expert)`` pairs land in ``sparse_matmul``'s sparsity and
1446
+ cost real expert weight reads and real math. Multiplying the routing
1447
+ vector by this mask takes them out of the sparsity instead
1448
+ (``decoder_layer_decode_multichip``).
1449
+
1450
+ **Why it is derived on device rather than passed in.** ``current_pos`` is
1451
+ already a persistent trace input, and the traced graph advances it with
1452
+ ``ttnn.plus_one(..., skip_negative_entries=True)`` -- an inactive row
1453
+ stays at ``-1`` through any number of replays, and a slot only becomes
1454
+ active through a host reinstall of ``current_pos``. So a mask computed
1455
+ from it inside the same graph is correct by construction on every replay,
1456
+ with no extra trace input to refresh and no way for it to go stale. A
1457
+ host-supplied mask would be one more thing that has to be right.
1458
+
1459
+ Returns ``None`` at ``max_batch_size == 1``, where there is no inactive
1460
+ row to skip: the graph is then byte-for-byte the one stage 08 shipped and
1461
+ the single-user headline cannot be perturbed by this change.
1462
+ """
1463
+ if self.decode_width <= 1 or not self.active_row_gating:
1464
+ return None
1465
+ row = ttnn.to_layout(ttnn.reshape(current_pos, (1, 1, 1, self.decode_width)), ttnn.TILE_LAYOUT)
1466
+ # bf16 cannot represent every position exactly at 262144, but it
1467
+ # represents every position's *sign* exactly, and ``gez`` only reads the
1468
+ # sign. -1 -> 0.0, everything >= 0 -> 1.0.
1469
+ as_float = ttnn.typecast(row, ttnn.bfloat16)
1470
+ ttnn.deallocate(row, True)
1471
+ live_row = ttnn.gez(as_float)
1472
+ ttnn.deallocate(as_float, True)
1473
+ mask = ttnn.transpose(live_row, -2, -1)
1474
+ ttnn.deallocate(live_row, True)
1475
+ return mask
1476
+
1477
+ def decode_hidden(
1478
+ self,
1479
+ tokens: ttnn.Tensor,
1480
+ *,
1481
+ current_pos: ttnn.Tensor,
1482
+ rotary_position: ttnn.Tensor,
1483
+ kv_cache: Sequence[KVCache] | None = None,
1484
+ ) -> ttnn.Tensor:
1485
+ caches = self.ensure_internal_kv_cache() if kv_cache is None else kv_cache
1486
+ if len(caches) != self.num_layers:
1487
+ raise ValueError(f"kv_cache has {len(caches)} layers, expected {self.num_layers}")
1488
+ hidden = self.embed_decode(tokens)
1489
+ cos, sin = self.rope_decode_tables(rotary_position)
1490
+ # Computed once per decode step and shared by all 48 layers.
1491
+ active_mask = self._decode_active_mask(current_pos)
1492
+ for layer_idx in range(self.num_layers):
1493
+ hidden = decoder_layer_decode_multichip(
1494
+ hidden,
1495
+ self.layers[layer_idx],
1496
+ self.config,
1497
+ self.ctx,
1498
+ cos,
1499
+ sin,
1500
+ caches[layer_idx],
1501
+ current_pos,
1502
+ 0, # token_index: unused by the rope seam below, see _rope_decode
1503
+ rope=self._rope_decode,
1504
+ precision=self.precision,
1505
+ active_mask=active_mask,
1506
+ )
1507
+ ttnn.deallocate(cos, True)
1508
+ ttnn.deallocate(sin, True)
1509
+ if active_mask is not None:
1510
+ ttnn.deallocate(active_mask, True)
1511
+ return hidden
1512
+
1513
+ def decode_terminal(self, hidden: ttnn.Tensor) -> ttnn.Tensor:
1514
+ """Final norm + column-parallel ``lm_head``, sampler-ready local logits.
1515
+
1516
+ The norm is the layer's own width-sharded decode kernel, and the shard
1517
+ it emits is exactly the width-sharded L1 config the projections read, so
1518
+ crossing into the head costs one sharded-to-interleaved.
1519
+ """
1520
+ normed_sharded = decode_residual_norm(hidden, self.final_norm_rm, self.rms_norm_eps, self.precision)
1521
+ normed = ttnn.sharded_to_interleaved(normed_sharded, ttnn.DRAM_MEMORY_CONFIG)
1522
+ ttnn.deallocate(normed_sharded, True)
1523
+ # ``ttnn.sampling`` addresses 32 fixed user slots, and it compares the
1524
+ # *logical* shapes of its values and indices, so the logits handed to it
1525
+ # must be logically 32 rows and not ``batch`` rows padded to a tile.
1526
+ # The rows are already physically there -- ``batch <= 32`` and decode is
1527
+ # one 32-row tile -- so this only rewrites the logical shape.
1528
+ rows = int(normed.shape[-2])
1529
+ if rows < SAMPLING_SLOTS:
1530
+ padded = ttnn.pad(
1531
+ normed,
1532
+ [(0, 0), (0, 0), (0, SAMPLING_SLOTS - rows), (0, 0)],
1533
+ value=0.0,
1534
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1535
+ )
1536
+ ttnn.deallocate(normed, True)
1537
+ normed = padded
1538
+ logits = self.local_logits(normed)
1539
+ ttnn.deallocate(normed, True)
1540
+ if self.precision.sampling_dtype != logits.dtype:
1541
+ # Equal on the shipped path, so this is dead code at the default and
1542
+ # the traced decode graph is byte-for-byte what stage 06 captured.
1543
+ cast = ttnn.typecast(logits, self.precision.sampling_dtype)
1544
+ ttnn.deallocate(logits, True)
1545
+ logits = cast
1546
+ self._observed_sampling_dtype = logits.dtype
1547
+ return logits
1548
+
1549
+ def decode_forward_from_ttnn_inputs(
1550
+ self,
1551
+ tokens: ttnn.Tensor,
1552
+ current_pos: ttnn.Tensor,
1553
+ *,
1554
+ rotary_position: ttnn.Tensor,
1555
+ kv_cache: Sequence[KVCache] | None = None,
1556
+ advance_position: bool = True,
1557
+ ) -> ttnn.Tensor:
1558
+ """Token in -> sampler-ready local logits out, entirely on device.
1559
+
1560
+ With ``advance_position`` the two position tensors are incremented
1561
+ **inside** this graph, so a captured trace steps its own positions on
1562
+ replay and the host never refreshes them per token.
1563
+ """
1564
+ hidden = self.decode_hidden(
1565
+ tokens,
1566
+ current_pos=current_pos,
1567
+ rotary_position=rotary_position,
1568
+ kv_cache=kv_cache,
1569
+ )
1570
+ logits = self.decode_terminal(hidden)
1571
+ if advance_position:
1572
+ ttnn.plus_one(current_pos, skip_negative_entries=True)
1573
+ ttnn.plus_one(rotary_position)
1574
+ return logits
1575
+
1576
+ # -- sampling -------------------------------------------------------------
1577
+
1578
+ def sample_split(self, logits, *, k, p, temp, seeds=None, tt_out_tok=None):
1579
+ """Canonical split sampling: local top-32 -> all-gather -> ``ttnn.sampling``.
1580
+
1581
+ ``k=1, p=0, temp=1`` is **semantically greedy**: the global argmax is by
1582
+ construction inside some die's local top-32, and the all-gather makes
1583
+ all four dies' candidates visible before the top-1 is taken.
1584
+ """
1585
+ return self.sampler.decode_forward(
1586
+ logits, k=k, p=p, temp=temp, seeds=seeds, tt_out_tok=tt_out_tok, enable_log_probs=False
1587
+ )[0]
1588
+
1589
+ def sample_greedy_argmax(self, logits, *, tt_out_tok=None):
1590
+ """``Sampling1D``'s force-argmax path, on this model's distributed override.
1591
+
1592
+ Still the common module, still on device, still traced, still writes the
1593
+ sampled token straight into ``tt_out_tok`` -- it is a different strategy
1594
+ inside the same implementation, not a custom sampler. The strategy body
1595
+ is ``_WatcherCleanSampling1D._sample_argmax``: reduce on each die, then
1596
+ all-gather the four survivors instead of all-gathering the vocabulary.
1597
+
1598
+ **This is what greedy uses**, because at this vocabulary it is 6.6x
1599
+ faster than the top-k/top-p split path (0.928 ms against 6.155 ms in the
1600
+ 48-layer model, both rows of
1601
+ ``doc/optimized_full_model/probes/perf_full_model.csv``, which is the
1602
+ **shipped** measurement) and produces the same token -- both rows sample
1603
+ token 16 on that run. Whole-model ``token_out`` on that same run is
1604
+ **19.693 ms, 50.78 t/s/u**.
1605
+
1606
+ Stage 05 shipped the same choice at 1.125 ms against 6.155 ms, and two
1607
+ changes inside this override moved the greedy row since, each with its
1608
+ own token-out delta at a like-for-like context:
1609
+
1610
+ * the **distributed reduction** above -- 22.079 ms to 21.461 ms
1611
+ (45.29 -> 46.60 t/s/u), both at ``context`` 4096,
1612
+ ``../full_model/probes/perf_full_model.json`` against
1613
+ ``doc/optimized_full_model/probes/perf_full_model_part1_preadoption.json``;
1614
+ * the **live-row slice** (reduce ``max_batch_size`` rows, not 32) --
1615
+ 20.146 ms to 19.693 ms at ``context`` 8192,
1616
+ ``doc/optimized_full_model/probes/perf_full_model_p128_after.json``
1617
+ against
1618
+ ``doc/optimized_full_model/probes/perf_full_model_p128_argmaxrows.json``.
1619
+
1620
+ The remaining step between them is the paged SDPA program config in
1621
+ ``tt/multichip_decoder.py`` and is not this sampler's. The moment any
1622
+ slot asks for ``top_k > 1`` or ``top_p > 0`` the generator switches back
1623
+ to ``sample_split``.
1624
+ """
1625
+ return self.sampler.decode_forward(logits, tt_out_tok=tt_out_tok, enable_log_probs=False)[0]
1626
+
1627
+ # -- audit ----------------------------------------------------------------
1628
+
1629
+ def runtime_fallback_audit(self, batch: int | None = None) -> dict:
1630
+ """The layer audit, plus the boundaries this wrapper owns."""
1631
+ batch = self.max_batch_size if batch is None else int(batch)
1632
+ audit = fallback_audit(self.layers[0], self.config, batch, self.precision)
1633
+ audit.update(
1634
+ {
1635
+ "num_layers": self.num_layers,
1636
+ "embedding": "replicated_bf16_no_collective",
1637
+ "residual_contract": "replicated [1,1,B,2048] bf16 TILE DRAM, no inter-layer collective",
1638
+ "final_norm": "replicated, width-sharded decode kernel",
1639
+ "lm_head_parallelism": "column_parallel_over_vocab",
1640
+ "lm_head_local_vocab": self.local_vocab_size,
1641
+ "lm_head_weight_dtype": str(self.lm_head.dtype),
1642
+ "embedding_weight_dtype": str(self.embed_tokens.dtype),
1643
+ "precision": self.precision.to_dict(),
1644
+ "vocab_padding": 0,
1645
+ "decode_rope": "rotary_embedding_hf(is_decode_mode=True), device position gather",
1646
+ "decode_rope_position_source": "device tensor advanced by ttnn.plus_one inside the trace",
1647
+ "sampling_greedy": (
1648
+ "Sampling1D force-argmax, distributed: per-die untilize/argmax/gather -> "
1649
+ "all-gather 4 candidates -> masked-min, traced, writes tt_out_tok"
1650
+ ),
1651
+ "sampling_topk_topp": "Sampling1D split (local topk -> all-gather 32 candidates -> ttnn.sampling)",
1652
+ "sampling_pad_to_power_of_2": False,
1653
+ "host_logit_readback_on_token_out_path": False,
1654
+ "host_argmax_on_token_out_path": False,
1655
+ # Read off the allocated cache when one exists, so a swept
1656
+ # ``kv_cache_dtype`` is *observed* rather than asserted. This
1657
+ # was a hard-coded "bfloat16" until stage 07's sweep, which
1658
+ # would have silently mislabelled every non-default KV row.
1659
+ # Falls back to the configured value before allocation.
1660
+ # Emitted as the PLAIN name ("bfloat16"), not ``str(dtype)``
1661
+ # ("DataType.BFLOAT16"), because that is the existing contract:
1662
+ # doc/optimized_full_model's committed runtime_fallback_audit.json
1663
+ # and check_published_figures.py both pin the plain spelling, and
1664
+ # they are stage evidence that must keep passing. The sibling
1665
+ # ``device_*`` fields use str(dtype) and are left alone.
1666
+ "kv_cache_dtype": dtype_to_name(
1667
+ self.kv_cache[0].k.dtype if self.kv_cache else self.precision.kv_cache_dtype
1668
+ ),
1669
+ "kv_cache_dtype_source": "device_readback" if self.kv_cache else "config_not_yet_allocated",
1670
+ # -- the four fields stage 07's selection proof could not check --
1671
+ #
1672
+ # Before the stage-07 review these were the only swept fields
1673
+ # with no audit entry at all, so ``R03_lmhead_lofi``,
1674
+ # ``R21_norm_hifi2`` and ``R22_logits_sampling_bfp8`` produced
1675
+ # ``device_audit`` blocks byte-identical to the baseline's and
1676
+ # "this lever does nothing" was indistinguishable from "this
1677
+ # lever is not wired up". For ``norm_fidelity`` it was the
1678
+ # second: ``decode_residual_norm`` built its compute config from
1679
+ # the module default and never saw ``self.precision``.
1680
+ #
1681
+ # The two fidelities are read off the ``compute_kernel_config``
1682
+ # objects the ops are actually handed (built here, passed at the
1683
+ # call site), so they verify the config -> compute-config
1684
+ # threading. The two dtypes are read off the **produced
1685
+ # tensors** and are ``None`` until a forward has run.
1686
+ "lm_head_math_fidelity": str(self.lm_head_compute_config.math_fidelity),
1687
+ "norm_math_fidelity": str(self.norm_compute_config.math_fidelity),
1688
+ "logits_dtype_observed": (
1689
+ None if self._observed_logits_dtype is None else dtype_to_name(self._observed_logits_dtype)
1690
+ ),
1691
+ "sampling_dtype_observed": (
1692
+ None if self._observed_sampling_dtype is None else dtype_to_name(self._observed_sampling_dtype)
1693
+ ),
1694
+ "terminal_dtype_source": ("device_readback" if self._observed_logits_dtype else "no_forward_yet"),
1695
+ "kv_cache_paged": True,
1696
+ "page_block_size": self.page_block_size,
1697
+ "collective_topology": str(TOPOLOGY),
1698
+ "prefill_num_links": self.ctx.num_links,
1699
+ "decode_num_links": self.ctx.decode_num_links,
1700
+ }
1701
+ )
1702
+ return audit
1703
+
1704
+ def teardown(self) -> None:
1705
+ if self.kv_cache is not None:
1706
+ for cache in self.kv_cache:
1707
+ ttnn.deallocate(cache.k, True)
1708
+ ttnn.deallocate(cache.v, True)
1709
+ self.kv_cache = None
1710
+
1711
+
1712
+ __all__ = [
1713
+ "DEFAULT_MAX_BATCH_SIZE",
1714
+ "DEFAULT_PAGE_BLOCK_SIZE",
1715
+ "DEFAULT_ROPE_CACHE_LEN",
1716
+ "DEFAULT_TRACE_REGION_SIZE",
1717
+ "HF_MODEL_ID",
1718
+ "HF_REVISION",
1719
+ "MAX_CONTEXT",
1720
+ "NUM_LAYERS",
1721
+ "Qwen3CoderModel",
1722
+ "ShardedCheckpoint",
1723
+ ]
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/multichip_decoder.py ADDED
@@ -0,0 +1,1982 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Multichip TTNN decoder layer for Qwen3-Coder-30B-A3B-Instruct, 4 Blackhole dies.
5
+
6
+ Stages 03 and 04. Stage 04 optimized this file **in place**; the parallelisation
7
+ below is stage 03's and is unchanged, and what stage 04 changed is where the
8
+ activations live inside a layer:
9
+
10
+ * **Both residual RMSNorms are width-sharded over 8 cores** rather than running
11
+ on one (``decode_residual_norm``). 19.82 -> 4.92 us each, and *more* accurate
12
+ than the call they replace. The shard spec is deliberately
13
+ ``_width_sharded_l1(2048)``, so the first norm's output feeds the qkv
14
+ projection with no conversion at all.
15
+ * **The router projection reads that L1 shard** instead of DRAM-interleaved.
16
+ 24.62 -> **5.85 us** at the shipped 8-core norm shard (the same sweep's 4-core
17
+ leg reads 4.30, but 4 cores is not what ships -- the norm shards over
18
+ ``_NORM_SHARD_CORES = 8``), output **bit-identical**, which is what keeps the
19
+ four dies agreeing on the top-8. In the layer it is row 182, 6.241 us.
20
+ * **The two collectives use caller-owned persistent buffers**
21
+ (``_decode_ccl_buffers``), so nothing in the forward path allocates inside the
22
+ trace.
23
+ * **Decode collectives use one ethernet link, not two** (``NUM_LINKS_DECODE``).
24
+ Stage 03 measured this at 0.6% and kept 2 for a single code path; against the
25
+ stage-04 layer it is **1.22%**, over six passes with the leg order alternating
26
+ so that a position effect cannot be read as a link effect. Prefill keeps both.
27
+
28
+ Decode layer device time 414.661 -> 362.828 us on device 0 (1.143x); traced
29
+ decode at ctx 128, 0.4767 -> 0.4286 ms (1.112x), and 0.4700 -> 0.4282 measured
30
+ before and after in one process by ``probes/layer_levers.py``. The inter-layer contract is untouched: a layer
31
+ takes and returns a replicated ``[1, 1, B, 2048]`` bf16 TILE DRAM tensor with no
32
+ collective, gather or reshard between layers. Everything is in
33
+ ``doc/optimized_multichip_decoder/``.
34
+
35
+ The single-chip baseline is ``optimized_decoder.py`` -- every program
36
+ config, dtype and fidelity constant it measured is imported rather than
37
+ re-derived, and the multichip path is the *same graph* with three changes:
38
+
39
+ 1. **Attention is tensor-parallel by 4.** Each die owns 8 Q heads, 1 K head and
40
+ 1 V head, so ``wqkv`` is ``[2048, 1280]`` per die and ``wo`` is
41
+ ``[1024, 2048]``. Both still satisfy ``_dram_sharded_ok`` (1280 = 5x256,
42
+ 1024 = 4x256), so stage 02's DRAM-sharded decode projections survive intact.
43
+ 2. **Experts are expert-parallel by 4.** Each die owns 32 whole experts. M, N
44
+ and K of both ``sparse_matmul`` calls are unchanged, which is the entire
45
+ reason EP was chosen over splitting ``moe_intermediate`` -- see
46
+ ``doc/multichip_decoder/mesh_plan.md`` section 2.
47
+ 3. **Two all-reduces per layer**, one after ``wo`` and one after the expert
48
+ reduce, so the residual stream stays a replicated ``[1, 1, B, 2048]``. That
49
+ makes the layer's input contract identical to its output contract and lets 48
50
+ of them stack with no boundary conversion.
51
+
52
+ Router, both RMSNorms and the residual are replicated. Routing needs a global
53
+ view of 128 logits for top-8, and a 128-wide ``topk`` occupies one core, so
54
+ there is nothing to fracture; the price is that 25.18% of the single-die decode
55
+ layer is replicated work -- 129.09 us of 512.65, the two residual RMSNorms
56
+ (40.23) plus the router block (88.86) -- which caps decode at 3.97x even at
57
+ infinite dies.
58
+
59
+ Mesh and fabric
60
+ ---------------
61
+ 1x4, ``FabricConfig.FABRIC_1D_RING`` before mesh open, ``Topology.Ring`` and
62
+ ``num_links=2`` on every collective. All three are deliberate:
63
+ ``tt_ccl.default_topology()`` returns ``Topology.Linear`` for a 4-device mesh
64
+ (it only special-cases 8-device T3K/Galaxy), and the cluster descriptor for this
65
+ host -- ``ClusterType.P300_X2``, two p300 boards -- shows a genuine closed
66
+ 4-ring with two ethernet links on every hop. Measured cost of taking the default
67
+ instead: 1.21x at decode size, 1.79x at 2 MB
68
+ (``doc/multichip_decoder/mesh_plan.md`` section 5).
69
+
70
+ Both all-reduces, in both modes, are reduce-scatter followed by all-gather. The
71
+ design phase expected decode to want AG-of-partials instead, on a standalone
72
+ sweep that measured 19.96 us against RS+AG's 23.69 at ``[1,1,32,2048]``; the
73
+ shipped decode tensor has **one** logical row rather than 32, which makes
74
+ ``ttnn.sum`` pull a ``FillPad``, and measured on the real layer the order
75
+ reverses. See ``all_reduce`` for the profile rows and the A/B.
76
+
77
+ The ``nnz`` contract, which is a device hang if you get it wrong
78
+ --------------------------------------------------------------
79
+ ``ttnn.sparse_matmul`` bakes ``nnz`` into the kernel as a compile-time arg and
80
+ requires ``count_nonzero(sparsity) == nnz`` exactly;
81
+ ``sparse_matmul_device_operation.cpp:205-211`` says a mismatch *deadlocks the
82
+ device* (tt-metal #45943), silently unless the watcher is on.
83
+
84
+ A TTNN mesh op is SPMD -- one program, one ``nnz``, four dies. Under EP the
85
+ number of locally-live experts is the number of the global top-8 that landed in
86
+ this die's 32-expert window: data-dependent, different on every die, anywhere in
87
+ 0..8. There is no single correct value, so **decode must pass ``nnz=None``**,
88
+ which switches the sender to reading the sparsity page at runtime. That costs
89
+ 0.79 us per slot per matmul, measured at the shipped shapes: the pair costs
90
+ 158.01 us dynamic against 107.73 with an exact nnz, **1.47x**
91
+ (``probes/nnz_cost_probe.py``). That is 50 us over 32 slots -- affordable only
92
+ because EP already cut E from 128 to 32, where the same rate would cost 200. The
93
+ two decisions are coupled.
94
+
95
+ **Prefill keeps an exact ``nnz``.** Its sparsity is per 32-token tile and with
96
+ 32 tokens x top-8 = 256 selections over 128 experts essentially every expert is
97
+ live, so the shipped path uses an all-ones mask; under EP that means all 32
98
+ local experts are live on every die, deterministically, and
99
+ ``nnz = 32 * group_size`` is exact and identical across dies.
100
+
101
+ Rejected alternatives, with the measurement, are in
102
+ ``doc/multichip_decoder/mesh_plan.md`` section 6.
103
+ """
104
+
105
+ from __future__ import annotations
106
+
107
+ import math
108
+ from dataclasses import dataclass, field, replace
109
+
110
+ import torch
111
+
112
+ import ttnn
113
+ from models.common.modules.tt_ccl import TT_CCL
114
+
115
+ from .functional_decoder import (
116
+ AttentionConfig,
117
+ AttentionWeights,
118
+ DecoderLayerConfig,
119
+ KVCache,
120
+ MoEConfig,
121
+ apply_rope_llama,
122
+ attention_prefill,
123
+ rope_transformation_matrix,
124
+ )
125
+
126
+ # The four precision constants this module used to import from here are gone:
127
+ # every one of them is now read off the ``PrecisionConfig`` threaded through the
128
+ # functions below, so importing the import-time default would have been the bug
129
+ # this stage exists to remove. ``tt/precision.py`` holds the values.
130
+ from .optimized_decoder import (
131
+ _DRAM_BANKS,
132
+ OptimizedWeights,
133
+ _attention_compute_kernel_config,
134
+ _bank_row,
135
+ _dram_sharded_ok,
136
+ _expert_compute_kernel_config,
137
+ _ones_column,
138
+ _tuned_sparse_matmul_config,
139
+ _width_sharded_l1,
140
+ attention_decode_optimized,
141
+ moe_prefill_optimized,
142
+ )
143
+ from .precision import DEFAULT_PRECISION, PrecisionConfig # noqa: F401 (re-exported)
144
+ from .weight_mapping import hf_to_meta_channels, permute_head_vector_to_meta, permute_wqkv_to_meta
145
+
146
+ # The target mesh. This module deliberately supports exactly one shape: the
147
+ # goal is the best use of *this* machine, and every constant below -- the head
148
+ # split, the 32-expert window, the ring topology, the two links -- is chosen
149
+ # against it. A 1x2 or 1x8 mesh would want different answers, not a scaled
150
+ # version of these ones.
151
+ MESH_SHAPE = (1, 4)
152
+ NUM_DEVICES = 4
153
+
154
+ # Ring, not Linear. See the module docstring; this must be passed explicitly
155
+ # because tt_ccl.default_topology() returns Linear for a 4-device mesh.
156
+ TOPOLOGY = ttnn.Topology.Ring
157
+
158
+ # Both ethernet links on every hop are used **in prefill**, where the payload is
159
+ # large enough to be bandwidth-bound: 1.84x at 2 MB.
160
+ NUM_LINKS = 2
161
+
162
+ # Decode uses **one**. Stage 03 measured 1 link at 0.4738 ms against 2 links'
163
+ # 0.4766, called 0.6% noise-level, and kept 2 for a single code path. Against
164
+ # the stage-04 layer -- where the collectives are a larger share because
165
+ # everything around them got smaller, and where they no longer allocate -- the
166
+ # gap is 1.22% and output is bit-identical (``probes/links_probe.py``, six
167
+ # passes with the **leg order alternating**, so that a position effect cannot be
168
+ # read as a link effect):
169
+ #
170
+ # posA 2 links 0.4342 0.4341 0.4340 1 link 0.4290 0.4288 0.4286
171
+ # posB 2 links 0.4341 0.4337 0.4339 1 link 0.4291 0.4283 0.4287
172
+ #
173
+ # mean 2 links 0.43400 1 link 0.42875 1.22%
174
+ #
175
+ # Each configuration reads the same at both positions, which is what rules out
176
+ # the alternative explanation -- that the leg running first in a pass is simply
177
+ # slower. That control was added because review found ``_links`` had stopped
178
+ # honouring an explicit ``num_links=2``, leaving the probe unable to tell its
179
+ # own legs apart; the figure survived the repair, its reproducibility did not
180
+ # and now does. 5.25 us on the layer against a leg-against-itself spread of
181
+ # 0.5-0.8 us. A decode
182
+ # collective moves 128 KB per die and is latency-bound, so the second link buys
183
+ # no bandwidth and costs the split and merge. ``all_reduce`` branches on the
184
+ # same ``S <= 32`` test ``_decode_ccl_buffers`` uses, so prefill keeps both.
185
+ NUM_LINKS_DECODE = 1
186
+
187
+
188
+ def _links(x: ttnn.Tensor, ctx: "MeshContext") -> int:
189
+ """``ctx.num_links`` for prefill, ``ctx.decode_num_links`` for decode.
190
+
191
+ The two counts are **separate fields** rather than one field plus a
192
+ "differs from the default means override" test. That test was the first
193
+ spelling here and it is not expressible: a caller asking explicitly for
194
+ ``num_links=2`` at decode passes ``ctx.num_links == NUM_LINKS``, which the
195
+ test read as "no override" and silently gave 1 link. ``links_probe.py``
196
+ builds its two-link leg exactly that way, so the probe that established
197
+ ``NUM_LINKS_DECODE`` could not have been re-run against it -- review caught
198
+ this. Two fields make each mode's count independently settable and the
199
+ probe's legs actually different.
200
+ """
201
+ return ctx.decode_num_links if int(x.shape[-2]) <= 32 else ctx.num_links
202
+
203
+
204
+ # Decode's two expert intermediates under EP:
205
+ #
206
+ # batch * 32 experts * 32 padded rows * (2*768 + 2048) cols * 2 B
207
+ # = batch * 7,340,032 B = batch * 7.34 MB
208
+ #
209
+ # a quarter of the single-die figure, because EP fractures the expert dimension
210
+ # these tensors are indexed by. Stage 02's 40 MB threshold sat between its batch
211
+ # 1 (29.4 MB) and batch 2 (58.8 MB) and its own comment calls it "asserted, not
212
+ # measured"; inherited here it would have admitted batch 5 by accident. Swept
213
+ # instead (``probes/l1_budget_probe.py``, eager, ms per decode step):
214
+ #
215
+ # batch intermediates L1 DRAM
216
+ # 1 7.34 MB 1.7336 1.6924
217
+ # 2 14.68 1.7816 1.7930
218
+ # 4 29.36 1.8887 1.9753
219
+ # 8 58.72 2.4969 2.9082
220
+ # 16 117.44 3.3040 4.0386
221
+ # 32 234.88 allocator refuses (bank_manager.cpp:462)
222
+ #
223
+ # L1 wins from batch 2 to 16 and stops being allocatable at 32, so the threshold
224
+ # goes between 117.44 and 234.88 MB. The batch-1 row of that sweep reads the
225
+ # other way, but it is an *eager* measurement where host dispatch is most of the
226
+ # 1.7 ms; the warmed traced A/B that decides the shipped configuration says L1,
227
+ # clearly -- **0.4766 ms against DRAM's 0.5128, 7.6%**
228
+ # (``probes/decode_levers.py``). Batch 1 is also the latency target.
229
+ _DECODE_EXPERT_L1_BUDGET_BYTES = 128 * 1024 * 1024
230
+
231
+
232
+ def _decode_expert_memory_config(batch: int, local_moe: MoEConfig) -> ttnn.MemoryConfig:
233
+ padded_rows = batch * local_moe.num_experts * 32
234
+ nbytes = padded_rows * (2 * local_moe.moe_intermediate_size + local_moe.hidden_size) * 2
235
+ return ttnn.L1_MEMORY_CONFIG if nbytes <= _DECODE_EXPERT_L1_BUDGET_BYTES else ttnn.DRAM_MEMORY_CONFIG
236
+
237
+
238
+ # --- Meta-ordered rotary for decode (stage 04) --------------------------------
239
+
240
+
241
+ def _meta_rope(ctx: MeshContext, cos_cache: ttnn.Tensor, sin_cache: ttnn.Tensor, head_dim: int):
242
+ """Return a ``rope(t, cos, sin, token_index)`` callable using the llama op.
243
+
244
+ **Measured, and not adopted.** Kept runnable rather than deleted, on the
245
+ same principle as ``router_forward_threshold``: the finding is the useful
246
+ part, and a future stage that changes prefill should not have to rediscover
247
+ it. Nothing on the shipped path calls this --
248
+ ``decoder_layer_decode_multichip`` ships the HF op, and
249
+ ``upload_multichip_weights`` builds the Meta weights only under
250
+ ``meta_rope=True``, so the shipped upload pays no DRAM for it.
251
+
252
+ ``rotary_embedding_llama`` costs **1.26 us** against the shipped HF op's
253
+ 3.84 at the per-die decode shape, with ``max|diff|`` exactly 0.0 and PCC
254
+ 1.0000000 (``probes/rope_probe.py``). Both run on one core: the llama
255
+ decode factory shards over *batch*, not heads, so at batch 1 none of the
256
+ 3.05x is parallelism -- it is the activation living in L1 and a kernel that
257
+ multiplies by a resident 32x32 matrix instead of gathering a cos/sin row out
258
+ of a DRAM cache. Same lever as the router projection, different op.
259
+
260
+ Two things are hoisted out of the forward path, both on the first (eager)
261
+ call and cached on ``ctx``:
262
+
263
+ * the **Meta cos/sin** for this ``token_index``, read off the HF device
264
+ cache once, permuted on the host and uploaded already sharded. This hoist
265
+ is what makes *this* wiring unreplayable, and it is a property of the
266
+ wiring rather than of the op: ``rotary_embedding_llama`` takes cos/sin as
267
+ tensors and no position argument at all
268
+ (``rotary_embedding_llama_nanobind.cpp:38-44``), so it can be driven from
269
+ a position tensor inside a trace. The shipped
270
+ ``rotary_embedding(..., token_index)`` genuinely cannot, which is why
271
+ stage 05 moved decode to ``rotary_embedding_hf``. An earlier revision of
272
+ this docstring claimed neither spelling could; that was wrong.
273
+ * the **transformation matrix**, which is position-independent.
274
+
275
+ The Meta *channel order* is not established here at all: it is a property of
276
+ ``ctx``-independent weights, applied once by
277
+ ``weight_mapping.permute_wqkv_to_meta`` at upload.
278
+
279
+ **Why it is not adopted.** RoPE runs *before* K is written to the cache, so
280
+ the cache inherits the rotary's channel convention. Prefill is untouched by
281
+ this lever and writes HF-ordered keys; a Meta-ordered decode Q then scores
282
+ against them, and the dot products are meaningless.
283
+ ``probes/rope_layer_probe.py``:
284
+
285
+ fresh KV cache PCC 0.9999697 the rotary itself is right
286
+ prefill-primed cache PCC 0.1932974 the cache convention is not
287
+
288
+ The op-level probe looked clean precisely because its cache was fresh. So
289
+ the lever is not decode-local: adopting it means adopting the llama rotary
290
+ in **prefill** as well, permuting the interleaved ``wqkv`` prefill copy, and
291
+ changing the KV cache's channel convention -- which
292
+ ``test_per_die_kv_heads_stitched`` compares against a single-chip cache and
293
+ which ``config/context_contract.json`` describes. That is a whole-layer change,
294
+ not the in-place decode optimization this stage is.
295
+
296
+ A second cost, smaller and independent: the qkv weight dtype
297
+ (``PrecisionConfig.attention_qkv_dtype``) is ``bfloat8_b`` by default, and
298
+ bfloat8_b's 16-element blocks share an exponent, so permuting
299
+ channels **regroups the blocks** and requantizes. The two paths therefore
300
+ are not bit-identical in the layer even where the ops are -- attention out
301
+ ``max|diff|`` 1.221e-04 on a fresh cache, and the K cache differs by
302
+ 3.125e-01 after permuting back. "Bit-identical" is a property of the op at
303
+ fixed input, not of the layer at permuted weights.
304
+ """
305
+ st = ctx.rope_meta
306
+
307
+ def trans_mat(batch: int):
308
+ # One 32x32 copy **per batch core**, because the decode factory shards
309
+ # over batch: at batch 1 that is a single tile on a single core, at
310
+ # batch 32 it is 32 of them. Keyed by batch for that reason.
311
+ t = st.get(("tm", batch))
312
+ if t is None:
313
+ t = st[("tm", batch)] = ttnn.from_torch(
314
+ rope_transformation_matrix().repeat(1, 1, batch, 1),
315
+ device=ctx.mesh,
316
+ layout=ttnn.TILE_LAYOUT,
317
+ dtype=ttnn.bfloat16,
318
+ memory_config=_head_shard(32, 32, batch),
319
+ mesh_mapper=ttnn.ReplicateTensorToMesh(ctx.mesh),
320
+ )
321
+ return t
322
+
323
+ def rope(t: ttnn.Tensor, _cos, _sin, token_index: int) -> ttnn.Tensor:
324
+ batch = int(t.shape[1])
325
+ key = (int(token_index), batch, head_dim)
326
+ pair = st.get(key)
327
+ if pair is None:
328
+ if "host" not in st:
329
+ # Read the HF cos/sin cache back once and permute on the host.
330
+ # Replicated, so die 0's copy is the whole tensor.
331
+ comp = ttnn.ConcatMeshToTensor(ctx.mesh, dim=0)
332
+ st["host"] = (
333
+ ttnn.to_torch(cos_cache, mesh_composer=comp)[:1].float(),
334
+ ttnn.to_torch(sin_cache, mesh_composer=comp)[:1].float(),
335
+ )
336
+ perm = hf_to_meta_channels(head_dim)
337
+ mem = _head_shard(32, head_dim, batch)
338
+ up = []
339
+ for c in st["host"]:
340
+ row = c[:, :, token_index : token_index + 1, :]
341
+ row = row.expand(1, 1, 32, head_dim).contiguous()[..., perm]
342
+ up.append(
343
+ ttnn.from_torch(
344
+ row.expand(1, batch, 32, head_dim).contiguous(),
345
+ device=ctx.mesh,
346
+ layout=ttnn.TILE_LAYOUT,
347
+ dtype=ttnn.bfloat16,
348
+ memory_config=mem,
349
+ mesh_mapper=ttnn.ReplicateTensorToMesh(ctx.mesh),
350
+ )
351
+ )
352
+ pair = st[key] = tuple(up)
353
+ sharded = ttnn.to_memory_config(t, _head_shard(32, head_dim, int(t.shape[1])))
354
+ out = apply_rope_llama(sharded, pair[0], pair[1], trans_mat(batch))
355
+ ttnn.deallocate(sharded)
356
+ return out
357
+
358
+ return rope
359
+
360
+
361
+ def _head_shard(rows: int, cols: int, batch: int) -> ttnn.MemoryConfig:
362
+ """The height-sharded L1 config ``nlp_create_qkv_heads_decode`` emits and
363
+ ``rotary_embedding_llama``'s decode factory requires: one core per user,
364
+ each holding that user's whole ``[32 padded heads, head_dim]`` block."""
365
+ gx = min(batch, 8)
366
+ while batch % gx:
367
+ gx -= 1
368
+ gy = batch // gx
369
+ return ttnn.MemoryConfig(
370
+ ttnn.TensorMemoryLayout.HEIGHT_SHARDED,
371
+ ttnn.BufferType.L1,
372
+ ttnn.ShardSpec(
373
+ ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))}),
374
+ [rows, cols],
375
+ ttnn.ShardOrientation.ROW_MAJOR,
376
+ ),
377
+ )
378
+
379
+
380
+ # --- mesh plumbing ------------------------------------------------------------
381
+
382
+
383
+ @dataclass
384
+ class MeshContext:
385
+ """Mesh, CCL semaphores and the collective parameters, owned explicitly.
386
+
387
+ ``TT_CCL`` is instantiated directly rather than through
388
+ ``tt_ccl.get_tt_ccl()``: that helper caches by ``mesh_device.id()`` in a
389
+ module-global dict, and the pytest ``mesh_device`` fixture is
390
+ function-scoped, so a later mesh can be handed a recycled id and inherit
391
+ semaphores belonging to a closed device. Ownership here is per-caller and
392
+ dies with the caller.
393
+
394
+ Global semaphores are hardware resources allocated at construction time,
395
+ which is also what makes this trace-safe: nothing in the forward path
396
+ allocates one.
397
+ """
398
+
399
+ mesh: ttnn.MeshDevice
400
+ ccl: TT_CCL
401
+ num_devices: int = NUM_DEVICES
402
+ num_links: int = NUM_LINKS
403
+ topology: ttnn.Topology = TOPOLOGY
404
+ # Links for a *decode* collective, separately settable. See ``_links``.
405
+ decode_num_links: int = NUM_LINKS_DECODE
406
+ # Stage 04. Meta-ordered rotary state for the decode path: the 32x32
407
+ # transformation matrix, the host-side Meta cos/sin caches, and the
408
+ # per-position sharded cos/sin pair. Keyed by ``token_index``, which is a
409
+ # Python int here exactly as it is for the shipped HF op -- the rotary
410
+ # position is baked into a traced program either way, so the gather is
411
+ # hoisted out of the forward path rather than run per token. Allocated on a
412
+ # miss, so the first call at each position must be eager, which is the same
413
+ # discipline ``ccl_buffers`` below already imposes. See ``_meta_rope``.
414
+ rope_meta: dict = field(default_factory=dict)
415
+ # Stage 04. Persistent collective buffers, keyed by (logical shape, padded
416
+ # shape, dtype), so that neither the reduce-scatter nor the all-gather
417
+ # allocates inside the trace. See ``_decode_ccl_buffers``. Owned by the
418
+ # context and therefore by the caller, exactly like the semaphores above.
419
+ ccl_buffers: dict = field(default_factory=dict)
420
+
421
+
422
+ def mesh_context(mesh_device) -> MeshContext:
423
+ """Build the CCL context for the 4-die mesh, asserting the shape."""
424
+ n = mesh_device.get_num_devices()
425
+ assert n == NUM_DEVICES, (
426
+ f"multichip_decoder targets exactly {NUM_DEVICES} dies (the full P300_X2 mesh); got {n}. "
427
+ "Smaller meshes are out of scope by design -- the head split, expert window and ring "
428
+ "topology are all chosen against the 4-die shape."
429
+ )
430
+ return MeshContext(mesh=mesh_device, ccl=TT_CCL(mesh_device))
431
+
432
+
433
+ def all_reduce(x: ttnn.Tensor, ctx: MeshContext, precision: PrecisionConfig = DEFAULT_PRECISION) -> ttnn.Tensor:
434
+ """All-reduce a ``[1, 1, ., H]`` partial as reduce-scatter then all-gather.
435
+
436
+ **One spelling for both modes, and that is a change from the plan.**
437
+ ``mesh_plan.md`` §5 chose AG-of-partials-plus-local-sum for decode on a
438
+ standalone sweep that measured 19.96 us against RS+AG's 23.69 at
439
+ ``[1,1,32,2048]``. That sweep used the wrong shape. The shipped decode tensor
440
+ is ``[1,1,1,2048]`` -- **one** logical row padded to a tile -- and
441
+ ``ttnn.sum`` over a tensor whose last two dims are not both tile-aligned
442
+ drags a ``FillPad`` behind it (``fill_pad.cpp:17-24``), which is precisely
443
+ the hazard stage 02 removed from the router. Read off
444
+ ``ops_perf_multichip_decode_agsum.csv`` -- a profile of this layer with the
445
+ plan's spelling, kept precisely because the shipped path no longer produces
446
+ those rows (``probes/profile_layer.py decode-agsum``) -- the local sum is not
447
+ one op but four:
448
+
449
+ AllGatherAsync 22.31 us + FillPad 5.89 + FastReduceNC 2.44 + Slice 1.32
450
+ = 31.96 us (attention all-reduce)
451
+ AllGatherAsync 18.65 + FillPad 5.64 + FastReduceNC 2.43 + Slice 1.31
452
+ = 28.04 us (expert all-reduce)
453
+
454
+ against the 19.96 the probe promised for each. Measured on the whole traced
455
+ layer at ctx 128, median of 100 (``probes/allreduce_ab.py``):
456
+
457
+ AG(dim 0) + ttnn.sum 0.4801 ms
458
+ reduce-scatter + all-gather **0.4760 ms** <- adopted
459
+
460
+ 0.9%, which is small -- but it is also three fewer ops, one code path
461
+ instead of two, and it is the direction the standalone probe got backwards.
462
+ A third leg that kept the single collective and reshaped the logical shape
463
+ up to the padded 32 rows to dodge the ``FillPad`` **did not run**:
464
+ ``reshape_common.cpp:50`` rejects it, ``new_volume == old_volume``.
465
+
466
+ Prefill was RS+AG already and for the reason that still holds: 76.85 us
467
+ against AG-of-partials' 121.72 at ``[1,1,512,2048]``, because past ~128 KB
468
+ per device the collective is bandwidth-bound and RS+AG moves a quarter of
469
+ the bytes on each of its two hops.
470
+
471
+ The scatter axis is dim 3 (hidden, 2048), which is independent of the
472
+ sequence length -- that is what keeps non-aligned S working through the
473
+ collective without any padding of its own.
474
+
475
+ ``precision.ccl_dtype`` is ``None`` on the shipped path, which means "run
476
+ the collective at whatever dtype the partial arrives in" -- no cast, no
477
+ extra op, the behaviour every stage-02..06 number was measured at. A named
478
+ dtype casts in before the reduce-scatter and back out after the all-gather,
479
+ so a sweep can price a narrower wire without touching the arithmetic that
480
+ feeds it. The cast is deliberately *outside* the buffer cache key's reach
481
+ only in the sense that the cache keys on ``x.dtype`` already -- casting
482
+ first means the cached buffers are allocated at the wire dtype, which is the
483
+ point.
484
+ """
485
+ # The cast allocates a *new* tensor and leaves ``x`` alone: every caller
486
+ # deallocates the partial it passed in, so freeing it here would be a double
487
+ # free the moment ``ccl_dtype`` was set.
488
+ wire_dtype = precision.ccl_dtype
489
+ restore_dtype = None
490
+ cast_in = None
491
+ if wire_dtype is not None and x.dtype != wire_dtype:
492
+ restore_dtype = x.dtype
493
+ cast_in = ttnn.typecast(x, wire_dtype)
494
+ x = cast_in
495
+ bufs = _decode_ccl_buffers(x, ctx)
496
+ num_links = _links(x, ctx)
497
+ scattered = ttnn.experimental.reduce_scatter_minimal_async(
498
+ x,
499
+ persistent_output_buffers=None if bufs is None else bufs[0],
500
+ dim=3,
501
+ multi_device_global_semaphore=ctx.ccl.get_and_cycle_rs_semaphore_handles(),
502
+ barrier_semaphore=ctx.ccl.get_and_cycle_barrier_semaphore_handle(),
503
+ num_links=num_links,
504
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
505
+ intermediate_memory_config=ttnn.DRAM_MEMORY_CONFIG,
506
+ topology=ctx.topology,
507
+ )
508
+ gathered = ttnn.experimental.all_gather_async(
509
+ scattered,
510
+ persistent_output_buffer=None if bufs is None else bufs[1],
511
+ dim=3,
512
+ multi_device_global_semaphore=ctx.ccl.get_and_cycle_ag_semaphore_handles(),
513
+ barrier_semaphore=ctx.ccl.get_and_cycle_barrier_semaphore_handle(),
514
+ num_links=num_links,
515
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
516
+ topology=ctx.topology,
517
+ )
518
+ if bufs is None:
519
+ ttnn.deallocate(scattered)
520
+ out = gathered
521
+ else:
522
+ # ``gathered`` *is* the persistent buffer, which the caller is about to
523
+ # deallocate. Hand back a copy so the buffer survives the next token.
524
+ out = ttnn.clone(gathered, memory_config=ttnn.DRAM_MEMORY_CONFIG)
525
+ if restore_dtype is not None:
526
+ cast_out = ttnn.typecast(out, restore_dtype)
527
+ ttnn.deallocate(out)
528
+ out = cast_out
529
+ if cast_in is not None:
530
+ ttnn.deallocate(cast_in)
531
+ return out
532
+
533
+
534
+ def _decode_ccl_buffers(x: ttnn.Tensor, ctx: MeshContext):
535
+ """``([rs intermediate, rs output, penult], ag output)`` for a decode-shaped
536
+ ``x``, allocated once per (logical shape, padded shape, dtype) and cached on
537
+ the context.
538
+
539
+ ``None`` for anything taller than one 32-row tile: prefill runs at a
540
+ different ``S`` on every call, so caching there would allocate a set per
541
+ sequence length for a lever worth 0.2% at decode and nothing measurable at
542
+ prefill (the prefill collective is bandwidth-bound, not allocation-bound).
543
+
544
+ **Allocated on the first call at each shape, so that call must be eager.**
545
+ ``ttnn.from_torch`` inside ``begin_trace_capture`` raises "Writes are not
546
+ supported during trace capture" and leaves the trace open, which is a hung
547
+ mesh and a ``tt-smi -r``. Every harness here runs the layer once before
548
+ capturing, which is also what the semaphores in ``MeshContext`` already
549
+ require; the constraint is not new, only wider.
550
+
551
+ All 48 layers of the stacked model share the cache, and so does every token.
552
+ That is safe because the trace serialises the collectives and each result is
553
+ cloned out before the next one starts -- but it is exactly the property a
554
+ future change has to preserve, so it is exercised by
555
+ ``test_multichip_decode_20_steps_deterministic`` running 20 tokens through
556
+ the same buffers and by ``test_two_layers_stacked``.
557
+
558
+ The layer's *two* all-reduces do **not** share a set; see the key below.
559
+
560
+ Measured: 0.4343 / 0.4337 ms against the allocating path's 0.4348 / 0.4346,
561
+ over two interleaved passes (``probes/layer_levers3.py``), and 0.4335 /
562
+ 0.4333 against 0.4348 / 0.4346 in ``probes/layer_levers2.py``.
563
+ """
564
+ if int(x.shape[-2]) > 32:
565
+ return None
566
+ # The key must carry the **logical** shape, not just the padded one.
567
+ #
568
+ # The layer's two all-reduces are both [1,1,batch,2048] and *do* share one
569
+ # set, correctly: the attention partial is ``batch`` rows because
570
+ # ``_concat_heads_decode`` slices the padded tile back before ``wo``, and the
571
+ # expert partial is ``batch`` by construction. What collides is the priming
572
+ # prefill: at ``S <= 32`` it takes this branch too, and a 32-token prefill
573
+ # and a decode at ``batch < 32`` have the same *padded* shape, one 32-row
574
+ # tile. A persistent output buffer imposes *its* logical shape on the op's
575
+ # result, so keyed on the padded shape alone the decode layer inherited the
576
+ # prefill's 32 rows and silently returned a 32-row tensor. Not hypothetical
577
+ # -- six decode tests caught it (``work_log.md`` section 5).
578
+ key = (tuple(int(v) for v in x.shape), tuple(int(v) for v in x.padded_shape), str(x.dtype))
579
+ entry = ctx.ccl_buffers.get(key)
580
+ if entry is None:
581
+ interm, penult = ttnn.experimental.reduce_scatter_minimal_async_create_intermediate_buffer(
582
+ x, dim=3, topology=ctx.topology, cluster_axis=None
583
+ )
584
+ shape = list(x.shape)
585
+ shape[3] //= ctx.num_devices
586
+
587
+ def zeros(s):
588
+ return ttnn.from_torch(
589
+ torch.zeros(s),
590
+ device=ctx.mesh,
591
+ layout=ttnn.TILE_LAYOUT,
592
+ dtype=x.dtype,
593
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
594
+ mesh_mapper=ttnn.ReplicateTensorToMesh(ctx.mesh),
595
+ )
596
+
597
+ rs_bufs = [interm, zeros(shape)] + ([penult] if penult is not None else [])
598
+ entry = (rs_bufs, zeros(list(x.shape)))
599
+ ctx.ccl_buffers[key] = entry
600
+ return entry
601
+
602
+
603
+ # Kept as names so callers read as prefill/decode; both are the same collective.
604
+ all_reduce_prefill = all_reduce
605
+ all_reduce_decode = all_reduce
606
+
607
+
608
+ # --- configuration ------------------------------------------------------------
609
+
610
+
611
+ @dataclass(frozen=True)
612
+ class MeshDecoderConfig:
613
+ """The global layer config plus the per-die views the kernels actually see.
614
+
615
+ ``local_attention`` carries 8 Q heads and 1 KV head; ``local_moe`` carries
616
+ 32 experts. Every op below is handed a *local* config, which is what makes
617
+ the multichip layer literally the single-chip code at a quarter of the
618
+ shape rather than a reimplementation of it.
619
+ """
620
+
621
+ global_config: DecoderLayerConfig
622
+ local_attention: AttentionConfig
623
+ local_moe: MoEConfig
624
+ num_devices: int = NUM_DEVICES
625
+
626
+ @classmethod
627
+ def from_hf(cls, hf_config, num_devices: int = NUM_DEVICES) -> "MeshDecoderConfig":
628
+ return cls.from_global(DecoderLayerConfig.from_hf(hf_config), num_devices)
629
+
630
+ @classmethod
631
+ def from_global(cls, config: DecoderLayerConfig, num_devices: int = NUM_DEVICES) -> "MeshDecoderConfig":
632
+ a, m = config.attention, config.moe
633
+ assert a.num_attention_heads % num_devices == 0, f"{a.num_attention_heads} Q heads / {num_devices}"
634
+ assert a.num_key_value_heads % num_devices == 0, (
635
+ f"{a.num_key_value_heads} KV heads / {num_devices} -- this is the hard cap on the TP factor; "
636
+ "TP=8 would need KV-head replication and would also take wqkv's N to 640, which is not a "
637
+ "multiple of 8 banks x 32 = 256, so the DRAM-sharded attention path would silently vanish"
638
+ )
639
+ assert m.num_experts % num_devices == 0, f"{m.num_experts} experts / {num_devices}"
640
+ local_attention = AttentionConfig(
641
+ hidden_size=a.hidden_size,
642
+ num_attention_heads=a.num_attention_heads // num_devices,
643
+ num_key_value_heads=a.num_key_value_heads // num_devices,
644
+ head_dim=a.head_dim,
645
+ rms_norm_eps=a.rms_norm_eps,
646
+ )
647
+ local_moe = MoEConfig(
648
+ hidden_size=m.hidden_size,
649
+ num_experts=m.num_experts // num_devices,
650
+ num_experts_per_tok=m.num_experts_per_tok,
651
+ moe_intermediate_size=m.moe_intermediate_size,
652
+ norm_topk_prob=m.norm_topk_prob,
653
+ )
654
+ return cls(
655
+ global_config=config,
656
+ local_attention=local_attention,
657
+ local_moe=local_moe,
658
+ num_devices=num_devices,
659
+ )
660
+
661
+
662
+ # --- weights ------------------------------------------------------------------
663
+
664
+
665
+ def head_interleaved_wqkv(wqkv: torch.Tensor, config: AttentionConfig, num_devices: int) -> torch.Tensor:
666
+ """Permute the fused QKV columns so a contiguous 4-way split is the TP split.
667
+
668
+ **This is the one weight transform that a naive port gets wrong.** The
669
+ checkpoint's fused weight is ``[Wq(4096) | Wk(512) | Wv(512)]``, so a plain
670
+ ``ShardTensorToMesh(dim=-1)`` hands die 0 nothing but Q heads and die 3
671
+ nothing but K and V. Die *d* must instead own Q heads ``8d..8d+7``, K head
672
+ ``d`` and V head ``d``, laid out as ``[Q_local | K_local | V_local]`` --
673
+ which is what ``nlp_create_qkv_heads_decode(num_heads=8, num_kv_heads=1)``
674
+ reads on the other side.
675
+
676
+ Rebuilding the tensor in that order here means the runtime split stays a
677
+ plain contiguous shard, and the failure mode -- which produces no shape
678
+ error, only a wrong answer -- cannot come back through a different mapper.
679
+
680
+ ``wqkv`` is ``[..., hidden, (n_heads + 2*n_kv) * head_dim]``; the return has
681
+ the same shape with its last dim permuted.
682
+ """
683
+ n_heads, n_kv, hd = config.num_attention_heads, config.num_key_value_heads, config.head_dim
684
+ q_per, kv_per = n_heads // num_devices, n_kv // num_devices
685
+ q_end, k_end = n_heads * hd, n_heads * hd + n_kv * hd
686
+
687
+ cols = []
688
+ for d in range(num_devices):
689
+ cols.append(wqkv[..., d * q_per * hd : (d + 1) * q_per * hd])
690
+ cols.append(wqkv[..., q_end + d * kv_per * hd : q_end + (d + 1) * kv_per * hd])
691
+ cols.append(wqkv[..., k_end + d * kv_per * hd : k_end + (d + 1) * kv_per * hd])
692
+ out = torch.cat(cols, dim=-1)
693
+ assert out.shape == wqkv.shape
694
+ return out
695
+
696
+
697
+ @dataclass
698
+ class MultichipWeights:
699
+ """Everything one multichip decoder layer reads.
700
+
701
+ ``experts`` is an ``OptimizedWeights`` whose tensors are mesh-sharded: the
702
+ two expert weights on the expert dimension, ``wqkv``/``wo`` on the
703
+ head-interleaved column and Q-head row split respectively. The dataclass is
704
+ reused unchanged so ``attention_decode_optimized`` and
705
+ ``moe_prefill_optimized`` can be called directly.
706
+
707
+ ``expert_window`` is the only genuinely *device-varying* constant in the
708
+ layer: a one-hot ``[1, 1, 128, 32]`` matrix, different on every die, that
709
+ slices this die's 32-expert column window out of the replicated dense
710
+ routing vector. See ``router_forward_multichip``.
711
+ """
712
+
713
+ input_layernorm: ttnn.Tensor
714
+ post_attention_layernorm: ttnn.Tensor
715
+ router: ttnn.Tensor
716
+ expert_window: ttnn.Tensor
717
+ experts: OptimizedWeights
718
+ # Stage 04. ``ttnn.rms_norm``'s sharded program factory reads its weight as a
719
+ # ROW_MAJOR ``[1, 1, dim/32, 32]`` tensor rather than the tiled ``[1,1,1,dim]``
720
+ # the interleaved factory takes, so decode carries a second copy of each of
721
+ # the two residual norm vectors. 4 KB each against the layer's 95.5 MB.
722
+ input_layernorm_rm: ttnn.Tensor | None = None
723
+ post_attention_layernorm_rm: ttnn.Tensor | None = None
724
+ # Stage 04. The same ``OptimizedWeights`` with the Q and K channels of
725
+ # ``wqkv_decode`` -- and of ``q_norm``/``k_norm``, which Qwen3 applies
726
+ # between the head split and RoPE -- reordered to the Meta convention
727
+ # ``rotary_embedding_llama`` requires. Decode-only: ``wqkv`` (prefill's
728
+ # interleaved copy), ``wo`` and the expert weights are the *same objects*,
729
+ # not copies, so this costs one extra DRAM-sharded qkv (11.14 MB/4 per die)
730
+ # and two 128-element vectors, and prefill cannot reach it.
731
+ experts_meta: OptimizedWeights | None = None
732
+
733
+
734
+ # SDPA-decode's tree reduction is capped at 6 rounds, i.e. 64 cores per KV head
735
+ # (``sdpa_decode_program_factory.cpp:245``). With no program config the op sets
736
+ # ``max_cores_per_head = num_cores_available``, so at TP=4 -- one KV head per die
737
+ # -- batch 1 asks for all 110 worker cores on that single head and the op raises
738
+ #
739
+ # Tree reduction max 6 rounds (64 cores/head), got 110 cores/head
740
+ #
741
+ # This is a *new* failure created by the head split: at the single-chip 4 KV
742
+ # heads the same arithmetic gives 27 cores/head. It only bites the contiguous
743
+ # cache path -- the paged one runs at the default -- and only at small batch,
744
+ # because ``num_cores_per_head`` divides by the batch. Capping the per-head core
745
+ # budget at the op's own limit fixes it without giving up any parallelism the op
746
+ # would have been allowed to use.
747
+ _SDPA_MAX_CORES_PER_HEAD = 64
748
+
749
+ # The **paged** path -- the one the full model actually runs -- had no program
750
+ # config at all until stage 06, because the cap above was added to clear a
751
+ # ``TT_FATAL`` the paged path never raised. Running at the op default is not
752
+ # free: with no config the op picks its own ``k_chunk_size`` and core split, and
753
+ # the result is a decode cost that is **linear in ``cur_pos``** rather than
754
+ # flat. Measured at the shipped per-die decode shapes -- 8 Q heads, 1 KV head,
755
+ # head_dim 128, page 32, batch 1, **bfloat16** cache -- with PCC taken against a
756
+ # float32 reference built from the same cache the kernel reads, not against the
757
+ # default leg (``probes/sdpa_sweep_confirm.py``, median of 5 blocks of 50):
758
+ #
759
+ # cur_pos default k256/c16 speedup default PCC k256/c16 PCC
760
+ # 127 23.72 us 19.00 us 1.25x 0.999734 0.999707
761
+ # 1023 120.51 22.02 5.47x 0.999714 0.999703
762
+ # 4095 451.85 30.75 14.69x 0.999519 0.999655
763
+ # 8191 893.14 38.15 23.41x 0.999024 0.999590
764
+ # 16383 1777.60 49.83 35.67x 0.993199 0.999577
765
+ # 32767 3545.05 74.13 47.82x 0.989703 0.999692
766
+ #
767
+ # Two things in that table, not one. The speed column is the expected one. The
768
+ # **PCC columns are the surprise**: the default's accuracy *decays with depth* --
769
+ # 0.9932 at 16k and 0.9897 at 32k, through this project's 0.995 layer bar --
770
+ # while the configured path holds 0.9996-0.9997 flat from 127 to 32767. So this
771
+ # is not a speed-for-accuracy trade. At the context this model advertises the
772
+ # config is strictly better on both axes, and the shipped default was the *less*
773
+ # accurate of the two.
774
+ #
775
+ # **The cache dtype is why this took two passes, and it is the lesson.** The
776
+ # stage-06 lever analysis recommended ``k_chunk_size=512`` on the strength of
777
+ # probes that allocated the cache as ``bfloat8_b``; ``create_mesh_kv_cache``
778
+ # allocates ``ttnn.bfloat16`` (see below, ~line 1167). Re-run at the real dtype,
779
+ # 512 loses its edge -- and, far worse, **512 is wrong in-model**:
780
+ # ``test_multichip_decode_batch`` (128-position paged cache, cur_pos 32) returns
781
+ # PCC **-0.04 to -0.17** against HF with it, nondeterministically in 2-3 of its 4
782
+ # batch sizes, across nine runs. Sweeping that real test pins the boundary
783
+ # exactly -- ``k_chunk`` in {32, 64, 128, 256} passes 4/4 at every
784
+ # ``max_cores_per_head_batch`` in {16, 32, 64}; only 512 fails -- so
785
+ # ``max_cores`` is innocent and ``k_chunk`` is the whole effect.
786
+ #
787
+ # No standalone construction reproduces it. ``probes/sdpa_kchunk_rule_probe.py``
788
+ # re-runs the op at bfloat16, at the failing 128-deep cache, with an 8-user paged
789
+ # page table laid out exactly as ``create_mesh_kv_cache`` lays it out, and reads
790
+ # PCC 0.9997 at k512 at every depth from 128 to 4096;
791
+ # ``probes/sdpa_shallow_cache_probe.py`` finds nothing either. The leading
792
+ # explanation is **L1 pressure**: standalone the op owns the whole of L1, while
793
+ # in-model it is co-resident with the layer's sharded activations, expert
794
+ # weights and CCL buffers, and a 512-deep bf16 K chunk is exactly the size that
795
+ # stops fitting. That the boundary is dtype-linked is independently visible --
796
+ # ``k1024/c64`` fails to *build* at bfloat16 (``program.cpp:1722``) and builds
797
+ # fine at bfloat8_b. It is recorded as unexplained-in-detail rather than argued;
798
+ # what is measured is that 512 is unsafe in-model and 256 is not.
799
+ #
800
+ # So: the sweep was redone at bfloat16, 6 x 4 points at five positions
801
+ # (``probes/sdpa_sweep_probe.py``), finalists re-timed over nine positions, and
802
+ # the choice restricted to the in-model-safe ``k_chunk <= 256``. **256/16 is the
803
+ # uniform winner** -- fastest of the safe configs at cur_pos 4095 and above,
804
+ # within 0.6% at 511-2047, and its worst point is +6.4% at cur_pos 127 (19.00 vs
805
+ # 17.86 us for 256/8, i.e. 0.05 ms on a 20 ms iteration). There is no
806
+ # context-dependence worth a runtime switch, so it is **fixed**; a traced decode
807
+ # could not vary it per step anyway. ``q_chunk_size`` stays 32: decode has one
808
+ # query row and 32 is the tile height.
809
+ _SDPA_PAGED_K_CHUNK = 256
810
+ _SDPA_PAGED_MAX_CORES_PER_HEAD = 16
811
+
812
+
813
+ def _paged_cache_depth(kv_cache) -> int:
814
+ """Positions allocated **per user** in a paged cache.
815
+
816
+ ``page_table`` is ``[max_batch, blocks_per_seq]`` and every block holds
817
+ ``block_size`` positions, so this is the length of the logical sequence the
818
+ cache can hold for one user -- which is the quantity ``k_chunk_size`` has to
819
+ respect. See ``_sdpa_k_chunk``.
820
+ """
821
+ return int(kv_cache.page_table.shape[-1]) * int(kv_cache.block_size)
822
+
823
+
824
+ def _sdpa_k_chunk(kv_cache) -> int:
825
+ """``_SDPA_PAGED_K_CHUNK``, clamped to what the cache can actually supply.
826
+
827
+ **``k_chunk_size`` must not exceed the cache's per-user allocated depth**, and
828
+ exceeding it does not raise -- it silently returns garbage. This is the whole
829
+ reason the first adoption of this lever failed its gates, and it is worth
830
+ stating precisely because nothing in the op signature hints at it.
831
+
832
+ How it presents: ``test_multichip_decode_batch`` allocates a **128**-position
833
+ paged cache. At ``k_chunk_size=256`` it returns PCC **-0.10 to +0.06** against
834
+ HF -- noise, not a degraded answer -- but *only when another test has run
835
+ before it in the same process*; run alone it passes. Run the same test after
836
+ ``test_router_windows_partition_global_routing`` and it fails 4/4, at every
837
+ ``max_cores_per_head_batch``. Sweeping ``k_chunk`` through that reproducer
838
+ puts the boundary exactly at the cache depth:
839
+
840
+ k_chunk 32 64 128 -> 7 passed, at max_cores in {8, 16, 32, 64}
841
+ k_chunk 256 -> 4 failed
842
+
843
+ That order-dependence is the tell, and it is what makes the bug so easy to
844
+ miss: the op reads a full ``k_chunk`` past the end of the cache buffer, and
845
+ whether that hurts depends on what the allocator last left there. On a fresh
846
+ device it is zeros and the softmax mask hides it; after another test has
847
+ allocated and freed tensors it is live garbage. **Every standalone probe
848
+ misses this by construction** -- ``probes/sdpa_shallow_cache_probe.py`` and
849
+ ``probes/sdpa_kchunk_rule_probe.py`` both reproduce the shapes, the dtype, the
850
+ 128-deep cache and the multi-user page table exactly, and both read PCC 0.9997
851
+ at k512, because in a probe the cache is the only thing allocated. This is the
852
+ same shape of miss as the stage-04 ``rotary_embedding_llama`` rejection: a
853
+ probe that structurally cannot see the state interaction.
854
+
855
+ So the clamp is not defensive coding, it is the operating range. At the
856
+ shipped ``max_context_len`` (4096 and up, contract 262144) it never binds and
857
+ the config is the tuned 256; at the tests' 128-deep caches it drops to 128,
858
+ which the sweep prices at +6% on the op at cur_pos 127 and 0% past 511.
859
+ """
860
+ depth = _paged_cache_depth(kv_cache)
861
+ # The ``max(32, ...)`` floor exists because SDPA will not take a chunk below
862
+ # one tile. It is the one input that could make this function *violate* the
863
+ # invariant in its own first line, and only when ``block_size < 32``, which
864
+ # this model never configures (the block size is 32 and the page table is at
865
+ # least one block per user). Assert it rather than leave a silent hole: a
866
+ # shallower cache than one tile would return a chunk deeper than the cache.
867
+ assert depth >= 32, (
868
+ f"paged cache depth {depth} is below one tile, so the 32-row floor below would return a "
869
+ "k_chunk_size deeper than the per-user allocated depth -- which SDPA reads past without "
870
+ "raising. Raise block_size (currently "
871
+ f"{int(kv_cache.block_size)}) or the page table width ({int(kv_cache.page_table.shape[-1])})."
872
+ )
873
+ chunk = min(_SDPA_PAGED_K_CHUNK, max(32, depth))
874
+ # SDPA wants a power-of-two chunk; take the largest one that still fits.
875
+ return 1 << (chunk.bit_length() - 1)
876
+
877
+
878
+ #: Program configs are immutable and there are at most a handful of distinct
879
+ #: ones, but they are built **per layer per call** -- 48 times a token on the
880
+ #: decode path and 48 times a prefill chunk. Each build calls
881
+ #: ``device.compute_with_storage_grid_size()``, which is a device query, not a
882
+ #: Python attribute. On the traced decode path that is capture-only and free; on
883
+ #: the *untraced* paths (``run_teacher_forcing``, ``run_prefill_check``, eager
884
+ #: decode) it is 96 device queries per token of pure host time. Memoised on the
885
+ #: grid size rather than the device handle so the cache survives device reopen.
886
+ _SDPA_CONFIG_CACHE: dict = {}
887
+
888
+
889
+ def _cached_sdpa_config(grid, q_chunk, k_chunk, max_cores=None):
890
+ key = (grid.x, grid.y, q_chunk, k_chunk, max_cores)
891
+ cfg = _SDPA_CONFIG_CACHE.get(key)
892
+ if cfg is None:
893
+ kwargs = {} if max_cores is None else {"max_cores_per_head_batch": max_cores}
894
+ cfg = _SDPA_CONFIG_CACHE[key] = ttnn.SDPAProgramConfig(
895
+ compute_with_storage_grid_size=grid, q_chunk_size=q_chunk, k_chunk_size=k_chunk, **kwargs
896
+ )
897
+ return cfg
898
+
899
+
900
+ def _sdpa_program_config(device, kv_cache=None):
901
+ """Program config for SDPA-decode; ``kv_cache`` selects the tuned paged form.
902
+
903
+ Both spellings are the same op family and the same maths; they differ only
904
+ in chunking and core budget. Neither touches dtype, fidelity, the KV cache
905
+ layout, or any collective -- this is a program config on a call the model
906
+ already makes.
907
+ """
908
+ paged = kv_cache is not None and kv_cache.is_paged
909
+ return _cached_sdpa_config(
910
+ device.compute_with_storage_grid_size(),
911
+ 32,
912
+ _sdpa_k_chunk(kv_cache) if paged else 32,
913
+ _SDPA_PAGED_MAX_CORES_PER_HEAD if paged else _SDPA_MAX_CORES_PER_HEAD,
914
+ )
915
+
916
+
917
+ # Prefill has the *same* gap and it is larger in absolute terms:
918
+ # ``attention_prefill`` also called SDPA with no program config, and the op
919
+ # default is quadratic-with-a-bad-constant in S. Same shapes, bfloat16
920
+ # (``doc/optimized_full_model/probes/sdpa_prefill_confirm.py``):
921
+ #
922
+ # S default q128/k128 q256/k256
923
+ # 128 23.92 us 25.72 us 32.68 us
924
+ # 512 58.96 54.14 87.18
925
+ # 1024 230.58 88.36 127.54
926
+ # 2048 741.08 216.03 207.28
927
+ # 4096 2850.25 882.61 451.04
928
+ # 8192 10938.43 2907.58 1956.15
929
+ # 16384 44456.67 11364.22 6527.48
930
+ #
931
+ # The winner *is* length-dependent here, unlike decode, and the two legs cross
932
+ # at S ~= 2048. Prefill could pick at call time -- it is eager, not traced, so
933
+ # the branch is a Python ``if`` and there is no captured trace to invalidate.
934
+ #
935
+ # **It is nevertheless NOT adopted.** ``decoder_layer_prefill_multichip`` passes
936
+ # ``sdpa_program_config=None`` and prefill runs at the op default. Two measured
937
+ # reasons, in this order:
938
+ #
939
+ # 1. **It costs accuracy on the one gate that can see it.** With this config
940
+ # wired, ``run_teacher_forcing`` reads top-1 **0.980** against a baseline of
941
+ # **0.990** on the same tree (top-5 and top-100 stay 1.000). Bisected: the
942
+ # *decode* config alone holds 0.990, the *prefill* config alone drops it to
943
+ # 0.980, so the flip is this and not the lever above
944
+ # (``logs/run_teacher_forcing_leg_prefill.log`` /
945
+ # ``logs/run_teacher_forcing_leg_decode.log``). One greedy token in a hundred
946
+ # is small, but the stage bar is "do not spend accuracy for speed", and here
947
+ # there is no speed to buy it with, which is reason 2.
948
+ # 2. **At the length actually being served it is a loss, not a win.** The
949
+ # readiness reference prompt is **158 tokens**. The table above says the
950
+ # config is *behind* the default below S ~= 384, and the measured TTFTs agree:
951
+ # 3448.79 ms baseline against 3445.31 ms with prefill configured -- noise. The
952
+ # 6.8x is real but it lives at S >= 4096, which nothing in the current gate
953
+ # set exercises.
954
+ #
955
+ # So this is a **verified-fast, accuracy-ungated** lever, left wired and
956
+ # documented rather than taken. What it needs before adoption is a readiness
957
+ # reference with a multi-thousand-token prompt, so the regime where it pays
958
+ # (S >= 4096, 6.3-6.8x on the SDPA op, and 48 of them per prefill) is the same
959
+ # regime the accuracy gate covers. The seam in ``attention_prefill`` exists for
960
+ # exactly that -- same pattern as ``_meta_rope``, which is also built, measured
961
+ # and not adopted.
962
+ #
963
+ # **Arbitrary S keeps working**, so nothing here is blocked on alignment. This
964
+ # was checked and not assumed: S in
965
+ # {1, 3, 31, 33, 100, 129, 255, 257, 1000, 1023, 1025, 2049, 4095, 4097, 5000}
966
+ # all build and run under both chunkings with PCC identical to the default's to
967
+ # five decimals. Prefill is not chunked in this model -- ``prefill_forward``
968
+ # feeds each user's whole logical length in -- so that property is load-bearing
969
+ # for the stage contract, not a nicety.
970
+ #
971
+ # ``q512/k512`` is **rejected** outright: it fails to build at *every* length
972
+ # including 128 (``program.cpp:1722``), so it is a resource limit, not an
973
+ # alignment rule.
974
+ _SDPA_PREFILL_CROSSOVER = 2048
975
+
976
+
977
+ def _sdpa_prefill_program_config(device, seq_len: int):
978
+ """Built and measured; **not wired in**. See the note above before adopting."""
979
+ q_chunk = 256 if seq_len >= _SDPA_PREFILL_CROSSOVER else 128
980
+ return _cached_sdpa_config(device.compute_with_storage_grid_size(), q_chunk, q_chunk)
981
+
982
+
983
+ # --- decode residual RMSNorm, width-sharded (stage 04) ------------------------
984
+ #
985
+ # The stage-03 decode layer spent 40.21 us -- 9.7% -- in its two residual
986
+ # RMSNorms, and the profile says why: both run on **one core**
987
+ # (``../doc/multichip_decoder/ops_perf_multichip_decode.csv.gz``, device 0, rows
988
+ # 134 and 159, 20.081 and 20.127 us, ``CORE COUNT`` 1). A 2048-wide bf16 norm
989
+ # over one 32-row tile is 128 KB in 20 us, i.e. 6.5 GB/s, which is a single
990
+ # core's share of L1 bandwidth and nothing else.
991
+ #
992
+ # ``ttnn.rms_norm`` has a sharded program factory that splits the row across a
993
+ # core grid. Feeding it the same L1 width-shard the DRAM-sharded qkv projection
994
+ # already wants -- 8 cores, one per DRAM bank, ``[32, 256]`` -- gives
995
+ # (``probes/norm_accuracy_probe.py``, trace slope, median of 30):
996
+ #
997
+ # interleaved, no compute config (shipped) 19.82 us max|err vs fp64| 6.711e-02
998
+ # sharded 4 cores, HiFi4 fp32acc 7.53 1.439e-02
999
+ # sharded 8 cores, default 4.26 3.586e-02
1000
+ # sharded 8 cores, HiFi4 fp32acc 4.92 1.686e-02
1001
+ # i2s 0.51 us, s2i 0.53 us
1002
+ #
1003
+ # 8 cores at HiFi4 with fp32 accumulation is **4.0x faster and 4.0x more
1004
+ # accurate** than the shipped call, which passes no compute config at all and so
1005
+ # accumulates the sum of squares in bf16. The reference is torch fp64 over the
1006
+ # bf16-rounded inputs the device actually sees, so "more accurate" is against
1007
+ # the mathematical answer rather than against the other kernel.
1008
+ #
1009
+ # 16 cores and beyond do not pay: the norm itself stops improving (a 2048-wide
1010
+ # row is 64 tiles, so 8 cores already hold 8 tiles each) while the resharding at
1011
+ # both ends grows with the core count.
1012
+ _NORM_SHARD_CORES = _DRAM_BANKS
1013
+
1014
+
1015
+ def _norm_shard_config(dim: int) -> ttnn.MemoryConfig:
1016
+ """The L1 width-shard the sharded norm reads and writes.
1017
+
1018
+ Deliberately ``_width_sharded_l1(dim)``'s spec: at ``dim == hidden_size``
1019
+ this is bit-for-bit the memory config ``attention_decode_optimized`` reshards
1020
+ its input into, so the first norm's output feeds the qkv projection with no
1021
+ conversion at all.
1022
+ """
1023
+ return ttnn.MemoryConfig(
1024
+ ttnn.TensorMemoryLayout.WIDTH_SHARDED,
1025
+ ttnn.BufferType.L1,
1026
+ ttnn.ShardSpec(_bank_row(_NORM_SHARD_CORES), [32, dim // _NORM_SHARD_CORES], ttnn.ShardOrientation.ROW_MAJOR),
1027
+ )
1028
+
1029
+
1030
+ def _norm_program_config(dim: int):
1031
+ block_w = dim // _NORM_SHARD_CORES // 32
1032
+ subblock_w = next(w for w in (4, 3, 2, 1) if block_w % w == 0)
1033
+ return ttnn.LayerNormShardedMultiCoreProgramConfig(
1034
+ compute_with_storage_grid_size=[_NORM_SHARD_CORES, 1],
1035
+ subblock_w=subblock_w,
1036
+ block_h=1, # decode's padded M is exactly one 32-row tile; batch is capped at 32
1037
+ block_w=block_w,
1038
+ inplace=False,
1039
+ )
1040
+
1041
+
1042
+ def _norm_compute_config(device, precision: PrecisionConfig = DEFAULT_PRECISION):
1043
+ return ttnn.init_device_compute_kernel_config(
1044
+ device.arch(),
1045
+ math_fidelity=precision.norm_fidelity,
1046
+ math_approx_mode=False,
1047
+ fp32_dest_acc_en=True,
1048
+ packer_l1_acc=True,
1049
+ )
1050
+
1051
+
1052
+ def decode_residual_norm(
1053
+ x: ttnn.Tensor, weight_rm: ttnn.Tensor, eps: float, precision: PrecisionConfig = DEFAULT_PRECISION
1054
+ ) -> ttnn.Tensor:
1055
+ """One residual-stream RMSNorm at decode shape, width-sharded across 8 cores.
1056
+
1057
+ Takes a DRAM-interleaved ``[1, 1, B, H]`` (B <= 32, padded to one tile) and
1058
+ returns an **L1 width-sharded** tensor in ``_norm_shard_config(H)``. Callers
1059
+ that need it interleaved say so; the first norm's consumer does not.
1060
+ """
1061
+ dim = int(x.shape[-1])
1062
+ assert int(x.shape[-2]) <= 32, (
1063
+ f"decode_residual_norm shards a single 32-row tile; got {int(x.shape[-2])} rows. "
1064
+ "Prefill uses the interleaved rms_norm."
1065
+ )
1066
+ mc = _norm_shard_config(dim)
1067
+ xs = ttnn.to_memory_config(x, mc)
1068
+ out = ttnn.rms_norm(
1069
+ xs,
1070
+ weight=weight_rm,
1071
+ epsilon=eps,
1072
+ program_config=_norm_program_config(dim),
1073
+ memory_config=mc,
1074
+ # ``precision`` rather than the default: this is the only site
1075
+ # ``norm_fidelity`` reaches. It was called with the module default until
1076
+ # the stage-07 review, which meant the field was a documented knob with
1077
+ # no effect and ``R21_norm_hifi2`` measured nothing. The prefill norms
1078
+ # (``decoder_layer_prefill_multichip``) pass no compute config at all and
1079
+ # still take the op default -- ``norm_fidelity`` is a decode-path field,
1080
+ # which is the path the stage ranks on.
1081
+ compute_kernel_config=_norm_compute_config(x.device(), precision),
1082
+ )
1083
+ ttnn.deallocate(xs)
1084
+ return out
1085
+
1086
+
1087
+ def _exact_matmul_config(device, precision: PrecisionConfig = DEFAULT_PRECISION):
1088
+ """HiFi4, so the one-hot window matmul is a copy rather than an approximation.
1089
+
1090
+ The matmul default is LoFi, which keeps ~5 mantissa bits, and that is fine
1091
+ for everything else in this layer -- but here the operand is 0/1 and the
1092
+ intent is to *select* a routing weight, not to compute with it. Measured, the
1093
+ LoFi spelling moved the stitched windows 9.77e-4 away from the single-chip
1094
+ dense routing (one bf16 ulp at these magnitudes) where HiFi4 reproduces them
1095
+ bit-for-bit. The tensor is 4 tiles by 1, so exactness is free.
1096
+ """
1097
+ return ttnn.init_device_compute_kernel_config(
1098
+ device.arch(),
1099
+ math_fidelity=precision.router_window_fidelity,
1100
+ math_approx_mode=False,
1101
+ fp32_dest_acc_en=False,
1102
+ packer_l1_acc=False,
1103
+ )
1104
+
1105
+
1106
+ def _expert_window_matrix(mesh_device, num_experts: int, num_devices: int) -> ttnn.Tensor:
1107
+ """Per-die one-hot selector ``[1, 1, E, E/num_devices]``.
1108
+
1109
+ A TTNN mesh op is SPMD: one program on four dies, so ``ttnn.slice`` cannot
1110
+ take a different start offset per die and there is no way to ask for
1111
+ "columns 32d..32d+31" directly. The device-varying constant is built the
1112
+ only way a mesh tensor can vary by device -- a leading dim of ``num_devices``
1113
+ sharded on dim 0 -- and applied as a matmul.
1114
+
1115
+ The matmul is exact, not approximate: the operand is 0/1, the accumulator is
1116
+ fp32 and the output is bf16, so each selected weight is copied bit-for-bit.
1117
+ K = 128 is 4 tiles and N = 32 is 1, so it is the cheapest op in the router
1118
+ block, and it *replaces* work rather than adding it -- the divide that
1119
+ follows now runs over 32 columns instead of 128.
1120
+ """
1121
+ local = num_experts // num_devices
1122
+ sel = torch.zeros(num_devices, 1, num_experts, local)
1123
+ for d in range(num_devices):
1124
+ for j in range(local):
1125
+ sel[d, 0, d * local + j, j] = 1.0
1126
+ return ttnn.from_torch(
1127
+ sel,
1128
+ dtype=ttnn.bfloat16,
1129
+ layout=ttnn.TILE_LAYOUT,
1130
+ device=mesh_device,
1131
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1132
+ mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=0),
1133
+ )
1134
+
1135
+
1136
+ def upload_multichip_weights(
1137
+ torch_weights: dict[str, torch.Tensor],
1138
+ mesh_device,
1139
+ config: MeshDecoderConfig,
1140
+ expert_dtype=None,
1141
+ meta_rope: bool = False,
1142
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1143
+ ) -> MultichipWeights:
1144
+ """Shard and upload one layer's weights across the mesh.
1145
+
1146
+ Per-die footprint at the shipped dtypes, which is the table
1147
+ ``doc/multichip_decoder/mesh_plan.md`` section 2 computes:
1148
+
1149
+ gate_up [1, 32, 2048, 1536] bfloat4_b 56.623 MB
1150
+ down [1, 32, 768, 2048] bfloat4_b 28.312 MB
1151
+ wqkv [2048, 1280] bfloat8_b x2 copies 5.570 MB
1152
+ wo [1024, 2048] bfloat8_b x2 copies 4.456 MB
1153
+ router [2048, 128] bf16 0.524 MB
1154
+ norms + qk-norms 0.009 MB
1155
+ total 95.49 MB / layer / die
1156
+
1157
+ Every division is exact -- 2048/4, 32/4, 4/4, 128/4 -- so this scheme needs
1158
+ **zero load-time padding**, and the contract's allowance for it goes unused.
1159
+ """
1160
+ a = config.global_config.attention
1161
+ n = config.num_devices
1162
+ # ``expert_dtype`` is the stage-04 spelling and still wins when given (the
1163
+ # multichip tests sweep it); otherwise every dtype below comes from
1164
+ # ``precision``, whose defaults are the values this docstring's table was
1165
+ # measured at.
1166
+ gate_up_dtype = expert_dtype if expert_dtype is not None else precision.experts_gate_up_dtype
1167
+ down_dtype = expert_dtype if expert_dtype is not None else precision.experts_down_dtype
1168
+
1169
+ def replicate(t: torch.Tensor, tensor_dtype, memory_config=ttnn.DRAM_MEMORY_CONFIG) -> ttnn.Tensor:
1170
+ return ttnn.from_torch(
1171
+ t.contiguous().float(),
1172
+ dtype=tensor_dtype,
1173
+ layout=ttnn.TILE_LAYOUT,
1174
+ device=mesh_device,
1175
+ memory_config=memory_config,
1176
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
1177
+ )
1178
+
1179
+ def shard(t: torch.Tensor, dim: int, tensor_dtype, memory_config=ttnn.DRAM_MEMORY_CONFIG) -> ttnn.Tensor:
1180
+ return ttnn.from_torch(
1181
+ t.contiguous().float(),
1182
+ dtype=tensor_dtype,
1183
+ layout=ttnn.TILE_LAYOUT,
1184
+ device=mesh_device,
1185
+ memory_config=memory_config,
1186
+ mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=dim),
1187
+ )
1188
+
1189
+ def as_4d(t: torch.Tensor, pad_to_4d: bool = False) -> torch.Tensor:
1190
+ if pad_to_4d:
1191
+ t = t.reshape(1, 1, 1, -1)
1192
+ while t.dim() < 4:
1193
+ t = t.unsqueeze(0)
1194
+ return t
1195
+
1196
+ wqkv = head_interleaved_wqkv(as_4d(torch_weights["wqkv"]), a, n)
1197
+ wo = as_4d(torch_weights["wo"])
1198
+
1199
+ # Per-die shapes, used by both the shard spec and the assertions below.
1200
+ k_qkv, n_qkv = int(wqkv.shape[-2]), int(wqkv.shape[-1]) // n
1201
+ k_o, n_o = int(wo.shape[-2]) // n, int(wo.shape[-1])
1202
+ assert _dram_sharded_ok(k_qkv, n_qkv), (
1203
+ f"per-die wqkv [{k_qkv}, {n_qkv}] is not bank-divisible; the DRAM-sharded decode "
1204
+ "projections would silently fall back to interleaved and give back stage 02's 1.11x"
1205
+ )
1206
+ assert _dram_sharded_ok(k_o, n_o), f"per-die wo [{k_o}, {n_o}] is not bank-divisible"
1207
+
1208
+ def dram_sharded(t: torch.Tensor, dim: int, k: int, n_local: int, tensor_dtype) -> ttnn.Tensor:
1209
+ """Width-shard the per-die weight one shard per DRAM bank, then mesh-shard it.
1210
+
1211
+ Two independent shardings compose here and it is worth being explicit
1212
+ about which is which: ``mesh_mapper`` fractures the tensor *across dies*
1213
+ (TP), while ``memory_config`` fractures each die's piece across that
1214
+ die's 8 DRAM banks (stage 02's decode projection layout). The shard spec
1215
+ is therefore written in per-die elements, not global ones.
1216
+ """
1217
+ return shard(
1218
+ t,
1219
+ dim,
1220
+ tensor_dtype,
1221
+ ttnn.MemoryConfig(
1222
+ ttnn.TensorMemoryLayout.WIDTH_SHARDED,
1223
+ ttnn.BufferType.DRAM,
1224
+ ttnn.ShardSpec(_bank_row(_DRAM_BANKS), [k, n_local // _DRAM_BANKS], ttnn.ShardOrientation.ROW_MAJOR),
1225
+ ),
1226
+ )
1227
+
1228
+ experts = OptimizedWeights(
1229
+ # [E, 2I, H] -> [1, E, H, 2I], sharded on the expert dim.
1230
+ gate_up_proj=shard(torch_weights["experts_gate_up"].transpose(-2, -1).unsqueeze(0), 1, gate_up_dtype),
1231
+ # [E, H, I] -> [1, E, I, H], sharded on the expert dim.
1232
+ down_proj=shard(torch_weights["experts_down"].transpose(-2, -1).unsqueeze(0), 1, down_dtype),
1233
+ attention=AttentionWeights(
1234
+ # Column split (head-interleaved, see head_interleaved_wqkv).
1235
+ wqkv=shard(wqkv, -1, precision.attention_qkv_dtype),
1236
+ # Row split by Q head. Contiguous *because* the Q head assignment
1237
+ # above is contiguous per die: die d owns rows 1024d..1024d+1023.
1238
+ wo=shard(wo, -2, precision.attention_wo_dtype),
1239
+ q_norm=replicate(as_4d(torch_weights["q_norm"], pad_to_4d=True), precision.norm_weight_dtype),
1240
+ k_norm=replicate(as_4d(torch_weights["k_norm"], pad_to_4d=True), precision.norm_weight_dtype),
1241
+ ),
1242
+ wqkv_decode=dram_sharded(wqkv, -1, k_qkv, n_qkv, precision.attention_qkv_dtype),
1243
+ wo_decode=dram_sharded(wo, -2, k_o, n_o, precision.attention_wo_dtype),
1244
+ )
1245
+
1246
+ # Stage 04. The Meta-ordered decode twin, built **only when asked**. It is
1247
+ # not the shipped path (see ``_meta_rope``); it exists so
1248
+ # ``probes/rope_layer_probe.py`` can re-measure the rejection rather than
1249
+ # cite it. Off by default, so the shipped upload pays no extra DRAM.
1250
+ experts_meta = None
1251
+ if meta_rope:
1252
+ # The channel permutation is applied to the *pre-interleave* wqkv, which
1253
+ # is safe because it reorders channels **within** a head and
1254
+ # ``head_interleaved_wqkv`` only reorders whole heads -- the two commute.
1255
+ # V is untouched, and so are ``wo`` and every expert weight, which are
1256
+ # shared objects here rather than copies.
1257
+ wqkv_meta = head_interleaved_wqkv(
1258
+ permute_wqkv_to_meta(
1259
+ as_4d(torch_weights["wqkv"]),
1260
+ n_heads=a.num_attention_heads,
1261
+ n_kv_heads=a.num_key_value_heads,
1262
+ head_dim=a.head_dim,
1263
+ ),
1264
+ a,
1265
+ n,
1266
+ )
1267
+ experts_meta = replace(
1268
+ experts,
1269
+ attention=replace(
1270
+ experts.attention,
1271
+ q_norm=replicate(
1272
+ as_4d(permute_head_vector_to_meta(torch_weights["q_norm"], head_dim=a.head_dim), pad_to_4d=True),
1273
+ precision.norm_weight_dtype,
1274
+ ),
1275
+ k_norm=replicate(
1276
+ as_4d(permute_head_vector_to_meta(torch_weights["k_norm"], head_dim=a.head_dim), pad_to_4d=True),
1277
+ precision.norm_weight_dtype,
1278
+ ),
1279
+ ),
1280
+ wqkv_decode=dram_sharded(wqkv_meta, -1, k_qkv, n_qkv, precision.attention_qkv_dtype),
1281
+ )
1282
+
1283
+ router = torch_weights["router"]
1284
+
1285
+ def norm_row_major(t: torch.Tensor) -> ttnn.Tensor:
1286
+ """The same vector the tiled copy holds, in the layout the sharded
1287
+ ``rms_norm`` program factory reads: ROW_MAJOR ``[1, 1, dim/32, 32]``."""
1288
+ flat = t.reshape(-1)
1289
+ return ttnn.from_torch(
1290
+ flat.reshape(1, 1, flat.numel() // 32, 32).contiguous().float(),
1291
+ dtype=precision.norm_weight_dtype,
1292
+ layout=ttnn.ROW_MAJOR_LAYOUT,
1293
+ device=mesh_device,
1294
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1295
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
1296
+ )
1297
+
1298
+ return MultichipWeights(
1299
+ input_layernorm=replicate(torch_weights["input_layernorm"].reshape(1, 1, 1, -1), precision.norm_weight_dtype),
1300
+ post_attention_layernorm=replicate(
1301
+ torch_weights["post_attention_layernorm"].reshape(1, 1, 1, -1), precision.norm_weight_dtype
1302
+ ),
1303
+ router=replicate(router.T.contiguous().reshape(1, 1, router.shape[1], router.shape[0]), precision.router_dtype),
1304
+ expert_window=_expert_window_matrix(mesh_device, config.global_config.moe.num_experts, n),
1305
+ experts=experts,
1306
+ experts_meta=experts_meta,
1307
+ input_layernorm_rm=norm_row_major(torch_weights["input_layernorm"]),
1308
+ post_attention_layernorm_rm=norm_row_major(torch_weights["post_attention_layernorm"]),
1309
+ )
1310
+
1311
+
1312
+ def create_mesh_kv_cache(
1313
+ mesh_device,
1314
+ config: MeshDecoderConfig,
1315
+ max_batch: int,
1316
+ max_seq_len: int,
1317
+ block_size: int | None = None,
1318
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1319
+ ) -> KVCache:
1320
+ """Allocate the *local* KV cache: 1 KV head per die, not 4.
1321
+
1322
+ This is where the TP factor buys capacity rather than speed. Per die the
1323
+ cache is ``[.., 1, .., 128]`` instead of ``[.., 4, .., 128]``, i.e. 512 B
1324
+ per token per layer instead of 2048 -- 6.44 GB at the advertised 262144
1325
+ context over 48 layers, against 25.77 GB on one die. One die cannot hold
1326
+ this model at full context; four can, with room to spare. See
1327
+ ``config/context_contract.json``.
1328
+
1329
+ The buffers are *replicated at allocation* because they are zeros, and
1330
+ diverge the moment the first token is written -- each die holds a different
1331
+ KV head. The page table is genuinely identical on every die: paging is a
1332
+ logical-to-physical block mapping and does not depend on which head lives
1333
+ where.
1334
+ """
1335
+ local = config.local_attention
1336
+ if block_size is None:
1337
+ shape = (max_batch, local.num_key_value_heads, max_seq_len, local.head_dim)
1338
+ page_table = None
1339
+ else:
1340
+ blocks_per_seq = math.ceil(max_seq_len / block_size)
1341
+ shape = (max_batch * blocks_per_seq, local.num_key_value_heads, block_size, local.head_dim)
1342
+ page_table = ttnn.from_torch(
1343
+ torch.arange(max_batch * blocks_per_seq, dtype=torch.int32).reshape(max_batch, blocks_per_seq),
1344
+ dtype=ttnn.int32,
1345
+ layout=ttnn.ROW_MAJOR_LAYOUT,
1346
+ device=mesh_device,
1347
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
1348
+ )
1349
+
1350
+ k, v = (
1351
+ ttnn.from_torch(
1352
+ torch.zeros(shape),
1353
+ dtype=precision.kv_cache_dtype,
1354
+ layout=ttnn.TILE_LAYOUT,
1355
+ device=mesh_device,
1356
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
1357
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
1358
+ )
1359
+ for _ in range(2)
1360
+ )
1361
+ return KVCache(k=k, v=v, page_table=page_table, block_size=block_size or 0)
1362
+
1363
+
1364
+ def build_local_sparsity(mesh_device, local_moe: MoEConfig) -> ttnn.Tensor:
1365
+ """All-ones prefill sparsity over this die's 32 experts, replicated."""
1366
+ return ttnn.from_torch(
1367
+ torch.ones(1, 1, 1, local_moe.num_experts, dtype=torch.bfloat16),
1368
+ layout=ttnn.ROW_MAJOR_LAYOUT,
1369
+ dtype=ttnn.bfloat16,
1370
+ device=mesh_device,
1371
+ mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
1372
+ )
1373
+
1374
+
1375
+ # --- router -------------------------------------------------------------------
1376
+
1377
+
1378
+ def router_forward_threshold(
1379
+ x: ttnn.Tensor,
1380
+ w_router: ttnn.Tensor,
1381
+ window: ttnn.Tensor,
1382
+ config: MoEConfig,
1383
+ local_moe: MoEConfig,
1384
+ ) -> ttnn.Tensor:
1385
+ """``router_forward_multichip`` with the dense vector built by a threshold
1386
+ comparison instead of a scatter, so nothing leaves TILE layout.
1387
+
1388
+ Stage 03 inherited stage 02's routing tail, whose shape is
1389
+ ``topk -> untilize(zeros) / untilize(indices) / untilize(values) ->
1390
+ scatter -> tilize``: ``ttnn.scatter`` only accepts ROW_MAJOR, and every
1391
+ consumer of the dense vector (both matmuls, the divide) needs TILE. Stage 02
1392
+ recorded that round trip as *not removable* on those grounds. It is
1393
+ removable -- by not scattering.
1394
+
1395
+ ``topk(sorted=True)`` already returns the 8th-largest logit in column 7, and
1396
+ the top-8 set is exactly ``{j : logit_j >= that}``. So the same dense vector
1397
+ is
1398
+
1399
+ dense = exp(logits - top_max) * (logits >= top_logits[..., 7])
1400
+
1401
+ computed over all 128 columns, entirely in TILE. The surviving values are
1402
+ ``ttnn.exp`` of the same fp32 inputs the scatter path fed it, so the result
1403
+ is bit-identical -- **unless two logits tie exactly at rank 8**, in which
1404
+ case this selects both and the scatter path selects one. With fp32 logits
1405
+ accumulated over K=2048 that does not happen on real weights, and
1406
+ ``test_router_windows_partition_global_routing`` asserts the equality at
1407
+ ``max |diff| = 0.0``, so a tie would fail loudly rather than drift.
1408
+
1409
+ **Measured, and rejected.** It removes rows 190-197 of the stage-04 decode
1410
+ profile -- ``zeros_like`` 1.210, two ``typecast`` 2.537, three ``untilize``
1411
+ 4.654, ``scatter`` 3.030, ``tilize`` 5.576 = **17.007 us** -- and is
1412
+ nonetheless **0.8% slower on the layer**: 0.4382 / 0.4382 ms against the
1413
+ shipped 0.4348 / 0.4346 over two interleaved passes
1414
+ (``doc/optimized_multichip_decoder/probes/layer_levers2.py``). Widening the
1415
+ softmax's ``sub`` and ``exp`` from 8 columns to 128, plus the ``ge`` and the
1416
+ ``mul``, costs more than the layout conversions save. The output is
1417
+ bit-identical on all four dies (``max|diff| 0.000e+00``), which is also the
1418
+ evidence that no two logits tie at rank 8.
1419
+
1420
+ Kept rather than deleted because the arithmetic is the useful part: stage
1421
+ 02 recorded this round trip as *not removable*, and it is.
1422
+ """
1423
+ assert config.norm_topk_prob
1424
+ e = config.num_experts
1425
+
1426
+ logits = ttnn.linear(x, w_router, dtype=ttnn.float32, memory_config=ttnn.DRAM_MEMORY_CONFIG)
1427
+ top_logits, _idx = ttnn.topk(logits, k=config.num_experts_per_tok, dim=-1, largest=True, sorted=True)
1428
+
1429
+ rows = top_logits.shape[2]
1430
+ top_max = ttnn.slice(top_logits, [0, 0, 0, 0], [1, 1, rows, 1])
1431
+ cutoff = ttnn.slice(top_logits, [0, 0, 0, config.num_experts_per_tok - 1], [1, 1, rows, config.num_experts_per_tok])
1432
+
1433
+ # exp(l - max) over the whole row; every entry is in (0, 1], so nothing can
1434
+ # overflow and the losers underflow towards zero before the mask even runs.
1435
+ weights = ttnn.exp(ttnn.sub(logits, top_max))
1436
+ dense = ttnn.typecast(ttnn.mul(weights, ttnn.ge(logits, cutoff)), ttnn.bfloat16)
1437
+
1438
+ total = ttnn.matmul(dense, _ones_column(x.device(), e), dtype=ttnn.bfloat16)
1439
+ local = ttnn.matmul(dense, window, dtype=ttnn.bfloat16, compute_kernel_config=_exact_matmul_config(x.device()))
1440
+ guarded = ttnn.maximum(total, 1e-30)
1441
+ normalised = ttnn.div(local, guarded)
1442
+ assert int(normalised.shape[-1]) == local_moe.num_experts
1443
+ for t in (logits, top_logits, _idx, top_max, cutoff, weights, dense, total, local, guarded):
1444
+ ttnn.deallocate(t)
1445
+ return normalised
1446
+
1447
+
1448
+ def router_forward_multichip(
1449
+ x: ttnn.Tensor,
1450
+ w_router: ttnn.Tensor,
1451
+ window: ttnn.Tensor,
1452
+ config: MoEConfig,
1453
+ local_moe: MoEConfig,
1454
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1455
+ ) -> ttnn.Tensor:
1456
+ """Replicated global routing, returning this die's ``[1, 1, S, 32]`` window.
1457
+
1458
+ Identical arithmetic to ``optimized_decoder.router_forward_optimized`` --
1459
+ selection on raw fp32 logits, softmax over the 8 survivors, neither keepdim
1460
+ reduction spelled as a ttnn reduction -- with one op added and one op made
1461
+ four times narrower.
1462
+
1463
+ **Why the whole router is replicated.** Top-8 of 128 needs the global logit
1464
+ vector, and there is nothing worth fracturing anyway: the router matmul has
1465
+ N = 128 = 4 tiles so it can use 4 cores, and ``ttnn.topk`` over a single
1466
+ 128-wide row occupies exactly 1. Splitting N four ways would give each die
1467
+ one tile and one core, and would additionally need a collective *inside* the
1468
+ routing path to reassemble the logits before the top-k. So each die computes
1469
+ the full 128-way routing on the bit-identical replicated activation and
1470
+ takes its own window out of the result -- no collective, at the price of
1471
+ 88.9 us of decode device time that four dies pay in full.
1472
+
1473
+ **The correctness assumption this makes, stated plainly.** The four windows
1474
+ are a partition of the global top-8 only if all four dies agree on which 8
1475
+ experts won. The inputs are bit-identical and the program is the same, so
1476
+ ``ttnn.topk`` should return identical indices -- but that is a *tie-breaking
1477
+ determinism* claim, and if it ever failed the layer would be silently wrong
1478
+ with no shape error and only a PCC drift to show for it. It is asserted
1479
+ directly by ``test_topk_is_identical_across_dies`` rather than argued.
1480
+ """
1481
+ assert config.norm_topk_prob, (
1482
+ "router selects on raw logits, which relies on the softmax denominator "
1483
+ "cancelling during top-k renormalisation; that only holds when norm_topk_prob is True"
1484
+ )
1485
+
1486
+ logits = ttnn.linear(x, w_router, dtype=ttnn.float32, memory_config=ttnn.DRAM_MEMORY_CONFIG)
1487
+ return _router_tail(logits, window, config, local_moe, x.device(), precision)
1488
+
1489
+
1490
+ def _router_tail(
1491
+ logits,
1492
+ window,
1493
+ config: MoEConfig,
1494
+ local_moe: MoEConfig,
1495
+ device,
1496
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1497
+ ) -> ttnn.Tensor:
1498
+ """Top-8, softmax over the survivors, this die's 32-expert window.
1499
+
1500
+ Split out of ``router_forward_multichip`` so a probe can vary where the
1501
+ logits are produced without duplicating the tail.
1502
+ """
1503
+ top_logits, top_indices = ttnn.topk(logits, k=config.num_experts_per_tok, dim=-1, largest=True, sorted=True)
1504
+
1505
+ top_max = ttnn.slice(top_logits, [0, 0, 0, 0], [1, 1, top_logits.shape[2], 1])
1506
+ exp_logits = ttnn.exp(ttnn.sub(top_logits, top_max))
1507
+
1508
+ zeros = ttnn.typecast(ttnn.zeros_like(logits), ttnn.bfloat16)
1509
+ dense = ttnn.scatter(zeros, dim=-1, index=top_indices, src=ttnn.typecast(exp_logits, ttnn.bfloat16))
1510
+
1511
+ # The denominator is the sum over all 128 -- which is the sum over the 8
1512
+ # survivors, since the scatter fills a field of exact zeros -- and must stay
1513
+ # global: normalising within a window would renormalise each die's share to
1514
+ # 1 and the four contributions would sum to 4.
1515
+ total = ttnn.matmul(dense, _ones_column(device, config.num_experts), dtype=ttnn.bfloat16)
1516
+ local = ttnn.matmul(
1517
+ dense, window, dtype=ttnn.bfloat16, compute_kernel_config=_exact_matmul_config(device, precision)
1518
+ )
1519
+ # Same clamp as the single-chip router, and for the same reason: after the
1520
+ # scatter the divide runs over whole tiles, and the tile row-padding has a
1521
+ # zero numerator *and* a zero denominator, which unguarded ttnn.div returns
1522
+ # as +inf. Every real row's denominator is >= 1 (sorted=True makes column 0
1523
+ # of exp_logits exactly exp(0) = 1), so the clamp cannot touch one.
1524
+ guarded = ttnn.maximum(total, 1e-30)
1525
+ normalised = ttnn.div(local, guarded)
1526
+ assert int(normalised.shape[-1]) == local_moe.num_experts
1527
+ for t in (logits, top_logits, top_indices, top_max, exp_logits, dense, total, local, guarded):
1528
+ ttnn.deallocate(t)
1529
+ return normalised
1530
+
1531
+
1532
+ # --- experts ------------------------------------------------------------------
1533
+
1534
+
1535
+ def moe_decode_multichip(
1536
+ x: ttnn.Tensor,
1537
+ routing: ttnn.Tensor,
1538
+ weights: OptimizedWeights,
1539
+ local_moe: MoEConfig,
1540
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1541
+ ) -> ttnn.Tensor:
1542
+ """Decode expert pass over this die's 32 experts. Returns a *partial* sum.
1543
+
1544
+ Structurally ``optimized_decoder.moe_decode_optimized`` with the local
1545
+ expert count, and one difference that is not cosmetic: **``nnz`` is
1546
+ ``None``**.
1547
+
1548
+ Stage 02 passes ``nnz = top_k * batch``, exact because every one of the
1549
+ global top-8 is computed on the single die. Under EP the count of live
1550
+ experts in *this* die's window is data-dependent -- 0 to 8, mean 2, and
1551
+ different on each die -- while a mesh op is SPMD and compiles one kernel for
1552
+ all four. Passing any host-computed value would deadlock the board the first
1553
+ time the routing was unbalanced (``sparse_matmul_device_operation.cpp``
1554
+ 205-211, tt-metal #45943), silently unless the watcher is on. ``nnz=None``
1555
+ switches the in0 sender to reading the sparsity page at runtime and
1556
+ multicasting a per-slot valid flag; the loop still visits all 32 slots but
1557
+ only reads weights and does math for the live ones.
1558
+
1559
+ Measured cost of dynamic mode, decode M=1, bfp4/LoFi, trace-slope:
1560
+
1561
+ E=128 nnz=8 (single-die baseline) 139.45 + 125.20 = 264.65 us
1562
+ E=32 nnz=None (this path) 60.67 + 63.29 = 123.96 us 2.13x
1563
+ E=32 nnz=8 (exact, illegal here) 82.72 + 58.21 = 140.93 us
1564
+ E=128 nnz=None (dynamic at full E) 243.08 + 249.03 = 492.11 us 0.54x
1565
+
1566
+ **That 2.13x did not survive.** Re-measured at the shipped shapes -- E=32,
1567
+ M=1, bfloat4_b, LoFi, L1 output, stage 02's tuned block widths -- dynamic
1568
+ mode costs 158.01 us against an exact ``nnz``'s 107.73, **1.47x**, and the
1569
+ multichip decode profile reads 82.65 us for the pair against the single
1570
+ chip's 92.06, i.e. **1.11x, not 2.13x**
1571
+ (``probes/nnz_cost_probe.py``, ``doc/multichip_decoder/work_log.md`` section
1572
+ 8). The sweep above was a DRAM-out, random-weight microbenchmark whose E=128
1573
+ baseline read 264.65 where the profiled layer reads 92.06; only ratios were
1574
+ taken from it, and the ratio was still wrong, because the overhead it hid is
1575
+ additive rather than proportional.
1576
+
1577
+ The E=32/nnz=8 row is also the answer to "why not capacity padding": the
1578
+ only capacity that can never be exceeded is 8, and building a fixed-count
1579
+ sparsity on device needs a second ``topk`` over the local 32 -- **26.32 us**
1580
+ on one core in the decode profile, more than the ~26 us it would save.
1581
+ Any smaller capacity drops experts, which changes the model output.
1582
+ """
1583
+ batch = x.shape[2]
1584
+ n_experts = local_moe.num_experts
1585
+ hidden_size = local_moe.hidden_size
1586
+ inter = local_moe.moe_intermediate_size
1587
+
1588
+ sparsity = ttnn.to_layout(routing, ttnn.ROW_MAJOR_LAYOUT)
1589
+ expert_memory_config = _decode_expert_memory_config(batch, local_moe)
1590
+ output_tile = ttnn.Tile([32, 32])
1591
+ compute_config = _expert_compute_kernel_config(x.device(), precision)
1592
+ gate_up_config = _tuned_sparse_matmul_config(1, 2 * inter, hidden_size, precision.experts_gate_up_in0_block_w)
1593
+ down_config = _tuned_sparse_matmul_config(1, hidden_size, inter, precision.experts_down_in0_block_w)
1594
+
1595
+ x_batched = ttnn.reshape(x, (1, batch, 1, hidden_size))
1596
+ fused = ttnn.sparse_matmul(
1597
+ x_batched,
1598
+ weights.gate_up_proj,
1599
+ sparsity=sparsity,
1600
+ nnz=None, # see docstring -- a host-computed nnz deadlocks the board here
1601
+ memory_config=expert_memory_config,
1602
+ output_tile=output_tile,
1603
+ program_config=gate_up_config,
1604
+ compute_kernel_config=compute_config,
1605
+ dtype=precision.activation_dtype,
1606
+ )
1607
+ packed_width = fused.shape[-1]
1608
+ fused = ttnn.reshape(fused, (batch, n_experts, packed_width))
1609
+
1610
+ half = packed_width // 2
1611
+ gate = ttnn.slice(fused, [0, 0, 0], [batch, n_experts, half])
1612
+ up = ttnn.slice(fused, [0, 0, half], [batch, n_experts, packed_width])
1613
+ ttnn.deallocate(fused)
1614
+
1615
+ down_input = ttnn.reshape(ttnn.mul(ttnn.silu(gate), up), (batch, n_experts, 1, half))
1616
+ ttnn.deallocate(gate)
1617
+ ttnn.deallocate(up)
1618
+
1619
+ down = ttnn.sparse_matmul(
1620
+ down_input,
1621
+ weights.down_proj,
1622
+ sparsity=sparsity,
1623
+ nnz=None,
1624
+ memory_config=expert_memory_config,
1625
+ output_tile=output_tile,
1626
+ program_config=down_config,
1627
+ is_input_a_sparse=True,
1628
+ is_input_b_sparse=False, # selects batch_length_A = B * E; see the single-chip docstring
1629
+ compute_kernel_config=compute_config,
1630
+ dtype=precision.activation_dtype,
1631
+ )
1632
+ ttnn.deallocate(down_input)
1633
+
1634
+ # The multiply by the routing weight is what makes a skipped slot harmless:
1635
+ # a die holding none of the global top-8 multiplies 32 untouched output
1636
+ # slots by exact zero and contributes an exact zero to the all-reduce.
1637
+ # test_expert_window_can_be_empty pins that, because "untouched" would not
1638
+ # be enough if the op left a NaN there.
1639
+ states = ttnn.reshape(down, (batch, n_experts, hidden_size))
1640
+ states = ttnn.mul(states, ttnn.reshape(routing, (batch, n_experts, 1)))
1641
+ states = ttnn.unsqueeze_to_4D(ttnn.sum(states, dim=1))
1642
+ return ttnn.reshape(states, (1, 1, batch, hidden_size), (1, 1, max(32, batch), hidden_size))
1643
+
1644
+
1645
+ # --- the layer ----------------------------------------------------------------
1646
+
1647
+
1648
+ def decoder_layer_prefill_multichip(
1649
+ x: ttnn.Tensor,
1650
+ weights: MultichipWeights,
1651
+ config: MeshDecoderConfig,
1652
+ ctx: MeshContext,
1653
+ cos_cache: ttnn.Tensor,
1654
+ sin_cache: ttnn.Tensor,
1655
+ sparsity: ttnn.Tensor,
1656
+ kv_cache: KVCache | None = None,
1657
+ user_id: int = 0,
1658
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1659
+ start_pos: int = 0,
1660
+ chunk_page_table=None,
1661
+ fill_page_table=None,
1662
+ fill_len: int | None = None,
1663
+ ) -> ttnn.Tensor:
1664
+ """Prefill one layer on the mesh. ``x`` / return replicated ``[1, 1, S, 2048]``.
1665
+
1666
+ ``S`` is arbitrary. Nothing in the multichip path adds an alignment
1667
+ constraint: the collectives scatter on dim 3 (hidden, 2048, fixed), and the
1668
+ only padding in play is ``moe_prefill_optimized``'s internal chunk padding,
1669
+ which is the single-chip behaviour and is sliced back inside that function.
1670
+ """
1671
+ eps = config.global_config.rms_norm_eps
1672
+
1673
+ normed = ttnn.rms_norm(x, weight=weights.input_layernorm, epsilon=eps)
1674
+ attn_partial = attention_prefill(
1675
+ normed,
1676
+ weights.experts.attention,
1677
+ config.local_attention,
1678
+ cos_cache,
1679
+ sin_cache,
1680
+ kv_cache,
1681
+ user_id,
1682
+ # ``None`` at the shipped precision, which is the op default and what
1683
+ # every prefill number was measured at; see
1684
+ # ``optimized_decoder._attention_compute_kernel_config``.
1685
+ compute_kernel_config=_attention_compute_kernel_config(x.device(), precision),
1686
+ activation_dtype=precision.activation_dtype,
1687
+ # NOT adopted -- see _sdpa_prefill_program_config. The seam is wired and
1688
+ # the config is built and measured; passing it costs a top-1 point on
1689
+ # run_teacher_forcing, so prefill stays at the op default.
1690
+ sdpa_program_config=None,
1691
+ start_pos=start_pos,
1692
+ chunk_page_table=chunk_page_table,
1693
+ fill_page_table=fill_page_table,
1694
+ fill_len=fill_len,
1695
+ )
1696
+ ttnn.deallocate(normed)
1697
+ attn_out = all_reduce_prefill(attn_partial, ctx, precision)
1698
+ ttnn.deallocate(attn_partial)
1699
+ hidden = ttnn.add(x, attn_out)
1700
+ ttnn.deallocate(attn_out)
1701
+
1702
+ normed = ttnn.rms_norm(hidden, weight=weights.post_attention_layernorm, epsilon=eps)
1703
+ routing = router_forward_multichip(
1704
+ normed, weights.router, weights.expert_window, config.global_config.moe, config.local_moe, precision
1705
+ )
1706
+ moe_partial = moe_prefill_optimized(normed, routing, weights.experts, config.local_moe, sparsity, precision)
1707
+ ttnn.deallocate(normed)
1708
+ ttnn.deallocate(routing)
1709
+ moe_out = all_reduce_prefill(moe_partial, ctx, precision)
1710
+ ttnn.deallocate(moe_partial)
1711
+
1712
+ out = ttnn.add(hidden, moe_out)
1713
+ ttnn.deallocate(hidden)
1714
+ ttnn.deallocate(moe_out)
1715
+ return out
1716
+
1717
+
1718
+ def decoder_layer_decode_multichip(
1719
+ x: ttnn.Tensor,
1720
+ weights: MultichipWeights,
1721
+ config: MeshDecoderConfig,
1722
+ ctx: MeshContext,
1723
+ cos_cache: ttnn.Tensor,
1724
+ sin_cache: ttnn.Tensor,
1725
+ kv_cache: KVCache,
1726
+ current_pos: ttnn.Tensor,
1727
+ token_index: int,
1728
+ rope=None,
1729
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1730
+ active_mask: ttnn.Tensor | None = None,
1731
+ ) -> ttnn.Tensor:
1732
+ """Decode one token per user on the mesh. ``x`` / return ``[1, 1, B, 2048]``.
1733
+
1734
+ Input and output layouts are the same replicated tensor, which is the point:
1735
+ 48 of these stack with no boundary conversion, and the stacked model pays
1736
+ the two all-reduces per layer and nothing else. **That is the inter-layer
1737
+ residual layout contract**, and stage 04 keeps it unchanged while moving
1738
+ every *intra*-layer boundary it can into L1 shards -- see
1739
+ ``doc/optimized_multichip_decoder/README.md``.
1740
+
1741
+ Both residual norms are width-sharded (``decode_residual_norm``). The first
1742
+ one's output is already in ``attention_decode_optimized``'s qkv input shard,
1743
+ so it crosses into attention with no conversion; the second one's output
1744
+ feeds the router projection sharded and the expert path interleaved.
1745
+ """
1746
+ eps = config.global_config.rms_norm_eps
1747
+
1748
+ normed = decode_residual_norm(x, weights.input_layernorm_rm, eps, precision)
1749
+ # ``rope`` is the stage-05 seam. It defaults to ``None`` and therefore to
1750
+ # ``_apply_rope`` -- ``ttnn.experimental.rotary_embedding`` with a **Python
1751
+ # int** ``token_index``, which is what every stage-03/04 number was measured
1752
+ # at and what the single-layer tests still exercise. That spelling cannot be
1753
+ # replayed: the position is a compile-time argument, so a captured trace
1754
+ # rotates every later token at the position it was captured at. The full
1755
+ # model therefore passes ``model._rope_decode``, which is
1756
+ # ``ttnn.experimental.rotary_embedding_hf(is_decode_mode=True)`` reading a
1757
+ # **per-user cos/sin pair gathered on device** from a position tensor the
1758
+ # trace itself advances. Same HF ``rotate_half`` channel convention, so the
1759
+ # KV cache convention, prefill, and every weight are untouched -- which is
1760
+ # exactly what stage 04's rejected ``rotary_embedding_llama`` lever could not
1761
+ # offer (README limitation 4).
1762
+ # The rotary stays the **HF** op. ``rotary_embedding_llama`` is 3.05x faster
1763
+ # standalone and bit-identical there, but it cannot be adopted for decode
1764
+ # alone: it needs Meta channel order, and the KV cache prefill already wrote
1765
+ # is in HF order, so SDPA would score a Meta-ordered Q against HF-ordered
1766
+ # keys. Measured, not argued -- PCC 0.193 against a prefill-primed cache
1767
+ # where a fresh cache reads 0.99997 (``probes/rope_layer_probe.py``). See
1768
+ # ``_meta_rope`` and ``README.md`` limitation 4.
1769
+ attn_partial = attention_decode_optimized(
1770
+ normed,
1771
+ weights.experts,
1772
+ config.local_attention,
1773
+ cos_cache,
1774
+ sin_cache,
1775
+ kv_cache,
1776
+ current_pos,
1777
+ token_index,
1778
+ # Both paths are configured now. The contiguous one needs the 64-core cap
1779
+ # to clear a TT_FATAL; the paged one -- what the full model runs -- takes
1780
+ # the swept k256/c16 config (k clamped to the cache depth), which is flat
1781
+ # in cur_pos where the op default is linear in it. See
1782
+ # _sdpa_program_config and _sdpa_k_chunk.
1783
+ sdpa_program_config=_sdpa_program_config(x.device(), kv_cache),
1784
+ rope=rope,
1785
+ precision=precision,
1786
+ )
1787
+ ttnn.deallocate(normed)
1788
+ attn_out = all_reduce_decode(attn_partial, ctx, precision)
1789
+ ttnn.deallocate(attn_partial)
1790
+ hidden = ttnn.add(x, attn_out)
1791
+ ttnn.deallocate(attn_out)
1792
+
1793
+ normed_sharded = decode_residual_norm(hidden, weights.post_attention_layernorm_rm, eps, precision)
1794
+ # The router projection reads the shard directly -- N = 128 is 4 tiles, so
1795
+ # the matmul uses 4 cores either way, but a width-sharded L1 in0 turns a
1796
+ # 24.62 us DRAM-interleaved read into 5.85 us of L1 with bit-identical
1797
+ # output (``probes/norm_router_probe.py``, max|diff| exactly 0.0).
1798
+ routing = router_forward_multichip(
1799
+ normed_sharded, weights.router, weights.expert_window, config.global_config.moe, config.local_moe, precision
1800
+ )
1801
+ if active_mask is not None:
1802
+ # Zero the routing weights of every slot that holds no live request.
1803
+ #
1804
+ # ``routing`` *is* ``sparse_matmul``'s sparsity tensor, and its nonzero
1805
+ # count is the amount of expert math the op does: with ``nnz=None`` the
1806
+ # kernel reads the sparsity page at runtime and only fetches weights and
1807
+ # multiplies for the live ``(row, expert)`` pairs. A serving decode batch
1808
+ # is padded to the configured ``max_num_seqs`` with inactive rows, and an
1809
+ # inactive row's garbage hidden state still routes to a full top-8 -- so
1810
+ # without this a 32-slot server does 32 rows of expert work no matter how
1811
+ # many users are actually connected. See
1812
+ # ``doc/optimized_vllm/probes/batch_decode_control.py``.
1813
+ #
1814
+ # ``active_mask`` is derived **on device** from ``current_pos`` inside the
1815
+ # same traced graph (``Qwen3CoderModel._decode_active_mask``), so it can
1816
+ # never be stale: ``ttnn.plus_one(..., skip_negative_entries=True)`` leaves
1817
+ # an inactive row at ``-1`` forever, and a row that changes hands only does
1818
+ # so through a host reinstall of ``current_pos``.
1819
+ gated = ttnn.mul(routing, active_mask)
1820
+ ttnn.deallocate(routing)
1821
+ routing = gated
1822
+ # ``sparse_matmul``'s in0 is DRAM-interleaved, so the expert path pays one
1823
+ # sharded-to-interleaved (0.53 us) rather than the norm paying 15.
1824
+ normed = ttnn.sharded_to_interleaved(normed_sharded, ttnn.DRAM_MEMORY_CONFIG)
1825
+ ttnn.deallocate(normed_sharded)
1826
+ moe_partial = moe_decode_multichip(normed, routing, weights.experts, config.local_moe, precision)
1827
+ ttnn.deallocate(normed)
1828
+ ttnn.deallocate(routing)
1829
+ moe_out = all_reduce_decode(moe_partial, ctx, precision)
1830
+ ttnn.deallocate(moe_partial)
1831
+
1832
+ out = ttnn.add(hidden, moe_out)
1833
+ ttnn.deallocate(hidden)
1834
+ ttnn.deallocate(moe_out)
1835
+ return out
1836
+
1837
+
1838
+ # Bytes one 32x32 tile occupies, per dtype. Spelled out because
1839
+ # ``Tensor.element_size()`` raises for the block-float types -- their storage is
1840
+ # a byte (or nibble) per element *plus* a shared exponent per 16-element face
1841
+ # row, i.e. 1024 + 64 for bfloat8_b and 512 + 64 for bfloat4_b -- and the
1842
+ # expert weights, which are the whole point of measuring this, are block-float.
1843
+ _TILE_BYTES = {
1844
+ str(ttnn.float32): 4096,
1845
+ str(ttnn.bfloat16): 2048,
1846
+ str(ttnn.bfloat8_b): 1088,
1847
+ str(ttnn.bfloat4_b): 576,
1848
+ }
1849
+
1850
+
1851
+ def _tensor_bytes(t: ttnn.Tensor) -> int | None:
1852
+ """Device bytes one mesh-sharded tensor occupies **per die**, or ``None``.
1853
+
1854
+ A mesh tensor's shape is already the *local* (per-die) shape, so this is the
1855
+ allocation a dtype change actually moves -- which is the observable
1856
+ ``tests/test_precision_config.py`` asserts on. ``None`` for a dtype with no
1857
+ entry above rather than a wrong number.
1858
+ """
1859
+ tile_bytes = _TILE_BYTES.get(str(t.dtype))
1860
+ if tile_bytes is None:
1861
+ return None
1862
+ shape = [int(v) for v in t.padded_shape]
1863
+ tiles = math.prod(shape[:-2]) * math.ceil(shape[-2] / 32) * math.ceil(shape[-1] / 32)
1864
+ return tiles * tile_bytes
1865
+
1866
+
1867
+ def fallback_audit(
1868
+ weights: MultichipWeights,
1869
+ config: MeshDecoderConfig,
1870
+ batch: int,
1871
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1872
+ ) -> dict:
1873
+ """Every runtime fallback the imported single-chip code can still take.
1874
+
1875
+ Three of stage 02's helpers choose a slower path silently rather than
1876
+ raising, and all three have different inputs under TP/EP than they were
1877
+ tuned against, so "it still passes PCC" would not notice any of them:
1878
+
1879
+ * ``_dram_sharded_usable`` -- falls back to the interleaved ``attention_decode``
1880
+ if the weight dims were not bank-divisible at upload or the batch exceeds
1881
+ 32. Per-die N is now 1280 rather than 5120 and per-die K 1024 rather than
1882
+ 4096, and 1280 = 5x256 is only one factor of two away from failing.
1883
+ * ``_tuned_sparse_matmul_config`` -- silently lowers ``in0_block_w`` to the
1884
+ largest divisor of K in tiles. EP leaves K alone (2048 and 768), so the
1885
+ tuned 16 and 12 must survive; if they did not, this would be scheme A's
1886
+ regression arriving by the back door.
1887
+ * ``_decode_expert_memory_config`` -- moves the expert intermediates from L1
1888
+ to DRAM past a byte budget, which EP shrank 4x.
1889
+
1890
+ Since stage 07 it also reports what the *precision config actually put on
1891
+ the device*: the dtypes read back off the uploaded tensors (not the config's
1892
+ own fields -- those would only prove the dataclass round-trips), the block
1893
+ widths the program configs resolved to, and the fidelities the compute
1894
+ configs carry. That is what ``tests/test_precision_config.py`` asserts
1895
+ against when it constructs at a non-default value.
1896
+
1897
+ Returned as data so a test can assert on it and the work log can quote it.
1898
+ """
1899
+ a = config.local_attention
1900
+ m = config.local_moe
1901
+ k_qkv = int(weights.experts.wqkv_decode.shape[-2]) if weights.experts.wqkv_decode is not None else None
1902
+ n_qkv = int(weights.experts.wqkv_decode.shape[-1]) if weights.experts.wqkv_decode is not None else None
1903
+ k_o = int(weights.experts.wo_decode.shape[-2]) if weights.experts.wo_decode is not None else None
1904
+ n_o = int(weights.experts.wo_decode.shape[-1]) if weights.experts.wo_decode is not None else None
1905
+ gate_up = _tuned_sparse_matmul_config(
1906
+ 1, 2 * m.moe_intermediate_size, m.hidden_size, precision.experts_gate_up_in0_block_w
1907
+ )
1908
+ down = _tuned_sparse_matmul_config(1, m.hidden_size, m.moe_intermediate_size, precision.experts_down_in0_block_w)
1909
+ return {
1910
+ "batch": batch,
1911
+ "dram_sharded_qkv": (k_qkv, n_qkv),
1912
+ "dram_sharded_wo": (k_o, n_o),
1913
+ "dram_sharded_taken": weights.experts.wqkv_decode is not None
1914
+ and weights.experts.wo_decode is not None
1915
+ and batch <= 32,
1916
+ "gate_up_in0_block_w": gate_up.in0_block_w,
1917
+ "down_in0_block_w": down.in0_block_w,
1918
+ "expert_intermediate_buffer": "L1"
1919
+ if _decode_expert_memory_config(batch, m) == ttnn.L1_MEMORY_CONFIG
1920
+ else "DRAM",
1921
+ "local_heads": (a.num_attention_heads, a.num_key_value_heads),
1922
+ "local_experts": m.num_experts,
1923
+ # Stage 04. Not a fallback in the "silently slower path" sense -- a
1924
+ # mismatch here raises rather than degrades -- but it is the same class
1925
+ # of risk, so it is reported as data: if the sharded norm's output shard
1926
+ # ever stops being *exactly* the one the DRAM-sharded qkv projection
1927
+ # wants, TTNN inserts a reshard between them and the layer gets slower
1928
+ # with no error at all. That single equality is what removed stage-03
1929
+ # row 135 from the profile.
1930
+ "norm_shard_cores": _NORM_SHARD_CORES,
1931
+ "norm_shard_feeds_qkv_directly": _norm_shard_config(m.hidden_size) == _width_sharded_l1(m.hidden_size),
1932
+ "decode_ccl_buffers_persistent": True,
1933
+ # -- what the precision config actually produced on device -------------
1934
+ # Read off the uploaded tensors, so these differ from
1935
+ # ``precision.<field>`` if any of the threading above is broken.
1936
+ "device_experts_gate_up_dtype": str(weights.experts.gate_up_proj.dtype),
1937
+ "device_experts_down_dtype": str(weights.experts.down_proj.dtype),
1938
+ "device_attention_qkv_dtype": str(weights.experts.attention.wqkv.dtype),
1939
+ "device_attention_wo_dtype": str(weights.experts.attention.wo.dtype),
1940
+ "device_attention_qkv_decode_dtype": (
1941
+ None if weights.experts.wqkv_decode is None else str(weights.experts.wqkv_decode.dtype)
1942
+ ),
1943
+ "device_router_dtype": str(weights.router.dtype),
1944
+ "device_norm_weight_dtype": str(weights.input_layernorm.dtype),
1945
+ # Bytes one layer's expert weights occupy per die -- the allocation-size
1946
+ # consequence of the two expert dtypes, in a form a sweep can diff.
1947
+ "device_expert_bytes_per_die": (
1948
+ _tensor_bytes(weights.experts.gate_up_proj) + _tensor_bytes(weights.experts.down_proj)
1949
+ ),
1950
+ "expert_math_fidelity": str(precision.experts_fidelity),
1951
+ "attention_math_fidelity": None if precision.attention_fidelity is None else str(precision.attention_fidelity),
1952
+ "router_window_math_fidelity": str(precision.router_window_fidelity),
1953
+ "ccl_dtype": str(precision.effective_ccl_dtype),
1954
+ "activation_dtype": str(precision.activation_dtype),
1955
+ }
1956
+
1957
+
1958
+ __all__ = [
1959
+ "MESH_SHAPE",
1960
+ "NUM_DEVICES",
1961
+ "NUM_LINKS",
1962
+ "NUM_LINKS_DECODE",
1963
+ "TOPOLOGY",
1964
+ "MeshContext",
1965
+ "MeshDecoderConfig",
1966
+ "MultichipWeights",
1967
+ "all_reduce",
1968
+ "all_reduce_decode",
1969
+ "all_reduce_prefill",
1970
+ "build_local_sparsity",
1971
+ "create_mesh_kv_cache",
1972
+ "decoder_layer_decode_multichip",
1973
+ "decoder_layer_prefill_multichip",
1974
+ "fallback_audit",
1975
+ "head_interleaved_wqkv",
1976
+ "mesh_context",
1977
+ "decode_residual_norm",
1978
+ "moe_decode_multichip",
1979
+ "router_forward_multichip",
1980
+ "router_forward_threshold",
1981
+ "upload_multichip_weights",
1982
+ ]
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/optimized_decoder.py ADDED
@@ -0,0 +1,1093 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Optimized TTNN decoder layer for Qwen3-Coder-30B-A3B-Instruct.
5
+
6
+ Same semantics as ``functional_decoder`` -- prefill/decode contract, paged KV
7
+ cache, non-aligned sequence lengths, determinism -- with the measured path
8
+ retuned. Every number below is from real checkpoint weights on a Blackhole
9
+ p300c, 1x1 mesh; the sweeps behind them are in ``doc/optimized_decoder/``.
10
+
11
+ prefill 536.54 -> 69.12 us/token at S=512 (7.76x)
12
+ decode 1.5655 -> 0.5634 ms/token traced at ctx128 (2.78x)
13
+
14
+ Both lines are cells of ``doc/{functional,optimized}_decoder/perf_prefill.csv``
15
+ and ``perf_decode.csv``, which every run of ``tests/test_perf.py`` rewrites;
16
+ the third significant figure moves between runs.
17
+
18
+ What changed, in order of how much it mattered
19
+ ----------------------------------------------
20
+ **1. ``in0_block_w`` (3.0x on prefill).** Stage 01 inherited ``in0_block_w=1``
21
+ from the exemplar's config helper. With K = 2048 (64 tiles) that feeds the
22
+ kernel one tile of the inner dimension at a time, which is what held the expert
23
+ matmuls at ~5.4% of peak FLOPs -- not the core count, and not precision.
24
+
25
+ **2. bfloat4_b expert weights (2.2x on prefill).** Only visible *after* the
26
+ block-width fix: at ``in0_block_w=1`` the kernel is latency-bound, so weight
27
+ dtype cannot matter. The two knobs interact and sweeping either alone finds the
28
+ wrong optimum -- see ``EXPERT_IN0_BLOCK_W_GATE_UP``.
29
+
30
+ **3. DRAM-sharded decode attention projections (1.11x on decode).** Once the
31
+ experts were fast, ``o_proj`` and ``qkv`` were 21% of decode device time and
32
+ the stage-01 "attention is 0.08% of prefill, no action" call went stale.
33
+ ``MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig`` with the weight
34
+ width-sharded across the 8 DRAM banks takes qkv 68.3 -> 46.8 us and wo
35
+ 96.0 -> 41.7 us at the op level, and the whole traced layer 0.6508 -> 0.5863 ms
36
+ at ctx128 measured like for like -- both legs otherwise at the configuration
37
+ shipped at the time, on the same bfloat8_b weights, and both before lever 7
38
+ below, which is why the fast leg reads 0.5863 rather than today's 0.5634. A
39
+ ``0.697 -> 0.587`` pair (1.19x)
40
+ that this file and the docs used to carry is withdrawn; see
41
+ ``attention_decode_optimized`` and ``doc/optimized_decoder/work_log.md`` §5.
42
+
43
+ **4. Packing gate and up (1.09x).** ``_sparse_matmul_config`` parallelises only
44
+ over N, so N tiles cap the usable cores: 768 -> 24, 1536 -> 48, 2048 -> 64.
45
+ Packing doubles gate/up's cores. Measured against a *properly tuned separate*
46
+ candidate it is only 1.09x (2 x 1.476 = 2.952 ms -> 2.699 ms); against the
47
+ untuned stage-01 candidate it looks like 1.66x, but most of that belongs to the
48
+ block-width fix. Re-confirmed at the end of the stage on the whole layer:
49
+ separate 0.735 ms vs packed 0.673 ms traced decode.
50
+
51
+ **5. LoFi on the expert matmuls (1.05x on prefill).** bfp4 weights carry 4
52
+ mantissa bits, so HiFi4's extra passes have nothing to work on. Prefill S=512
53
+ 72.46 -> 69.13 us/token, decode 0.6746 -> 0.6638 ms, and layer PCC is 0.99910 at
54
+ LoFi vs 0.99909 at HiFi4 -- i.e. very slightly *better*.
55
+
56
+ **6. bfloat8_b attention projections (1.02x on decode).** 0.6726 -> 0.6605 ms at
57
+ PCC 0.99906 vs 0.99909. bfloat4_b is 0.6595 ms but drops layer PCC to 0.9928,
58
+ below the 0.995 bar, so it is rejected.
59
+
60
+ **7. The router's two keepdim reductions (1.045x on decode).** The router and
61
+ its routing prep were 111.6 us of decode device time -- 20.9% of it, more than
62
+ either matmul family -- and had no audit finding of their own until the fourth
63
+ review. Two thirds of the removable part was not arithmetic at all: ``ttnn.max``
64
+ and ``ttnn.sum`` each pull a ``FillPad`` behind them on a tensor whose last two
65
+ dims are not tile-aligned, 10.42 and 10.41 us against 1.43 and 1.41 us of actual
66
+ reduction. ``router_forward_optimized`` deletes both -- the max is column 0 of
67
+ the sorted top-k, and the sum moves after the scatter where the reduction length
68
+ is a whole number of tiles -- for 0.5866 -> 0.5615 ms traced at ctx128 (both
69
+ legs in one process), with the routing itself unchanged
70
+ (``test_optimized_router_matches_functional`` asserts identical expert
71
+ selection). Moving the sum past the scatter puts the divide over whole tiles,
72
+ whose row-padding then divides 0 by 0; the divisor is clamped so that padding
73
+ stays exactly zero, which costs +1.6 us and is why ``perf_decode.csv`` reads
74
+ 0.5634 rather than 0.5615. See that function and ``work_log.md`` §7.
75
+
76
+
77
+ Rejected, with measurements
78
+ ---------------------------
79
+ **Per-token sparsity in prefill.** Prefill hands ``sparse_matmul`` a sparsity
80
+ tensor of shape ``[1, 1, group_size, E]`` whose rows are 32-token *tiles*, so an
81
+ expert counts as active if any of the tile's 32 tokens picked it. With 256
82
+ selections landing across 128 slots essentially every expert is hit -- hence
83
+ ``active=128/128``. Making it per-token requires tokens to be *batch* indices
84
+ (``sparse_matmul`` indexes sparsity by batch dims, not by M), i.e. a
85
+ ``[1, T, 1, H]`` layout. Measured rather than assumed: it runs, cuts nnz 16x
86
+ from 4096 to 256, and is **2.1x slower** (14.35 ms vs 6.70 ms), because M
87
+ collapses to 1 and the op pads M to a full 32-row tile. Decode keeps real
88
+ per-token sparsity, which is free there because M is genuinely 1.
89
+
90
+ **1x32 output tiles on the decode sparse matmuls.** The M padding above is the
91
+ single largest remaining inefficiency: at decode M=1 the gate/up matmul writes
92
+ 12 MB and ``down`` writes 16 MB where 0.4/0.5 MB is real, and the reshapes that
93
+ compact it away cost 31 + 33 + 46 us. ``output_tile=ttnn.Tile([1, 32])``
94
+ removes the padding at the source and is 1.07x faster end to end -- but no
95
+ downstream op consumes the result correctly. Measured, in this order:
96
+ ``slice`` rejects it (``slice_device_operation.cpp:165`` hardcodes
97
+ ``TILE_HEIGHT``), ``ttnn.sum`` and ``ttnn.reshape`` raise
98
+ ``MeshBuffer must be large enough``, ``untilize`` returns wrong data without
99
+ erroring, and ``fast_reduce_nc`` returns all zeros. Only eltwise ops read it
100
+ correctly, and they immediately re-pad to 32 rows. Blocked on TTNN support for
101
+ non-32 tile heights outside matmul, not on this model.
102
+
103
+ **Folding the routing weight in before ``down``.** ``down`` is linear, so
104
+ scaling its *input* by the routing probability is equivalent to scaling its
105
+ output, and the input is the compact ``[B, E, I]`` tensor rather than the
106
+ 32x-row-padded ``[B, E, 1, H]`` one. It looked like it should collapse the
107
+ whole tail into one reduce. Measured at ctx128 in one run, before the §7 router
108
+ change (so its shipped-tail leg is the 0.5862 ms configuration of the time, not
109
+ today's 0.5634): shipped tail 0.5862 ms, folded with a compact sum 0.5852 ms,
110
+ folded with ``fast_reduce_nc`` straight off the padded tensor 0.6316 ms. A
111
+ second run of the same three legs read 0.5870 / 0.5853 / 0.6326 and was quoted
112
+ in parallel with this one; this triple is the one whose shipped leg matches the
113
+ ``perf_decode.csv`` ctx128 cell of that day, and it is now the only one quoted
114
+ anywhere. The first two are a tie and the third is 8% *worse* --
115
+ ``fast_reduce_nc`` also promotes ``down``'s tile padding into the logical shape,
116
+ so recovering ``[1,1,B,H]`` needs a permute plus a slice that together cost more
117
+ than the ops they replaced. The shipped tail stays. (The intermediate version of
118
+ this that used a plain reshape instead of the permute was faster still, 0.550 ms
119
+ -- and silently wrong for every user but the first, which is how the permute
120
+ came to be needed. ``test_optimized_decode_batch`` caught it.)
121
+
122
+ **Keeping the expert path rank-6.** The obvious reading of the profile is that
123
+ the rank-changing reshapes are pure overhead. They are not: they compact the 32x
124
+ M padding away, so the elementwise ops that follow touch 192 tiles instead of
125
+ 6144. Staying rank-6 and dropping all three reshapes measured **0.713 ms vs
126
+ 0.673** -- 6% slower.
127
+
128
+ **Everything else tt-perf-report suggested**, each measured on the traced layer
129
+ against the tuned baseline: in0 in L1 on the sparse rows 1.001x, on the
130
+ attention rows 0.998x, ``out_subblock_w=2`` 1.001x, ``=4`` 1.001x. All noise;
131
+ none adopted. HiFi2 on the sparse rows is covered by lever 5 -- LoFi is both
132
+ faster and no less accurate.
133
+ """
134
+
135
+ from __future__ import annotations
136
+
137
+ import math
138
+ from dataclasses import dataclass
139
+
140
+ import torch
141
+
142
+ import ttnn
143
+
144
+ from .functional_decoder import ( # noqa: F401 (re-exported for callers)
145
+ AttentionConfig,
146
+ AttentionWeights,
147
+ DecoderLayerConfig,
148
+ DecoderLayerWeights,
149
+ KVCache,
150
+ MoEConfig,
151
+ _apply_rope,
152
+ _concat_heads_decode,
153
+ _per_head_rms_norm,
154
+ _sparse_matmul_config,
155
+ attention_decode,
156
+ attention_prefill,
157
+ build_expert_sparsity,
158
+ build_rope_cache,
159
+ create_kv_cache,
160
+ upload_layer_weights,
161
+ upload_router_weight,
162
+ )
163
+ from .precision import DEFAULT_PRECISION, PrecisionConfig # noqa: F401 (re-exported)
164
+
165
+ # Tokens per expert-path chunk. Kept at one tile: sparse_matmul folds the group
166
+ # dimension into M, so a larger chunk grows num_blocks_y and can overflow the
167
+ # core grid. Chunking at 32 keeps all blocking in N.
168
+ EXPERT_CHUNK_SIZE = 32
169
+
170
+ # Expert matmul precision and inner-block width.
171
+ #
172
+ # These two knobs INTERACT, and sweeping either alone finds the wrong optimum.
173
+ # At in0_block_w=1 the kernel is latency-bound, so weight dtype makes no
174
+ # measurable difference -- which is exactly the null result stage 02 first
175
+ # recorded, and wrongly concluded from. Widening the block makes the matmul
176
+ # bandwidth-bound, at which point precision becomes the dominant lever.
177
+ #
178
+ # Joint sweep, real weights, M=32, ms (grid 8x6 for gate/up, 8x8 for down):
179
+ #
180
+ # packed gate+up (K=2048, 64 tiles) down (K=768, 24 tiles)
181
+ # blk bf16 bfp8 bfp4 bfp4/LoFi blk bf16 bfp8 bfp4
182
+ # 4 2.697 2.371 2.367 2.366 4 1.489 1.062 1.058
183
+ # 8 2.726 1.669 1.445 1.431 6 1.507 1.002 0.785
184
+ # 16 2.932 1.806 1.259 1.149 8 1.507 0.993 0.712
185
+ # 32 2.910 1.734 1.158 1.153 12 1.513 0.941 0.654
186
+ # 64 3.180 1.789 1.372 1.204 24 1.622 0.982 0.733
187
+ #
188
+ # bf16's best is 2.697 + 1.489 = 4.186 ms; bfp4's is 1.259 + 0.654 = 1.913 ms.
189
+ # Block width must divide K in tiles, and the two matmuls have different K, so
190
+ # the widths are per-role rather than one shared constant.
191
+ #
192
+ # The table above is a matmul microbenchmark; the widths were re-confirmed on
193
+ # the whole layer, where the interaction with fidelity reverses the gate/up
194
+ # choice (prefill S=512 us/token, real weights):
195
+ #
196
+ # blk 8 16 32 64
197
+ # HiFi4 78.27 72.46 69.96 75.56
198
+ # LoFi 78.12 69.13 69.27 70.88
199
+ #
200
+ # 16 at LoFi is the minimum, so 16 stays.
201
+ #
202
+ # **These five names are now aliases, not the source of truth.** The values
203
+ # themselves live in ``precision.PrecisionConfig``, whose defaults are exactly
204
+ # what was written here before stage 07; the names survive because probes under
205
+ # ``doc/`` and the stage-02/04 tests import them, and because a reader arriving
206
+ # at the sweep comments above should find the value they describe next to them.
207
+ # Anything that needs to *vary* the policy must take a ``PrecisionConfig``
208
+ # instead -- these are bound at import and cannot follow a non-default model.
209
+ EXPERT_WEIGHT_DTYPE = DEFAULT_PRECISION.experts_gate_up_dtype
210
+ EXPERT_IN0_BLOCK_W_GATE_UP = DEFAULT_PRECISION.experts_gate_up_in0_block_w # divides 2048/32 = 64
211
+ EXPERT_IN0_BLOCK_W_DOWN = DEFAULT_PRECISION.experts_down_in0_block_w # divides 768/32 = 24
212
+
213
+ # bfp4 weights carry 4 mantissa bits, so HiFi4's extra passes have nothing left
214
+ # to resolve. LoFi is 4.6% faster on prefill and 1.6% on decode at PCC 0.99910
215
+ # vs HiFi4's 0.99909. This also answers tt-perf-report's "HiFi2 may also work"
216
+ # on the sparse rows: HiFi2 measured 69.78 us/token, between the two.
217
+ EXPERT_MATH_FIDELITY = DEFAULT_PRECISION.experts_fidelity
218
+
219
+ # Attention projections. bf16 -> bfloat8_b costs 0.00003 PCC and buys 1.8% of
220
+ # decode; bfloat4_b buys another 0.1% but drops layer PCC to 0.9928, under the
221
+ # 0.995 bar, so it is rejected. q_norm/k_norm stay bf16 -- they are norms, not
222
+ # projections, and weigh 4 KB.
223
+ ATTENTION_WEIGHT_DTYPE = DEFAULT_PRECISION.attention_qkv_dtype
224
+
225
+ # Blackhole p300c has 8 DRAM banks. The DRAM-sharded matmul wants the weight
226
+ # width-sharded one shard per bank, and both the activation and the output
227
+ # width-sharded in L1 over the matching core row.
228
+ _DRAM_BANKS = 8
229
+
230
+
231
+ def _expert_compute_kernel_config(device, precision: PrecisionConfig = DEFAULT_PRECISION):
232
+ """LoFi, and ``fp32_dest_acc_en`` deliberately OFF.
233
+
234
+ ``fp32_dest_acc_en`` looks like the natural next lever but must not be used
235
+ here: it halves the matmul dest from 8 tiles to 4, which corrupts expert
236
+ output on Blackhole (tt-metal #49068, hit on BH-QB-2). It is therefore
237
+ **not** a ``PrecisionConfig`` field -- a sweep must not be able to turn it
238
+ on.
239
+ """
240
+ return ttnn.init_device_compute_kernel_config(
241
+ device.arch(),
242
+ math_fidelity=precision.experts_fidelity,
243
+ math_approx_mode=False,
244
+ fp32_dest_acc_en=False,
245
+ packer_l1_acc=False,
246
+ )
247
+
248
+
249
+ def _attention_compute_kernel_config(device, precision: PrecisionConfig = DEFAULT_PRECISION):
250
+ """``None`` at the default, which is what the projections have always passed.
251
+
252
+ ``attention_fidelity=None`` means "leave the op at its own default", so this
253
+ returns ``None`` and the ``compute_kernel_config=`` argument is a no-op. Any
254
+ named fidelity builds a real config; the remaining flags mirror
255
+ ``_expert_compute_kernel_config``'s, which is the closest measured
256
+ neighbour.
257
+ """
258
+ if precision.attention_fidelity is None:
259
+ return None
260
+ return ttnn.init_device_compute_kernel_config(
261
+ device.arch(),
262
+ math_fidelity=precision.attention_fidelity,
263
+ math_approx_mode=False,
264
+ fp32_dest_acc_en=False,
265
+ packer_l1_acc=False,
266
+ )
267
+
268
+
269
+ def _tuned_sparse_matmul_config(m: int, n: int, k: int, target_blk: int):
270
+ """``_sparse_matmul_config`` with a tuned inner block width.
271
+
272
+ ``k`` is the inner dimension in elements; the block width must divide it in
273
+ tiles, so this falls back to the largest legal divisor at or below the
274
+ target rather than failing.
275
+ """
276
+ k_tiles = max(1, k // 32)
277
+ blk = min(target_blk, k_tiles)
278
+ while blk > 1 and k_tiles % blk:
279
+ blk -= 1
280
+ return _sparse_matmul_config(m, n, in0_block_w=blk)
281
+
282
+
283
+ # Column of ones used to sum the dense routing row (see ``router_forward_optimized``).
284
+ # Cached per (device, length) because it is a constant, and because allocating a
285
+ # tensor inside a trace capture is illegal -- every caller runs the layer eagerly
286
+ # once to compile before capturing, which is what populates this.
287
+ #
288
+ # The key carries ``id(device)`` but the *value* carries the device object itself.
289
+ # ``mesh_device`` is function-scoped in ``conftest.py`` and is closed and deleted
290
+ # after each test, and CPython reuses freed addresses, so a later device could be
291
+ # handed the same id and collide with an entry bound to a destroyed one. Holding
292
+ # the object in the value makes the pin explicit: the address cannot be recycled
293
+ # while the entry lives, so equal ids imply the same live device. The identity
294
+ # check below is then a belt-and-braces assertion, not a hope.
295
+ _ONES_COLUMN: dict[tuple[int, int], tuple[object, ttnn.Tensor]] = {}
296
+
297
+
298
+ def _ones_column(device, n: int) -> ttnn.Tensor:
299
+ key = (id(device), n)
300
+ entry = _ONES_COLUMN.get(key)
301
+ if entry is not None and entry[0] is device:
302
+ return entry[1]
303
+ cached = ttnn.from_torch(
304
+ torch.ones(1, 1, n, 1),
305
+ dtype=ttnn.bfloat16,
306
+ layout=ttnn.TILE_LAYOUT,
307
+ device=device,
308
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
309
+ )
310
+ _ONES_COLUMN[key] = (device, cached)
311
+ return cached
312
+
313
+
314
+ def router_forward_optimized(x: ttnn.Tensor, w_router: ttnn.Tensor, config: MoEConfig) -> ttnn.Tensor:
315
+ """``router_forward`` with both keepdim reductions removed. Same result.
316
+
317
+ The routing maths is unchanged from ``functional_decoder.router_forward``,
318
+ including the part that is load-bearing for correctness: selection happens
319
+ on the **raw fp32 logits** and the softmax is taken over the 8 survivors
320
+ only. A 128-wide bf16 softmax misroutes 83/128 tokens and is not an option
321
+ here; see that function's docstring for the algebra.
322
+
323
+ What changes is how the two reductions are spelled. ``ttnn.max``/``ttnn.sum``
324
+ call ``fill_implicit_tile_padding`` whenever *either* of the last two dims is
325
+ unaligned (``fill_pad.cpp:17-24``); the top-k tensor is 8 wide, and decode's
326
+ is 1 row tall, so each keepdim reduction dragged a ``FillPad`` behind it --
327
+ **10.421 µs and 10.413 µs** in the archived stage-01 decode profile
328
+ (``doc/functional_decoder/ops_perf_decode_paged32.csv`` rows 73 and 77),
329
+ against **1.432 and 1.407 µs** for the reductions themselves (rows 74 and
330
+ 78). Both are avoided rather than tuned:
331
+
332
+ * the **max** is column 0 of the top-k output, which ``sorted=True``
333
+ guarantees is the largest, so one 0.87 µs ``slice`` replaces
334
+ ``FillPad + Reduce``;
335
+ * the **sum** moves *after* the scatter and becomes a matmul against a
336
+ column of ones. Over the dense row the reduction length is
337
+ ``num_experts`` = 128 — a whole number of tiles — so no padding lane can
338
+ enter the sum. That is also why this is preferred to the same matmul over
339
+ the 8-wide tensor, whose K padding would carry whatever ``topk`` left
340
+ behind. Normalising after the scatter is legal because the scatter is a
341
+ permutation of the 8 survivors into a field of exact zeros, so the sum
342
+ over 128 *is* the sum over the 8.
343
+
344
+ **The padding cost of moving the sum.** Dividing after the scatter divides
345
+ over whole tiles, and the tile *row* padding -- rows S..ceil(S/32)*32 -- has
346
+ ``dense`` = 0 and therefore ``total`` = 0 too. Unguarded, ``ttnn.div``
347
+ returns **+inf** there (not NaN; measured at S = 33 and 100, where every one
348
+ of the 31 and 28 padding rows came back +inf), where the functional router,
349
+ which divided before the scatter, returned exact zeros. Nothing observable
350
+ leaked -- ``ttnn.to_torch`` returns the logical shape, the sparsity path
351
+ drops the padding in ``to_layout(ROW_MAJOR)``, and the scale multiply,
352
+ ``rms_norm`` and ``fast_reduce_nc`` all reduce along axes that are either
353
+ tile-aligned or not the padded one -- but it is a hazard the functional path
354
+ did not have, so the divisor is clamped (``ttnn.maximum(total, 1e-30)``).
355
+ Decode is the one case that was already clean: at S = 1 the padding came
356
+ back exact zero unguarded. The clamp is one extra op -- 1.12 µs in the
357
+ decode profile, +1.6 µs on the traced layer, 0.28% -- and
358
+ ``test_optimized_router_padding_is_zero`` stops it being optimized back out.
359
+ See ``work_log.md`` §7 for the padding table and the rejected free version.
360
+
361
+ Measured on the whole traced layer at ctx128, real weights, median of 100:
362
+ **0.5866 -> 0.5615 ms** (``perf_decode.csv`` reads 0.5634, which is that
363
+ configuration in its own run), and the router block **111.6 -> 87.8 µs** of
364
+ decode device time -- rows 68-88 of
365
+ ``doc/functional_decoder/ops_perf_decode_paged32.csv``, which still holds the
366
+ pre-fix block, against the same block in the pre-guard optimized profile.
367
+ With the divisor guard the block is 88.9 µs, rows 69-88 of
368
+ ``doc/optimized_decoder/ops_perf_optimized_decode.csv``. Layer PCC is
369
+ 0.99901 either way (prefill S=128 vs HF: 0.9990057 after, 0.9990050
370
+ before). ``doc/optimized_decoder/work_log.md`` §7
371
+ carries the rejected variants, including ``ttnn.softmax`` over the 8
372
+ survivors — faster still, and wrong: it reduces over the whole 32-wide tile,
373
+ so the weights sum to 0.9736 instead of 1.
374
+ """
375
+ assert config.norm_topk_prob, (
376
+ "router selects on raw logits, which relies on the softmax denominator "
377
+ "cancelling during top-k renormalisation; that only holds when "
378
+ "norm_topk_prob is True"
379
+ )
380
+
381
+ logits = ttnn.linear(x, w_router, dtype=ttnn.float32, memory_config=ttnn.DRAM_MEMORY_CONFIG)
382
+ top_logits, top_indices = ttnn.topk(logits, k=config.num_experts_per_tok, dim=-1, largest=True, sorted=True)
383
+
384
+ # Subtracting the max is for exp() range only; any shared shift cancels in
385
+ # the division. sorted=True means column 0 already is that max.
386
+ top_max = ttnn.slice(top_logits, [0, 0, 0, 0], [1, 1, top_logits.shape[2], 1])
387
+ exp_logits = ttnn.exp(ttnn.sub(top_logits, top_max))
388
+
389
+ zeros = ttnn.typecast(ttnn.zeros_like(logits), ttnn.bfloat16)
390
+ dense = ttnn.scatter(
391
+ zeros,
392
+ dim=-1,
393
+ index=top_indices,
394
+ src=ttnn.typecast(exp_logits, ttnn.bfloat16),
395
+ )
396
+ # Sum over the dense row == sum over the 8 survivors; see the docstring.
397
+ total = ttnn.matmul(dense, _ones_column(x.device(), config.num_experts), dtype=ttnn.bfloat16)
398
+ # Guard the divisor's tile row-padding, which is 0 where ``dense`` is also 0
399
+ # (see the padding note in the docstring). Every real row's denominator is
400
+ # >= 1, because sorted=True makes column 0 of ``exp_logits`` exactly
401
+ # exp(0) = 1, so clamping at 1e-30 cannot touch one: measured bit-identical
402
+ # on the real rows at S = 1, 33, 100. In the padding it turns 0/0 into
403
+ # 0/1e-30 = 0, which is what the functional router returned there.
404
+ guarded = ttnn.maximum(total, 1e-30)
405
+ normalised = ttnn.div(dense, guarded)
406
+ for t in (logits, top_logits, top_indices, top_max, exp_logits, dense, total, guarded):
407
+ ttnn.deallocate(t)
408
+ return normalised
409
+
410
+
411
+ def _bank_row(n: int) -> ttnn.CoreRangeSet:
412
+ return ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(n - 1, 0))})
413
+
414
+
415
+ # The L1 shard height below, and ``per_core_M=1`` in the program config, are
416
+ # both decode's padded M of one 32-row tile. That is what caps this path at
417
+ # batch 32; see ``_dram_sharded_usable``.
418
+ _DRAM_SHARDED_MAX_BATCH = 32
419
+
420
+
421
+ def _width_sharded_l1(width: int) -> ttnn.MemoryConfig:
422
+ """L1 width-sharded over one core per DRAM bank, 32 rows (decode's padded M)."""
423
+ return ttnn.MemoryConfig(
424
+ ttnn.TensorMemoryLayout.WIDTH_SHARDED,
425
+ ttnn.BufferType.L1,
426
+ ttnn.ShardSpec(_bank_row(_DRAM_BANKS), [32, width // _DRAM_BANKS], ttnn.ShardOrientation.ROW_MAJOR),
427
+ )
428
+
429
+
430
+ def _dram_sharded_ok(k: int, n: int) -> bool:
431
+ """Both dims must split evenly into whole tiles across the banks."""
432
+ return k % (_DRAM_BANKS * 32) == 0 and n % (_DRAM_BANKS * 32) == 0
433
+
434
+
435
+ # Decode's two expert intermediates are 97% M padding -- ``sparse_matmul`` pads
436
+ # M=1 back to a 32-row tile -- so they are large in absolute terms:
437
+ #
438
+ # batch * 128 experts * 32 rows * (1536 + 2048) cols * 2 B = batch * 29.4 MB
439
+ #
440
+ # Blackhole offers ~160 MB of allocatable L1 (110 banks x 1.46 MB, as the
441
+ # allocator reports it on this p300c). Holding both in L1 is therefore a
442
+ # batch-1 affordance, not a general one: at batch 8 the allocator rejects
443
+ # ``down``'s 134 MB output outright. Past the budget the pair goes to DRAM,
444
+ # which is what prefill already does at every length.
445
+ #
446
+ # The 40 MB threshold itself is **asserted, not measured**: it is one comfortable
447
+ # step above batch 1's 29.4 MB and below batch 2's 58.8 MB, so it separates the
448
+ # only two cases that exist here, and no sweep was run to find where L1 actually
449
+ # stops paying. What is measured is the pair of endpoints -- B=1 in L1 is the
450
+ # shipped, profiled configuration, and B=8 in L1 does not allocate at all.
451
+ _DECODE_EXPERT_L1_BUDGET_BYTES = 40 * 1024 * 1024
452
+
453
+
454
+ def _decode_expert_memory_config(batch: int, config: MoEConfig) -> ttnn.MemoryConfig:
455
+ """L1 for the intermediates while they fit the budget above, else DRAM."""
456
+ padded_rows = batch * config.num_experts * 32
457
+ nbytes = padded_rows * (2 * config.moe_intermediate_size + config.hidden_size) * 2
458
+ return ttnn.L1_MEMORY_CONFIG if nbytes <= _DECODE_EXPERT_L1_BUDGET_BYTES else ttnn.DRAM_MEMORY_CONFIG
459
+
460
+
461
+ def _dram_sharded_usable(weights: "OptimizedWeights", batch: int) -> bool:
462
+ """Whether decode may take the DRAM-sharded projections at this batch.
463
+
464
+ Two independent conditions:
465
+
466
+ * the weight dims divided evenly across the banks at upload time, so a
467
+ sharded copy exists at all (``_dram_sharded_ok``);
468
+ * the batch still fits decode's single 32-row M tile. ``_width_sharded_l1``
469
+ hardcodes a 32-row shard and ``_dram_sharded_program_config`` sets
470
+ ``per_core_M=1``, so at ``batch > 32`` the activation no longer matches
471
+ its shard spec. Without this check that surfaces as a shard-shape
472
+ mismatch deep in the matmul rather than as a fallback.
473
+ """
474
+ if weights.wqkv_decode is None or weights.wo_decode is None:
475
+ return False
476
+ return batch <= _DRAM_SHARDED_MAX_BATCH
477
+
478
+
479
+ def _dram_sharded_program_config(k: int, n: int):
480
+ return ttnn.MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig(
481
+ in0_block_w=k // _DRAM_BANKS // 32,
482
+ per_core_M=1,
483
+ per_core_N=n // _DRAM_BANKS // 32,
484
+ fused_activation=None,
485
+ )
486
+
487
+
488
+ @dataclass
489
+ class OptimizedWeights:
490
+ """Device weights for the optimized layer.
491
+
492
+ Two things live here that ``upload_layer_weights`` does not provide:
493
+
494
+ * experts with gate and up kept as one weight. ``weight_mapping`` already
495
+ produces the checkpoint's fused ``[E, 2I, H]`` tensor; stage 01 split it
496
+ apart at upload time to mirror the exemplars, so packing is *undoing*
497
+ that split rather than inventing a layout.
498
+ * two copies of the attention projections. Decode uses a DRAM
499
+ width-sharded copy for ``MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig``;
500
+ prefill cannot -- a plain ``ttnn.linear`` on a DRAM-sharded weight throws
501
+ ``Only L1 buffers can have an associated circular buffer`` -- so an
502
+ interleaved copy is kept for it. At bfloat8_b (1.0625 B/elem, because
503
+ each 16-element block carries its own exponent byte) wqkv is 11.14 MB and
504
+ wo 8.91 MB, so the duplicate copy costs **20.05 MB** and the pair 40.11 MB,
505
+ against 24 GB available. An 18.9 MB figure that this file used to carry
506
+ came from rounding bfloat8_b to 1 B/elem and is withdrawn.
507
+ ``config/context_contract.json`` now carries 20.05 MB too; an earlier
508
+ revision of it made the same rounding error and called the pair a wash
509
+ against stage 01's single bf16 copy, which its ``optimized_note`` records.
510
+ """
511
+
512
+ gate_up_proj: ttnn.Tensor # [1, num_experts, hidden, 2 * intermediate]
513
+ down_proj: ttnn.Tensor # [1, num_experts, intermediate, hidden]
514
+ attention: AttentionWeights # interleaved, for prefill
515
+ wqkv_decode: ttnn.Tensor | None # DRAM width-sharded, for decode
516
+ wo_decode: ttnn.Tensor | None
517
+
518
+
519
+ # Stage-01 name, kept so existing callers and docs still resolve.
520
+ PackedExpertWeights = OptimizedWeights
521
+
522
+
523
+ def upload_optimized_weights(
524
+ torch_weights,
525
+ device,
526
+ config: MoEConfig,
527
+ dtype=None,
528
+ *,
529
+ precision: PrecisionConfig = DEFAULT_PRECISION,
530
+ ) -> OptimizedWeights:
531
+ """Upload experts packed along the output dim, plus both attention copies.
532
+
533
+ ``precision`` supplies every weight dtype. ``dtype``, the stage-02 spelling,
534
+ still overrides **both** expert dtypes when given -- several stage-02 tests
535
+ sweep it directly -- but new callers should pass a ``PrecisionConfig``,
536
+ which can also move gate/up and down apart.
537
+
538
+ The expert dtype is a parameter at all because it and ``in0_block_w`` are
539
+ **not** independent: precision only pays once the block width is wide enough
540
+ for the matmul to become bandwidth-bound. Sweeping either alone finds the
541
+ wrong optimum, which is why both live in the same config object.
542
+ """
543
+ fused = torch_weights["experts_gate_up"] # [E, 2I, H], gate first
544
+ gate_up_dtype = dtype if dtype is not None else precision.experts_gate_up_dtype
545
+ down_dtype = dtype if dtype is not None else precision.experts_down_dtype
546
+
547
+ def up(t: torch.Tensor, tensor_dtype, memory_config=ttnn.DRAM_MEMORY_CONFIG) -> ttnn.Tensor:
548
+ return ttnn.from_torch(
549
+ t.contiguous().float(),
550
+ dtype=tensor_dtype,
551
+ layout=ttnn.TILE_LAYOUT,
552
+ device=device,
553
+ memory_config=memory_config,
554
+ )
555
+
556
+ def as_4d(t: torch.Tensor, pad_to_4d: bool = False) -> torch.Tensor:
557
+ if pad_to_4d:
558
+ t = t.reshape(1, 1, 1, -1)
559
+ while t.dim() < 4:
560
+ t = t.unsqueeze(0)
561
+ return t
562
+
563
+ wqkv, wo = as_4d(torch_weights["wqkv"]), as_4d(torch_weights["wo"])
564
+
565
+ def dram_sharded(t: torch.Tensor, tensor_dtype) -> ttnn.Tensor | None:
566
+ k, n = int(t.shape[-2]), int(t.shape[-1])
567
+ if not _dram_sharded_ok(k, n):
568
+ return None
569
+ return up(
570
+ t,
571
+ tensor_dtype,
572
+ ttnn.MemoryConfig(
573
+ ttnn.TensorMemoryLayout.WIDTH_SHARDED,
574
+ ttnn.BufferType.DRAM,
575
+ ttnn.ShardSpec(_bank_row(_DRAM_BANKS), [k, n // _DRAM_BANKS], ttnn.ShardOrientation.ROW_MAJOR),
576
+ ),
577
+ )
578
+
579
+ return OptimizedWeights(
580
+ gate_up_proj=up(fused.transpose(-2, -1).unsqueeze(0), gate_up_dtype),
581
+ down_proj=up(torch_weights["experts_down"].transpose(-2, -1).unsqueeze(0), down_dtype),
582
+ attention=AttentionWeights(
583
+ wqkv=up(wqkv, precision.attention_qkv_dtype),
584
+ wo=up(wo, precision.attention_wo_dtype),
585
+ q_norm=up(as_4d(torch_weights["q_norm"], pad_to_4d=True), precision.norm_weight_dtype),
586
+ k_norm=up(as_4d(torch_weights["k_norm"], pad_to_4d=True), precision.norm_weight_dtype),
587
+ ),
588
+ wqkv_decode=dram_sharded(wqkv, precision.attention_qkv_dtype),
589
+ wo_decode=dram_sharded(wo, precision.attention_wo_dtype),
590
+ )
591
+
592
+
593
+ # Stage-01 name, kept so existing callers still resolve.
594
+ upload_packed_expert_weights = upload_optimized_weights
595
+
596
+
597
+ def attention_decode_optimized(
598
+ x: ttnn.Tensor,
599
+ weights: OptimizedWeights,
600
+ config: AttentionConfig,
601
+ cos_cache: ttnn.Tensor,
602
+ sin_cache: ttnn.Tensor,
603
+ kv_cache: KVCache,
604
+ current_pos: ttnn.Tensor,
605
+ token_index: int,
606
+ sdpa_program_config=None,
607
+ rope=None,
608
+ precision: PrecisionConfig = DEFAULT_PRECISION,
609
+ ) -> ttnn.Tensor:
610
+ """``attention_decode`` with the two projections run DRAM-sharded.
611
+
612
+ ``sdpa_program_config`` is passed straight through to the SDPA-decode op and
613
+ defaults to ``None``, which is what every single-chip caller uses and what
614
+ every number in this file was measured at. It exists for the multichip path:
615
+ at one KV head the op's default core assignment tries to put all 110 worker
616
+ cores on the single head and its tree reduction refuses past 64
617
+ (``sdpa_decode_program_factory.cpp:245``). See
618
+ ``multichip_decoder._sdpa_program_config``.
619
+
620
+ At decode M=1 both projections are pure weight reads, so what limits them is
621
+ how well the read spreads over the DRAM banks. The interleaved layout gave
622
+ 383 GB/s on qkv and 235 GB/s on wo; sharding the weight one shard per bank
623
+ and keeping the activation and output width-sharded in L1 measures
624
+
625
+ qkv (K=2048, N=5120) 68.3 -> 46.8 us 1.46x
626
+ wo (K=4096, N=2048) 96.0 -> 41.7 us 2.30x
627
+
628
+ at the op level, and 0.6508 -> 0.5863 ms (1.11x) on the whole traced layer
629
+ (both legs measured before the §7 router change, hence 0.5863 rather than
630
+ today's 0.5634)
631
+ at ctx128 -- both legs otherwise at the shipped configuration and on the
632
+ same bfloat8_b weights, so only the program config and shard layout differ.
633
+ Core count was swept: 8 (one per bank) beats 16, 32 and 64 on both matmuls,
634
+ because past one shard per bank the extra cores only add mcast traffic. The
635
+ 8 is the tuned quantity; the profiler reports ``CORE COUNT`` 80 for these
636
+ rows and ``tt-perf-report`` prints 12, and neither of those was chosen.
637
+
638
+ In the archived profiles the two projections go 57.06 -> 27.33 us (qkv) and
639
+ 72.80 -> 21.91 us (wo). Only ``wo`` was SLOW-classified interleaved; qkv was
640
+ already DRAM-classified, so its gain is duration, not a change of class.
641
+
642
+ Batch is capped at 32 here -- see ``_dram_sharded_usable`` -- which is where
643
+ ``nlp_create_qkv_heads_decode`` caps it anyway, on either path.
644
+
645
+ Everything between the two projections is identical to ``attention_decode``,
646
+ including the Blackhole staging workarounds, so the two stay diffable.
647
+ """
648
+ if not _dram_sharded_usable(weights, int(x.shape[-2])):
649
+ return attention_decode(x, weights.attention, config, cos_cache, sin_cache, kv_cache, current_pos, token_index)
650
+
651
+ k_cache, v_cache, page_table = kv_cache.k, kv_cache.v, kv_cache.page_table
652
+ k_qkv, n_qkv = int(weights.wqkv_decode.shape[-2]), int(weights.wqkv_decode.shape[-1])
653
+ k_o, n_o = int(weights.wo_decode.shape[-2]), int(weights.wo_decode.shape[-1])
654
+
655
+ attn_compute_config = _attention_compute_kernel_config(x.device(), precision)
656
+ x_sharded = ttnn.to_memory_config(x, _width_sharded_l1(k_qkv))
657
+ xqkv = ttnn.linear(
658
+ x_sharded,
659
+ weights.wqkv_decode,
660
+ program_config=_dram_sharded_program_config(k_qkv, n_qkv),
661
+ memory_config=_width_sharded_l1(n_qkv),
662
+ dtype=precision.activation_dtype,
663
+ compute_kernel_config=attn_compute_config,
664
+ )
665
+ ttnn.deallocate(x_sharded)
666
+
667
+ # nlp_create_qkv_heads_decode wants interleaved L1. (It also must not be
668
+ # handed a DRAM tensor at all on Blackhole -- tt-metal #16667 zeroes
669
+ # odd-indexed Q rows via a NoC DRAM-read alignment violation.)
670
+ xqkv = ttnn.to_memory_config(xqkv, ttnn.L1_MEMORY_CONFIG)
671
+ q, k, v = ttnn.experimental.nlp_create_qkv_heads_decode(
672
+ xqkv,
673
+ num_heads=config.num_attention_heads,
674
+ num_kv_heads=config.num_key_value_heads,
675
+ memory_config=ttnn.L1_HEIGHT_SHARDED_MEMORY_CONFIG,
676
+ )
677
+ ttnn.deallocate(xqkv)
678
+
679
+ # rms_norm wants interleaved DRAM while paged_update_cache requires a
680
+ # *sharded* update tensor, so remember the split's layout and restore it.
681
+ kv_sharded_mem = k.memory_config()
682
+ q = _per_head_rms_norm(
683
+ ttnn.to_memory_config(q, ttnn.DRAM_MEMORY_CONFIG), weights.attention.q_norm, config.rms_norm_eps
684
+ )
685
+ k = _per_head_rms_norm(
686
+ ttnn.to_memory_config(k, ttnn.DRAM_MEMORY_CONFIG), weights.attention.k_norm, config.rms_norm_eps
687
+ )
688
+ # ``rope`` defaults to ``None`` and therefore to ``_apply_rope``, which is
689
+ # what every caller uses -- including the shipped multichip decode path --
690
+ # and what every number in this file was measured at. It is a seam, not a
691
+ # switch: stage 04 used it to build and measure a Meta-ordered
692
+ # ``rotary_embedding_llama`` alternative (3.05x faster standalone and
693
+ # bit-identical) without disturbing the 1x1 baseline the multichip documents
694
+ # compare against. That alternative is **rejected** -- the KV cache carries
695
+ # the rotary's channel convention and prefill writes HF-ordered keys, so it
696
+ # is not a decode-local change. See ``multichip_decoder._meta_rope`` and
697
+ # ``doc/optimized_multichip_decoder/README.md`` limitation 4.
698
+ _rope = _apply_rope if rope is None else rope
699
+ q = _rope(q, cos_cache, sin_cache, token_index)
700
+ k = ttnn.to_memory_config(_rope(k, cos_cache, sin_cache, token_index), kv_sharded_mem)
701
+
702
+ # Deliberately NOT cast to the cache dtype, unlike the prefill fill writers.
703
+ # ``paged_update_cache`` requires a FLOAT32/BFLOAT16 update and converts into
704
+ # the cache itself (measured: bfp8 cache + bf16 update round-trips at PCC
705
+ # 0.999969, bfp8 update is rejected at
706
+ # ``paged_update_cache_device_operation.cpp:296``). See
707
+ # ``functional_decoder.match_cache_dtype`` for the full table.
708
+ ttnn.experimental.paged_update_cache(k_cache, k, update_idxs_tensor=current_pos, page_table=page_table)
709
+ ttnn.experimental.paged_update_cache(v_cache, v, update_idxs_tensor=current_pos, page_table=page_table)
710
+ ttnn.deallocate(k)
711
+ ttnn.deallocate(v)
712
+
713
+ if kv_cache.is_paged:
714
+ attn = ttnn.transformer.paged_scaled_dot_product_attention_decode(
715
+ q,
716
+ k_cache,
717
+ v_cache,
718
+ page_table_tensor=page_table,
719
+ cur_pos_tensor=current_pos,
720
+ scale=config.head_dim**-0.5,
721
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
722
+ program_config=sdpa_program_config,
723
+ )
724
+ else:
725
+ attn = ttnn.transformer.scaled_dot_product_attention_decode(
726
+ q,
727
+ k_cache,
728
+ v_cache,
729
+ cur_pos_tensor=current_pos,
730
+ scale=config.head_dim**-0.5,
731
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
732
+ program_config=sdpa_program_config,
733
+ )
734
+ ttnn.deallocate(q)
735
+
736
+ attn = ttnn.to_memory_config(_concat_heads_decode(attn, config), _width_sharded_l1(k_o))
737
+ out = ttnn.linear(
738
+ attn,
739
+ weights.wo_decode,
740
+ program_config=_dram_sharded_program_config(k_o, n_o),
741
+ memory_config=_width_sharded_l1(n_o),
742
+ dtype=precision.activation_dtype,
743
+ compute_kernel_config=attn_compute_config,
744
+ )
745
+ ttnn.deallocate(attn)
746
+ return ttnn.to_memory_config(out, ttnn.DRAM_MEMORY_CONFIG)
747
+
748
+
749
+ def _experts_chunk_packed(
750
+ hidden: ttnn.Tensor,
751
+ routing: ttnn.Tensor,
752
+ weights: OptimizedWeights,
753
+ config: MoEConfig,
754
+ sparsity_base: ttnn.Tensor,
755
+ precision: PrecisionConfig = DEFAULT_PRECISION,
756
+ ) -> ttnn.Tensor:
757
+ """One 32-token chunk with gate and up computed in a single matmul.
758
+
759
+ The win is core occupancy, not the saved kernel launch.
760
+ ``_sparse_matmul_config`` parallelises only over N, so the usable core count
761
+ is capped by the number of N tiles:
762
+
763
+ gate or up alone N = 768 -> 24 tiles -> 24 cores
764
+ gate+up packed N = 1536 -> 48 tiles -> 48 cores
765
+ down N = 2048 -> 64 tiles -> 64 cores
766
+
767
+ which is also why the stage-01 profile showed down running at 127 GB/s
768
+ while gate/up sat at 64 GB/s. Worth 1.09x against a *tuned* separate
769
+ candidate (2 x 1.476 = 2.952 ms -> 2.699 ms); the larger figure it shows
770
+ against an untuned one belongs to the block-width fix, not to packing.
771
+ """
772
+ chunk_len = hidden.shape[2]
773
+ n_experts = config.num_experts
774
+ hidden_size = config.hidden_size
775
+ inter = config.moe_intermediate_size
776
+ group_size = chunk_len // EXPERT_CHUNK_SIZE
777
+
778
+ device = hidden.device()
779
+ compute_config = _expert_compute_kernel_config(device, precision)
780
+ output_tile = ttnn.Tile([32, 32])
781
+ # PREFILL blocking. Separate from the decode fields (see precision.py): these
782
+ # matmuls run at M = EXPERT_CHUNK_SIZE here and M = 1 in decode, and the
783
+ # shipped values were tuned only for the latter. Defaults are identical, so
784
+ # this is a no-op until the prefill fields are changed.
785
+ gate_up_config = _tuned_sparse_matmul_config(
786
+ EXPERT_CHUNK_SIZE, 2 * inter, hidden_size, precision.prefill_experts_gate_up_in0_block_w
787
+ )
788
+ down_config = _tuned_sparse_matmul_config(
789
+ EXPERT_CHUNK_SIZE, hidden_size, inter, precision.prefill_experts_down_in0_block_w
790
+ )
791
+
792
+ hidden_grouped = ttnn.reshape(hidden, (1, group_size, EXPERT_CHUNK_SIZE, hidden_size))
793
+ sparsity = ttnn.repeat(sparsity_base, (1, 1, group_size, 1))
794
+ nnz = n_experts * group_size
795
+
796
+ fused = ttnn.sparse_matmul(
797
+ hidden_grouped,
798
+ weights.gate_up_proj,
799
+ sparsity=sparsity,
800
+ nnz=nnz,
801
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
802
+ output_tile=output_tile,
803
+ program_config=gate_up_config,
804
+ compute_kernel_config=compute_config,
805
+ dtype=precision.activation_dtype,
806
+ )
807
+ ttnn.deallocate(hidden_grouped)
808
+ packed_width = fused.shape[-1]
809
+ fused = ttnn.reshape(ttnn.transpose(fused, 1, 3), (1, n_experts, chunk_len, packed_width))
810
+
811
+ # gate is the first half -- matches Qwen3MoeExperts.forward's chunk(2, dim=-1)
812
+ half = packed_width // 2
813
+ gate = ttnn.slice(fused, [0, 0, 0, 0], [1, n_experts, chunk_len, half])
814
+ up = ttnn.slice(fused, [0, 0, 0, half], [1, n_experts, chunk_len, packed_width])
815
+ ttnn.deallocate(fused)
816
+
817
+ down_input = ttnn.reshape(ttnn.mul(ttnn.silu(gate), up), (1, n_experts, chunk_len, half))
818
+ ttnn.deallocate(gate)
819
+ ttnn.deallocate(up)
820
+
821
+ down = ttnn.sparse_matmul(
822
+ down_input,
823
+ weights.down_proj,
824
+ sparsity=sparsity_base,
825
+ nnz=n_experts,
826
+ memory_config=ttnn.DRAM_MEMORY_CONFIG,
827
+ output_tile=output_tile,
828
+ program_config=down_config,
829
+ is_input_a_sparse=True,
830
+ compute_kernel_config=compute_config,
831
+ dtype=precision.activation_dtype,
832
+ )
833
+ ttnn.deallocate(down_input)
834
+
835
+ states = ttnn.reshape(down, (1, n_experts, chunk_len, hidden_size))
836
+ states = ttnn.mul(states, ttnn.permute(routing, (0, 3, 2, 1)))
837
+ states = ttnn.unsqueeze_to_4D(ttnn.experimental.fast_reduce_nc(states, dims=[1]))
838
+ return ttnn.reshape(states, (1, 1, chunk_len, hidden_size))
839
+
840
+
841
+ def moe_prefill_optimized(
842
+ x: ttnn.Tensor,
843
+ routing: ttnn.Tensor,
844
+ weights: OptimizedWeights,
845
+ config: MoEConfig,
846
+ sparsity_base: ttnn.Tensor,
847
+ precision: PrecisionConfig = DEFAULT_PRECISION,
848
+ ) -> ttnn.Tensor:
849
+ """Expert pass over a sequence. ``x`` ``[1, 1, S, H]``, any S.
850
+
851
+ Non-aligned lengths are zero-padded to a chunk boundary and sliced back;
852
+ padded rows carry an all-zero routing vector so they contribute nothing.
853
+ """
854
+ seq_len = x.shape[2]
855
+ padded_len = math.ceil(seq_len / EXPERT_CHUNK_SIZE) * EXPERT_CHUNK_SIZE
856
+
857
+ if padded_len != seq_len:
858
+ pad = [(0, 0), (0, 0), (0, padded_len - seq_len), (0, 0)]
859
+ x = ttnn.pad(x, pad, value=0.0)
860
+ routing = ttnn.pad(routing, pad, value=0.0)
861
+
862
+ outputs = []
863
+ for start in range(0, padded_len, EXPERT_CHUNK_SIZE):
864
+ end = start + EXPERT_CHUNK_SIZE
865
+ outputs.append(
866
+ _experts_chunk_packed(
867
+ ttnn.slice(x, [0, 0, start, 0], [1, 1, end, config.hidden_size]),
868
+ ttnn.slice(routing, [0, 0, start, 0], [1, 1, end, config.num_experts]),
869
+ weights,
870
+ config,
871
+ sparsity_base,
872
+ precision,
873
+ )
874
+ )
875
+ out = outputs[0] if len(outputs) == 1 else ttnn.concat(outputs, dim=2)
876
+ if padded_len != seq_len:
877
+ out = ttnn.slice(out, [0, 0, 0, 0], [1, 1, seq_len, config.hidden_size])
878
+ return out
879
+
880
+
881
+ def moe_decode_optimized(
882
+ x: ttnn.Tensor,
883
+ routing: ttnn.Tensor,
884
+ weights: OptimizedWeights,
885
+ config: MoEConfig,
886
+ precision: PrecisionConfig = DEFAULT_PRECISION,
887
+ ) -> ttnn.Tensor:
888
+ """Decode MoE with gate/up packed and per-token sparsity. ``x`` ``[1, 1, batch, H]``.
889
+
890
+ Tokens are carried as *batch* indices (``[1, B, 1, H]``) rather than along M.
891
+ ``sparse_matmul`` indexes its sparsity tensor by batch dims, so this is what
892
+ makes the pattern per-token -- and unlike prefill it costs nothing here,
893
+ because decode's M is genuinely 1.
894
+
895
+ The two matmuls need different sparsity flags, which is not obvious and is
896
+ what previously limited this path to a single user. From
897
+ ``sparse_matmul_device_operation.cpp``::
898
+
899
+ a_sparse && b_sparse -> batch_length = batch_length_B
900
+ a_sparse -> batch_length = batch_length_A
901
+ neither -> batch_length = batch_length_A * batch_length_B
902
+
903
+ ``is_input_b_sparse`` defaults to true. gate/up takes the third branch and
904
+ gets ``B * E``, which matches a ``[1, 1, B, E]`` sparsity tensor. The down
905
+ projection has a sparse activation, so it would take the *first* branch and
906
+ get just ``E`` -- ignoring the batch entirely and rejecting any B > 1.
907
+ Passing ``is_input_b_sparse=False`` selects the second branch instead, so
908
+ down sees ``batch_length_A = B * E`` and matches.
909
+
910
+ The rank juggling is deliberate. ``sparse_matmul`` returns
911
+ ``[1, 1, B, E, 1, N]`` and pads that M=1 to a full 32-row tile, so the
912
+ result is 97% padding; reshaping to a compact ``[B, E, N]`` before the
913
+ elementwise work makes those ops touch 192 tiles instead of 6144. Dropping
914
+ the reshapes and staying rank-6 measured 6% *slower*.
915
+
916
+ Applying the routing weight to ``down``'s *input* instead -- equivalent,
917
+ since ``down`` is linear, and the input is the compact tensor -- was
918
+ measured and is a tie (0.5852 vs 0.5862 ms, one run, at the pre-§7
919
+ configuration); see the module docstring for why the variant that looked
920
+ much better than that was not.
921
+ """
922
+ batch = x.shape[2]
923
+ n_experts = config.num_experts
924
+ hidden_size = config.hidden_size
925
+ inter = config.moe_intermediate_size
926
+ nnz = config.num_experts_per_tok * batch
927
+
928
+ sparsity = ttnn.to_layout(routing, ttnn.ROW_MAJOR_LAYOUT)
929
+ expert_memory_config = _decode_expert_memory_config(batch, config)
930
+ output_tile = ttnn.Tile([32, 32])
931
+ compute_config = _expert_compute_kernel_config(x.device(), precision)
932
+ gate_up_config = _tuned_sparse_matmul_config(1, 2 * inter, hidden_size, precision.experts_gate_up_in0_block_w)
933
+ down_config = _tuned_sparse_matmul_config(1, hidden_size, inter, precision.experts_down_in0_block_w)
934
+
935
+ x_batched = ttnn.reshape(x, (1, batch, 1, hidden_size))
936
+ fused = ttnn.sparse_matmul(
937
+ x_batched,
938
+ weights.gate_up_proj,
939
+ sparsity=sparsity,
940
+ nnz=nnz,
941
+ memory_config=expert_memory_config,
942
+ output_tile=output_tile,
943
+ program_config=gate_up_config,
944
+ compute_kernel_config=compute_config,
945
+ dtype=precision.activation_dtype,
946
+ )
947
+ packed_width = fused.shape[-1]
948
+ fused = ttnn.reshape(fused, (batch, n_experts, packed_width))
949
+
950
+ # gate is the first half -- matches Qwen3MoeExperts.forward's chunk(2, dim=-1)
951
+ half = packed_width // 2
952
+ gate = ttnn.slice(fused, [0, 0, 0], [batch, n_experts, half])
953
+ up = ttnn.slice(fused, [0, 0, half], [batch, n_experts, packed_width])
954
+ ttnn.deallocate(fused)
955
+
956
+ down_input = ttnn.reshape(ttnn.mul(ttnn.silu(gate), up), (batch, n_experts, 1, half))
957
+ ttnn.deallocate(gate)
958
+ ttnn.deallocate(up)
959
+
960
+ down = ttnn.sparse_matmul(
961
+ down_input,
962
+ weights.down_proj,
963
+ sparsity=sparsity,
964
+ nnz=nnz,
965
+ memory_config=expert_memory_config,
966
+ output_tile=output_tile,
967
+ program_config=down_config,
968
+ is_input_a_sparse=True,
969
+ is_input_b_sparse=False, # see docstring: selects batch_length_A = B * E
970
+ compute_kernel_config=compute_config,
971
+ dtype=precision.activation_dtype,
972
+ )
973
+ ttnn.deallocate(down_input)
974
+
975
+ states = ttnn.reshape(down, (batch, n_experts, hidden_size))
976
+ states = ttnn.mul(states, ttnn.reshape(routing, (batch, n_experts, 1)))
977
+ states = ttnn.unsqueeze_to_4D(ttnn.sum(states, dim=1))
978
+ return ttnn.reshape(states, (1, 1, batch, hidden_size), (1, 1, max(32, batch), hidden_size))
979
+
980
+
981
+ def decoder_layer_prefill_optimized(
982
+ x: ttnn.Tensor,
983
+ weights: DecoderLayerWeights,
984
+ config: DecoderLayerConfig,
985
+ cos_cache: ttnn.Tensor,
986
+ sin_cache: ttnn.Tensor,
987
+ sparsity: ttnn.Tensor,
988
+ packed_experts: OptimizedWeights,
989
+ kv_cache: KVCache | None = None,
990
+ user_id: int = 0,
991
+ precision: PrecisionConfig = DEFAULT_PRECISION,
992
+ ) -> ttnn.Tensor:
993
+ """Optimized prefill. Same contract as ``decoder_layer_prefill``.
994
+
995
+ Attention runs the interleaved bfloat8_b copy: the DRAM-sharded program
996
+ config is decode-only (``per_core_M=1``), and a plain ``ttnn.linear`` cannot
997
+ read a DRAM-sharded weight at all.
998
+ """
999
+ eps = config.rms_norm_eps
1000
+
1001
+ normed = ttnn.rms_norm(x, weight=weights.input_layernorm, epsilon=eps)
1002
+ attn_out = attention_prefill(
1003
+ normed, packed_experts.attention, config.attention, cos_cache, sin_cache, kv_cache, user_id
1004
+ )
1005
+ ttnn.deallocate(normed)
1006
+ hidden = ttnn.add(x, attn_out)
1007
+ ttnn.deallocate(attn_out)
1008
+
1009
+ normed = ttnn.rms_norm(hidden, weight=weights.post_attention_layernorm, epsilon=eps)
1010
+ routing = router_forward_optimized(normed, weights.router, config.moe)
1011
+ moe_out = moe_prefill_optimized(normed, routing, packed_experts, config.moe, sparsity, precision)
1012
+ ttnn.deallocate(normed)
1013
+ ttnn.deallocate(routing)
1014
+
1015
+ out = ttnn.add(hidden, moe_out)
1016
+ ttnn.deallocate(hidden)
1017
+ ttnn.deallocate(moe_out)
1018
+ return out
1019
+
1020
+
1021
+ def decoder_layer_decode_optimized(
1022
+ x: ttnn.Tensor,
1023
+ weights: DecoderLayerWeights,
1024
+ config: DecoderLayerConfig,
1025
+ cos_cache: ttnn.Tensor,
1026
+ sin_cache: ttnn.Tensor,
1027
+ kv_cache: KVCache,
1028
+ current_pos: ttnn.Tensor,
1029
+ token_index: int,
1030
+ *,
1031
+ packed_experts: OptimizedWeights,
1032
+ precision: PrecisionConfig = DEFAULT_PRECISION,
1033
+ ) -> ttnn.Tensor:
1034
+ """Optimized decode. Decode already used per-token sparsity in stage 01."""
1035
+ eps = config.rms_norm_eps
1036
+
1037
+ normed = ttnn.rms_norm(x, weight=weights.input_layernorm, epsilon=eps)
1038
+ attn_out = attention_decode_optimized(
1039
+ normed,
1040
+ packed_experts,
1041
+ config.attention,
1042
+ cos_cache,
1043
+ sin_cache,
1044
+ kv_cache,
1045
+ current_pos,
1046
+ token_index,
1047
+ precision=precision,
1048
+ )
1049
+ ttnn.deallocate(normed)
1050
+ hidden = ttnn.add(x, attn_out)
1051
+ ttnn.deallocate(attn_out)
1052
+
1053
+ normed = ttnn.rms_norm(hidden, weight=weights.post_attention_layernorm, epsilon=eps)
1054
+ routing = router_forward_optimized(normed, weights.router, config.moe)
1055
+ moe_out = moe_decode_optimized(normed, routing, packed_experts, config.moe, precision)
1056
+ ttnn.deallocate(normed)
1057
+ ttnn.deallocate(routing)
1058
+
1059
+ out = ttnn.add(hidden, moe_out)
1060
+ ttnn.deallocate(hidden)
1061
+ ttnn.deallocate(moe_out)
1062
+ return out
1063
+
1064
+
1065
+ __all__ = [
1066
+ "EXPERT_CHUNK_SIZE",
1067
+ "router_forward_optimized",
1068
+ "EXPERT_WEIGHT_DTYPE",
1069
+ "EXPERT_MATH_FIDELITY",
1070
+ "ATTENTION_WEIGHT_DTYPE",
1071
+ "PrecisionConfig",
1072
+ "DEFAULT_PRECISION",
1073
+ "moe_prefill_optimized",
1074
+ "moe_decode_optimized",
1075
+ "attention_decode_optimized",
1076
+ "OptimizedWeights",
1077
+ "PackedExpertWeights",
1078
+ "upload_optimized_weights",
1079
+ "upload_packed_expert_weights",
1080
+ "decoder_layer_prefill_optimized",
1081
+ "decoder_layer_decode_optimized",
1082
+ "build_rope_cache",
1083
+ "build_expert_sparsity",
1084
+ "create_kv_cache",
1085
+ "upload_layer_weights",
1086
+ "upload_router_weight",
1087
+ "DecoderLayerConfig",
1088
+ "DecoderLayerWeights",
1089
+ "KVCache",
1090
+ "MoEConfig",
1091
+ "AttentionConfig",
1092
+ "AttentionWeights",
1093
+ ]
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/precision.py ADDED
@@ -0,0 +1,378 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """The model's precision policy, as one value that construction consumes.
5
+
6
+ Before this module the policy was module-level constants in
7
+ ``optimized_decoder.py`` (``EXPERT_WEIGHT_DTYPE``, ``EXPERT_MATH_FIDELITY``,
8
+ ``ATTENTION_WEIGHT_DTYPE``, the two ``EXPERT_IN0_BLOCK_W_*``) plus a handful of
9
+ literal ``dtype=ttnn.bfloat16`` arguments scattered through ``model.py`` and
10
+ ``multichip_decoder.py``. That is enough to *ship* a policy but not enough to
11
+ *sweep* one: varying any of it meant editing source between runs, and a JSON
12
+ file written next to unchanged source would be a claim rather than a
13
+ configuration.
14
+
15
+ ``PrecisionConfig`` is the value those constants became. It is frozen, it
16
+ round-trips through JSON losslessly (:meth:`to_json` / :meth:`from_json`), and
17
+ every field below is read at model-construction or forward time by code that
18
+ would behave differently if the field changed. The module constants still exist
19
+ -- probes and stage-02/04 tests import them by name -- but they are now
20
+ *derived* from ``DEFAULT_PRECISION`` rather than being the source of truth, so
21
+ there is exactly one place a shipped value is written down.
22
+
23
+ **The default is the shipped policy, and stage 07 moved it.** When this module
24
+ was introduced its default reproduced stages 02-06 exactly. It no longer does:
25
+ the stage-07 sweep selected new values for the two expert block widths
26
+ (``experts_gate_up_in0_block_w`` 16 -> 64, ``experts_down_in0_block_w`` 12 ->
27
+ 24) and ``DEFAULT_PRECISION`` carries them, because the goal requires the
28
+ selection to be what the default construction path consumes. **Every other
29
+ field still reads back exactly what stages 02-06 measured**, and the two that
30
+ moved are scheduling choices rather than numerical ones -- the graph is the same
31
+ graph and the tokens are the same tokens; only the expert matmuls' inner
32
+ blocking differs. See the block-width fields below and
33
+ ``doc/datatype_sweep/README.md``.
34
+
35
+ Three fields are ``None`` by default and that is deliberate rather than an
36
+ omission:
37
+
38
+ ``attention_fidelity``
39
+ The attention projections pass no ``compute_kernel_config`` today, so they
40
+ take the op default. ``None`` reproduces that exactly; any other value
41
+ builds a config and passes it. Encoding "op default" as an explicit
42
+ ``MathFidelity`` would be a guess about what the op picks.
43
+ ``ccl_dtype``
44
+ The collectives run at whatever dtype the activation arrives in. ``None``
45
+ means "inherit", which is today's behaviour and costs no ops; a named dtype
46
+ casts into and out of the collective.
47
+ ``experts_gate_up_fidelity`` has no ``None`` counterpart -- the expert matmuls
48
+ have always passed an explicit config -- so it is a plain value.
49
+
50
+ The block widths live here too. They are not dtypes, but they were tuned
51
+ *against* the dtype, so a sweep that varies expert dtype and cannot vary the
52
+ block width alongside it would be measuring a mis-tuned point.
53
+
54
+ That mattered more than expected. Stage 07's sweep found the two block widths
55
+ to be **the only fields worth moving in the entire twenty-field config**: taking
56
+ each to its full-K ceiling (gate_up 16 -> 64, down 12 -> 24) bought +2.83%
57
+ traced decode at bit-identical accuracy, while every dtype lever the sweep tried
58
+ either regressed, landed inside the run-to-run band, or hit a TTNN blocker. The
59
+ old values were inherited from single-chip stage-02 tuning and expert
60
+ parallelism had since cut per-die N four-fold, changing which blocking the
61
+ matmul wants. See ``doc/datatype_sweep/README.md``.
62
+ """
63
+
64
+ from __future__ import annotations
65
+
66
+ import json
67
+ from dataclasses import asdict, dataclass, fields, replace
68
+ from pathlib import Path
69
+
70
+ import ttnn
71
+
72
+ __all__ = [
73
+ "PrecisionConfig",
74
+ "DEFAULT_PRECISION",
75
+ "dtype_from_name",
76
+ "dtype_to_name",
77
+ "fidelity_from_name",
78
+ "fidelity_to_name",
79
+ ]
80
+
81
+
82
+ # --- name <-> object tables ---------------------------------------------------
83
+ #
84
+ # Spelled out rather than derived from ``str(dtype)`` so the JSON is a stable
85
+ # contract: a rename inside ttnn's binding would silently invalidate every
86
+ # archived config if the names were scraped.
87
+
88
+ _DTYPES: dict[str, "ttnn.DataType"] = {
89
+ "bfloat16": ttnn.bfloat16,
90
+ "bfloat8_b": ttnn.bfloat8_b,
91
+ "bfloat4_b": ttnn.bfloat4_b,
92
+ "float32": ttnn.float32,
93
+ "uint8": ttnn.uint8,
94
+ "uint16": ttnn.uint16,
95
+ "int32": ttnn.int32,
96
+ "uint32": ttnn.uint32,
97
+ }
98
+ _DTYPE_NAMES: dict["ttnn.DataType", str] = {v: k for k, v in _DTYPES.items()}
99
+
100
+ _FIDELITIES: dict[str, "ttnn.MathFidelity"] = {
101
+ "LoFi": ttnn.MathFidelity.LoFi,
102
+ "HiFi2": ttnn.MathFidelity.HiFi2,
103
+ "HiFi3": ttnn.MathFidelity.HiFi3,
104
+ "HiFi4": ttnn.MathFidelity.HiFi4,
105
+ }
106
+ _FIDELITY_NAMES: dict["ttnn.MathFidelity", str] = {v: k for k, v in _FIDELITIES.items()}
107
+
108
+
109
+ def dtype_from_name(name):
110
+ """``"bfloat4_b"`` -> ``ttnn.bfloat4_b``. ``None`` and ttnn dtypes pass through."""
111
+ if name is None or isinstance(name, ttnn.DataType):
112
+ return name
113
+ try:
114
+ return _DTYPES[str(name)]
115
+ except KeyError:
116
+ raise ValueError(f"unknown dtype {name!r}; known: {sorted(_DTYPES)}") from None
117
+
118
+
119
+ def dtype_to_name(dtype):
120
+ if dtype is None:
121
+ return None
122
+ try:
123
+ return _DTYPE_NAMES[dtype]
124
+ except KeyError:
125
+ raise ValueError(f"dtype {dtype!r} has no serialised name; add it to precision._DTYPES") from None
126
+
127
+
128
+ def fidelity_from_name(name):
129
+ """``"LoFi"`` -> ``ttnn.MathFidelity.LoFi``. ``None`` and fidelities pass through."""
130
+ if name is None or isinstance(name, ttnn.MathFidelity):
131
+ return name
132
+ try:
133
+ return _FIDELITIES[str(name)]
134
+ except KeyError:
135
+ raise ValueError(f"unknown math fidelity {name!r}; known: {sorted(_FIDELITIES)}") from None
136
+
137
+
138
+ def fidelity_to_name(fidelity):
139
+ if fidelity is None:
140
+ return None
141
+ try:
142
+ return _FIDELITY_NAMES[fidelity]
143
+ except KeyError:
144
+ raise ValueError(f"fidelity {fidelity!r} has no serialised name") from None
145
+
146
+
147
+ # Which coercion each field takes on the way in from JSON. Every field of
148
+ # ``PrecisionConfig`` must appear here or in ``_INT_FIELDS``; ``__post_init__``
149
+ # asserts that, so a field added without a serialisation rule fails at import
150
+ # rather than producing a JSON file that silently drops it.
151
+ _DTYPE_FIELDS = frozenset(
152
+ {
153
+ "experts_gate_up_dtype",
154
+ "experts_down_dtype",
155
+ "attention_qkv_dtype",
156
+ "attention_wo_dtype",
157
+ "lm_head_dtype",
158
+ "router_dtype",
159
+ "embedding_dtype",
160
+ "norm_weight_dtype",
161
+ "activation_dtype",
162
+ "ccl_dtype",
163
+ "kv_cache_dtype",
164
+ "logits_dtype",
165
+ "sampling_dtype",
166
+ }
167
+ )
168
+ _FIDELITY_FIELDS = frozenset(
169
+ {
170
+ "experts_fidelity",
171
+ "attention_fidelity",
172
+ "router_window_fidelity",
173
+ "lm_head_fidelity",
174
+ "norm_fidelity",
175
+ }
176
+ )
177
+ _INT_FIELDS = frozenset(
178
+ {
179
+ "experts_gate_up_in0_block_w",
180
+ "experts_down_in0_block_w",
181
+ "prefill_experts_gate_up_in0_block_w",
182
+ "prefill_experts_down_in0_block_w",
183
+ }
184
+ )
185
+
186
+
187
+ @dataclass(frozen=True)
188
+ class PrecisionConfig:
189
+ """One model's dtype / fidelity policy.
190
+
191
+ Constructed with no arguments this is ``DEFAULT_PRECISION``: the shipped
192
+ stage-06 policy. Vary a field with :meth:`with_overrides` (or plain
193
+ ``dataclasses.replace``) to get a different model out of the same source.
194
+ """
195
+
196
+ # -- weight dtypes, per group ---------------------------------------------
197
+ # The two expert weights are separate fields even though they ship at the
198
+ # same dtype: they have different K, different tuned block widths, and
199
+ # different sensitivity (down feeds the residual directly), so a sweep that
200
+ # could only move them together would be unable to price them apart.
201
+ experts_gate_up_dtype: "ttnn.DataType" = ttnn.bfloat4_b
202
+ experts_down_dtype: "ttnn.DataType" = ttnn.bfloat4_b
203
+ # Likewise qkv and wo: both bfloat8_b today, but wo is the one whose output
204
+ # goes straight into the attention all-reduce.
205
+ attention_qkv_dtype: "ttnn.DataType" = ttnn.bfloat8_b
206
+ attention_wo_dtype: "ttnn.DataType" = ttnn.bfloat8_b
207
+ lm_head_dtype: "ttnn.DataType" = ttnn.bfloat8_b
208
+ router_dtype: "ttnn.DataType" = ttnn.bfloat16
209
+ embedding_dtype: "ttnn.DataType" = ttnn.bfloat16
210
+ # RMSNorm weights and the per-head q_norm/k_norm vectors. 4 KB each; here
211
+ # for completeness of the picture rather than because it is a lever.
212
+ norm_weight_dtype: "ttnn.DataType" = ttnn.bfloat16
213
+
214
+ # -- per-group compute fidelity -------------------------------------------
215
+ experts_fidelity: "ttnn.MathFidelity" = ttnn.MathFidelity.LoFi
216
+ # ``None`` == the op default, which is what the projections take today.
217
+ attention_fidelity: "ttnn.MathFidelity | None" = None
218
+ # HiFi4 so the one-hot expert-window matmul selects rather than approximates
219
+ # -- see ``multichip_decoder._exact_matmul_config``. Lowering this is a
220
+ # correctness change, not a speed/accuracy trade; it is configurable so a
221
+ # sweep can *demonstrate* that rather than assert it.
222
+ router_window_fidelity: "ttnn.MathFidelity" = ttnn.MathFidelity.HiFi4
223
+ lm_head_fidelity: "ttnn.MathFidelity" = ttnn.MathFidelity.HiFi2
224
+ norm_fidelity: "ttnn.MathFidelity" = ttnn.MathFidelity.HiFi4
225
+
226
+ # -- expert matmul inner block widths -------------------------------------
227
+ # Tuned against ``experts_fidelity``; see the module docstring.
228
+ #
229
+ # **Stage 07 moved these, and they are the only fields the sweep moved.**
230
+ # They were 16 and 12, inherited from the single-chip stage-02 tuning. The
231
+ # 48-layer sweep measured both brackets end to end and found each monotonic
232
+ # upward to its full-K ceiling, at *identical* accuracy:
233
+ #
234
+ # gate_up (K = hidden_size 2048 = 64 tiles): 8 -> 41.33, 16 -> 42.34,
235
+ # 32 -> 42.94, 64 -> 43.23 t/s/u
236
+ # down (K = moe_intermediate_size 768 = 24 tiles):
237
+ # 6 -> 41.67, 12 -> 42.34,
238
+ # 24 -> 42.99 t/s/u
239
+ #
240
+ # and the combination at both ceilings measured 43.54 t/s/u -- +2.83% over
241
+ # the shipped default, top-1 0.990 / top-5 1.000 / top-100 1.000, i.e. no
242
+ # accuracy cost at all, because a block width is a *scheduling* choice and
243
+ # not a numerical one. The stage-02 comment that 16 wins at LoFi predates
244
+ # expert parallelism, which cut per-die N four-fold and changed which
245
+ # blocking the matmul wants.
246
+ #
247
+ # Both values are exact divisors of K in tiles, so
248
+ # ``_tuned_sparse_matmul_config`` does not clamp them; ``fallback_audit``
249
+ # reports the resolved widths and stage 07 asserts on those, not on these.
250
+ # See ``doc/datatype_sweep/README.md``.
251
+ experts_gate_up_in0_block_w: int = 64
252
+ experts_down_in0_block_w: int = 24
253
+
254
+ # PREFILL-only inner block widths. Default to the decode values above, so
255
+ # the shipped graph is byte-identical until a sweep says otherwise.
256
+ #
257
+ # Why they exist separately: the two widths above were selected by stage 07,
258
+ # the DATATYPE sweep, optimised against single-token DECODE (52.05 t/s/u),
259
+ # where the expert matmuls run at M = 1. Prefill runs the same matmuls at
260
+ # M = EXPERT_CHUNK_SIZE = 32 and simply inherited decode's tuning; an
261
+ # optimum for a latency-bound M=1 matmul has no reason to be the optimum at
262
+ # M=32. Consumed at optimized_decoder.py:782-784 (prefill) while 926-927
263
+ # (decode, M=1) keeps reading the fields above.
264
+ #
265
+ # Scheduling only, not numerics: the graph and the tokens are the same, only
266
+ # the matmuls' inner blocking differs.
267
+ #
268
+ # SHIPPED 16/12, which are the PRE-stage-07 values. Stage 07 moved the decode
269
+ # fields 16 -> 64 and 12 -> 24 and, because prefill shared them, regressed
270
+ # prefill by ~4 % without measuring it. Sweeping prefill on its own lands
271
+ # exactly back on the old pair:
272
+ #
273
+ # 64/24 (stage 07) 3.741 s 16/24 3.664 s
274
+ # 32/12 3.604 s 16/12 3.596 s <- 1.0404x
275
+ #
276
+ # at 4,096 tokens, 15 configurations, every one PCC 1.0000000000000058 with
277
+ # identical greedy tokens -- scheduling only, no numerical dimension. The 4 %
278
+ # is modest but free, and decode keeps its own 64/24 above.
279
+ # See doc/batch_scaling/README.md, "Where prefill time is NOT going".
280
+ prefill_experts_gate_up_in0_block_w: int = 16
281
+ prefill_experts_down_in0_block_w: int = 12
282
+
283
+ # -- activations ----------------------------------------------------------
284
+ # The dtype of every hidden state, including the inter-layer residual. The
285
+ # residual *layout* (replicated ``[1, 1, rows, 2048]``, DRAM interleaved) is
286
+ # a contract and is not configurable here.
287
+ activation_dtype: "ttnn.DataType" = ttnn.bfloat16
288
+ # ``None`` == run the collective at the activation dtype, no cast. A named
289
+ # dtype casts in and out around the reduce-scatter/all-gather pair.
290
+ ccl_dtype: "ttnn.DataType | None" = None
291
+
292
+ # -- kv cache -------------------------------------------------------------
293
+ kv_cache_dtype: "ttnn.DataType" = ttnn.bfloat16
294
+
295
+ # -- terminal path --------------------------------------------------------
296
+ logits_dtype: "ttnn.DataType" = ttnn.bfloat16
297
+ # What the sampler is handed. Equal to ``logits_dtype`` by default, so the
298
+ # shipped path casts nothing.
299
+ sampling_dtype: "ttnn.DataType" = ttnn.bfloat16
300
+
301
+ def __post_init__(self) -> None:
302
+ names = {f.name for f in fields(self)}
303
+ unclassified = names - _DTYPE_FIELDS - _FIDELITY_FIELDS - _INT_FIELDS
304
+ assert not unclassified, f"PrecisionConfig fields with no serialisation rule: {sorted(unclassified)}"
305
+ for name in _DTYPE_FIELDS:
306
+ object.__setattr__(self, name, dtype_from_name(getattr(self, name)))
307
+ for name in _FIDELITY_FIELDS:
308
+ object.__setattr__(self, name, fidelity_from_name(getattr(self, name)))
309
+ for name in _INT_FIELDS:
310
+ value = int(getattr(self, name))
311
+ if value < 1:
312
+ raise ValueError(f"{name} must be >= 1, got {value}")
313
+ object.__setattr__(self, name, value)
314
+ # ``None`` is legal only where the docstring says it is.
315
+ for name in _DTYPE_FIELDS | _FIDELITY_FIELDS:
316
+ if getattr(self, name) is None and name not in ("ccl_dtype", "attention_fidelity"):
317
+ raise ValueError(f"{name} may not be None")
318
+
319
+ # -- convenience ----------------------------------------------------------
320
+
321
+ def with_overrides(self, **overrides) -> "PrecisionConfig":
322
+ """A copy with ``overrides`` applied; values may be names or objects."""
323
+ unknown = set(overrides) - {f.name for f in fields(self)}
324
+ if unknown:
325
+ raise ValueError(f"unknown precision fields: {sorted(unknown)}")
326
+ return replace(self, **overrides)
327
+
328
+ @property
329
+ def effective_ccl_dtype(self):
330
+ """The dtype the collectives actually run at, resolving ``None``."""
331
+ return self.activation_dtype if self.ccl_dtype is None else self.ccl_dtype
332
+
333
+ # -- serialisation --------------------------------------------------------
334
+
335
+ def to_dict(self) -> dict:
336
+ """JSON-ready ``{field: name}``. Every field of the dataclass appears."""
337
+ out: dict = {}
338
+ for name, value in asdict(self).items():
339
+ if name in _DTYPE_FIELDS:
340
+ out[name] = dtype_to_name(value)
341
+ elif name in _FIDELITY_FIELDS:
342
+ out[name] = fidelity_to_name(value)
343
+ else:
344
+ out[name] = int(value)
345
+ return out
346
+
347
+ @classmethod
348
+ def from_dict(cls, data: dict) -> "PrecisionConfig":
349
+ known = {f.name for f in fields(cls)}
350
+ unknown = set(data) - known
351
+ if unknown:
352
+ raise ValueError(f"unknown precision fields in config: {sorted(unknown)}")
353
+ return cls(**{k: v for k, v in data.items() if k in known})
354
+
355
+ def to_json(self, indent: int = 2) -> str:
356
+ return json.dumps(self.to_dict(), indent=indent, sort_keys=True) + "\n"
357
+
358
+ @classmethod
359
+ def from_json(cls, text: str) -> "PrecisionConfig":
360
+ return cls.from_dict(json.loads(text))
361
+
362
+ def write_json(self, path: str | Path) -> Path:
363
+ path = Path(path)
364
+ path.parent.mkdir(parents=True, exist_ok=True)
365
+ path.write_text(self.to_json())
366
+ return path
367
+
368
+ @classmethod
369
+ def read_json(cls, path: str | Path) -> "PrecisionConfig":
370
+ return cls.from_json(Path(path).read_text())
371
+
372
+
373
+ #: The shipped policy, as stage 07 selected it.
374
+ #:
375
+ #: Every stage-02..06 number was measured at this config **except for the two
376
+ #: expert block widths**, which stage 07 moved from 16/12 to 64/24; those stages
377
+ #: ran at 16/12. Nothing else here has changed since stage 02.
378
+ DEFAULT_PRECISION = PrecisionConfig()
code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/weight_mapping.py ADDED
@@ -0,0 +1,184 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """HuggingFace -> TTNN weight mapping for Qwen3-Coder-30B-A3B-Instruct.
5
+
6
+ RoPE convention (the decision this file encodes)
7
+ ------------------------------------------------
8
+ TTNN offers two rotary embedding ops, and each demands a different weight
9
+ layout. They are interchangeable only as a *pair*; crossing them runs fine and
10
+ silently produces garbage.
11
+
12
+ Meta-style ``rotary_embedding_llama`` head channels interleaved
13
+ -> q/k rows must be reordered (``reverse_permute``), and because
14
+ that reorders channels *within* a head, Qwen3's per-head
15
+ QK-norm weights must be reordered to match as well.
16
+
17
+ HF-style ``ttnn.experimental.rotary_embedding`` HF's native layout
18
+ -> no weight transformation at all.
19
+
20
+ This port uses **HF-style**, matching models/demos/gemma4 (the Blackhole
21
+ exemplar this attention block is written against, and the only other supported
22
+ model with Qwen3-shaped per-head QK-norm). Keeping the checkpoint layout
23
+ untouched removes both permutation steps, so the QK-norm weights are copied
24
+ verbatim. If the RoPE op in functional_decoder.py is ever swapped for the llama
25
+ variant, both permutations have to come back together.
26
+
27
+ Expert fusion (MoE)
28
+ -------------------
29
+ The checkpoint stores 3 tensors per expert. TTNN wants them batched, with gate
30
+ and up fused as ``[gate ; up]`` along the output dim, matching
31
+ ``Qwen3MoeExperts.forward``'s ``chunk(2, dim=-1)``.
32
+
33
+ The fused QKV layout follows models/tt_transformers/tt/attention.py: transpose
34
+ each projection to ``[in, out]``, then concatenate ``[q, k, v]`` along the
35
+ output dim.
36
+ """
37
+
38
+ from __future__ import annotations
39
+
40
+ import torch
41
+
42
+
43
+ def convert_attention_weights(
44
+ sd: dict[str, torch.Tensor],
45
+ *,
46
+ n_heads: int,
47
+ n_kv_heads: int,
48
+ head_dim: int,
49
+ ) -> dict[str, torch.Tensor]:
50
+ """Return the fused ``wqkv`` plus ``wo`` and the QK-norm weights.
51
+
52
+ Channel order is left exactly as HuggingFace stores it -- see the RoPE
53
+ convention note in the module docstring. The only structural change is
54
+ fusing Q, K and V into one matmul.
55
+
56
+ ``sd`` uses layer-relative keys (``self_attn.q_proj.weight``, ...).
57
+ Single-device layout; tensor-parallel chunking is a later stage.
58
+ """
59
+ q = sd["self_attn.q_proj.weight"].float()
60
+ k = sd["self_attn.k_proj.weight"].float()
61
+ v = sd["self_attn.v_proj.weight"].float()
62
+ o = sd["self_attn.o_proj.weight"].float()
63
+
64
+ assert q.shape[0] == n_heads * head_dim, f"q_proj {tuple(q.shape)} != {n_heads}x{head_dim}"
65
+ assert k.shape[0] == n_kv_heads * head_dim, f"k_proj {tuple(k.shape)} != {n_kv_heads}x{head_dim}"
66
+
67
+ # torch stores nn.Linear as [out, in]; TTNN matmuls right-multiply, so
68
+ # transpose to [in, out] before concatenating along the output dim.
69
+ wqkv = torch.cat([q.T, k.T, v.T], dim=-1).unsqueeze(0).unsqueeze(0)
70
+
71
+ out = {
72
+ "wqkv": wqkv,
73
+ "wo": o.T.contiguous(),
74
+ }
75
+
76
+ for name, key in (("q_norm", "self_attn.q_norm.weight"), ("k_norm", "self_attn.k_norm.weight")):
77
+ if key in sd:
78
+ out[name] = sd[key].float()
79
+
80
+ return out
81
+
82
+
83
+ def hf_to_meta_channels(head_dim: int) -> torch.Tensor:
84
+ """Index vector mapping an HF-ordered head to Meta (llama) channel order.
85
+
86
+ HF's ``rotate_half`` pairs channel ``i`` with ``i + head_dim/2``; the Meta
87
+ convention that ``rotary_embedding_llama``'s transformation matrix encodes
88
+ pairs ``2i`` with ``2i + 1``. So ``meta[2i] = hf[i]`` and
89
+ ``meta[2i+1] = hf[i + head_dim/2]``, i.e. an interleave of the two halves.
90
+
91
+ The *same* vector converts a cos/sin row, because HF stores
92
+ ``[c0 .. c_{d/2-1}, c0 .. c_{d/2-1}]`` and Meta stores ``[c0, c0, c1, c1, ...]``.
93
+
94
+ Applied per head, so it commutes with any reordering of whole heads --
95
+ which is what lets ``multichip_decoder`` apply it before
96
+ ``head_interleaved_wqkv`` splits the heads across dies.
97
+ """
98
+ assert head_dim % 2 == 0, head_dim
99
+ half = head_dim // 2
100
+ return torch.stack([torch.arange(half), torch.arange(half) + half], dim=1).reshape(-1)
101
+
102
+
103
+ def permute_wqkv_to_meta(wqkv: torch.Tensor, *, n_heads: int, n_kv_heads: int, head_dim: int) -> torch.Tensor:
104
+ """Reorder the Q and K channels of a fused ``[..., in, out]`` wqkv to Meta order.
105
+
106
+ **V is deliberately untouched** -- RoPE is applied to Q and K only, so V's
107
+ channels keep their HF meaning, and so does ``wo``, which consumes the
108
+ attention output in V's space.
109
+
110
+ This is the ``reverse_permute`` step the module docstring says comes back
111
+ the moment the llama rotary op is used. It is a *checkpoint layout* change
112
+ and is applied once at upload, never per token.
113
+ """
114
+ out = wqkv.clone()
115
+ perm = hf_to_meta_channels(head_dim)
116
+ q_end = n_heads * head_dim
117
+ k_end = q_end + n_kv_heads * head_dim
118
+ for start, count in ((0, n_heads), (q_end, n_kv_heads)):
119
+ for h in range(count):
120
+ lo = start + h * head_dim
121
+ out[..., lo : lo + head_dim] = wqkv[..., lo : lo + head_dim][..., perm]
122
+ assert torch.equal(out[..., k_end:], wqkv[..., k_end:]), "V must not be permuted"
123
+ return out
124
+
125
+
126
+ def permute_head_vector_to_meta(vec: torch.Tensor, *, head_dim: int) -> torch.Tensor:
127
+ """Reorder a per-head vector (Qwen3's ``q_norm`` / ``k_norm``) to Meta order.
128
+
129
+ Qwen3 applies these **between the head split and RoPE**, so they index the
130
+ same channels the rotary op does. Permuting Q/K without permuting these
131
+ scales the wrong channel and is silent -- it does not change any shape and
132
+ it does not raise. ``test_meta_rope_weights_match_hf`` is the assertion that
133
+ catches it.
134
+ """
135
+ flat = vec.reshape(-1)
136
+ assert flat.numel() == head_dim, (flat.shape, head_dim)
137
+ return flat[hf_to_meta_channels(head_dim)].reshape(vec.shape)
138
+
139
+
140
+ def convert_moe_weights(sd: dict[str, torch.Tensor], *, n_experts: int) -> dict[str, torch.Tensor]:
141
+ """Batch the per-expert checkpoint tensors and fuse gate/up.
142
+
143
+ gate first, then up -- ``Qwen3MoeExperts.forward`` chunks the matmul output
144
+ in half and treats the FIRST half as gate.
145
+ """
146
+ gate_up = torch.stack(
147
+ [
148
+ torch.cat(
149
+ [
150
+ sd[f"mlp.experts.{e}.gate_proj.weight"].float(),
151
+ sd[f"mlp.experts.{e}.up_proj.weight"].float(),
152
+ ],
153
+ dim=0,
154
+ )
155
+ for e in range(n_experts)
156
+ ]
157
+ )
158
+ down = torch.stack([sd[f"mlp.experts.{e}.down_proj.weight"].float() for e in range(n_experts)])
159
+ return {
160
+ "router": sd["mlp.gate.weight"].float(),
161
+ "experts_gate_up": gate_up,
162
+ "experts_down": down,
163
+ }
164
+
165
+
166
+ def convert_norm_weights(sd: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
167
+ """Decoder-layer RMSNorm weights. Plain RMSNorm -- no zero-centering."""
168
+ return {
169
+ "input_layernorm": sd["input_layernorm.weight"].float(),
170
+ "post_attention_layernorm": sd["post_attention_layernorm.weight"].float(),
171
+ }
172
+
173
+
174
+ def convert_layer_weights(sd: dict[str, torch.Tensor], config) -> dict[str, torch.Tensor]:
175
+ """Full layer conversion: attention + MoE + norms."""
176
+ weights = convert_attention_weights(
177
+ sd,
178
+ n_heads=config.num_attention_heads,
179
+ n_kv_heads=config.num_key_value_heads,
180
+ head_dim=config.head_dim,
181
+ )
182
+ weights.update(convert_moe_weights(sd, n_experts=config.num_experts))
183
+ weights.update(convert_norm_weights(sd))
184
+ return weights
code/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle/qwen3_coder_30b_a3b_instruct/tt_qwen3_coder_30b_a3b_instruct.py ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Bundle entry point for ``EXTRA_MODELS_DIR``.
5
+
6
+ ``vllm_tt_plugin.platform.register_tt_models()`` runs
7
+ ``_register_models_from_extra_dir(ModelRegistry)`` as its **first** action --
8
+ "so a distributed bundle can supply a model without touching this file". That
9
+ hook appends this folder to ``sys.path`` and lazily registers
10
+ ``vllm_metadata.json``'s ``main_class`` under the plugin's ``TT``-prefixed
11
+ convention, so the model ends up registered *by* ``register_tt_models()`` with
12
+ no edit to the plugin checkout.
13
+
14
+ Registration is lazy: vLLM resolves the ``"module:Class"`` string later, in the
15
+ API-server process and again in each EngineCore worker. This module therefore
16
+ has to be importable on its own, which means it cannot assume the tt-metal
17
+ checkout is already on ``sys.path`` -- an EngineCore worker's working directory
18
+ is not guaranteed. It appends the repository root (never ``insert(0)``, matching
19
+ the hook's own rule that an installed package of the same name must still win)
20
+ and re-exports the real adapter, which lives with the model it adapts.
21
+ """
22
+
23
+ from __future__ import annotations
24
+
25
+ import sys
26
+ from pathlib import Path
27
+
28
+ # .../models/demos/blackhole/<model>/vllm_bundle/<bundle>/this_file.py
29
+ _REPO_ROOT = Path(__file__).resolve().parents[6]
30
+ if str(_REPO_ROOT) not in sys.path:
31
+ sys.path.append(str(_REPO_ROOT))
32
+
33
+ from models.demos.blackhole.qwen3_coder_30b_a3b.tt.generator_vllm import Qwen3CoderForCausalLM # noqa: E402
34
+
35
+ __all__ = ["Qwen3CoderForCausalLM"]
code/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle/qwen3_coder_30b_a3b_instruct/vllm_metadata.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "arch": "Qwen3MoeForCausalLM",
3
+ "main_class": "tt_qwen3_coder_30b_a3b_instruct:Qwen3CoderForCausalLM",
4
+ "hf_model": "Qwen/Qwen3-Coder-30B-A3B-Instruct",
5
+ "adapter": "models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator_vllm.py",
6
+ "registered_as": "TTQwen3MoeForCausalLM",
7
+ "notes": "Self-contained bundle for vllm_tt_plugin.platform.register_tt_models() -> _register_models_from_extra_dir(). Point EXTRA_MODELS_DIR at the parent vllm_bundle/ directory. 'arch' is this checkout's config.json architecture; the plugin prefixes it with TT. 'main_class' resolves through the sibling shim module, which is importable because the plugin appends this folder to sys.path."
8
+ }
image/blobs/sha256/0926a8eb0e608a5c6888d1cd5594184bdf3ed3aa311dba5b42a547caefdc6f2e ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0926a8eb0e608a5c6888d1cd5594184bdf3ed3aa311dba5b42a547caefdc6f2e
3
+ size 29752807
image/blobs/sha256/24cba7375920bef8d4cc4f0ce4294f8f70c65b9a37c4ed6f6c2d63405c76ba3c ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:24cba7375920bef8d4cc4f0ce4294f8f70c65b9a37c4ed6f6c2d63405c76ba3c
3
+ size 1385020
image/blobs/sha256/3de1f5eb93e54b4561afa733cf844d1914a9ba260f49e86fe600f344c5dd025c ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3de1f5eb93e54b4561afa733cf844d1914a9ba260f49e86fe600f344c5dd025c
3
+ size 35141596
image/blobs/sha256/530b0e35f44c6f963e06fdaacdbfecb2021d4c55d49fe4c0019e161d94c18de3 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:530b0e35f44c6f963e06fdaacdbfecb2021d4c55d49fe4c0019e161d94c18de3
3
+ size 10547489
image/blobs/sha256/540cf00275e913a9bccc49fe7beba58037696661bde6ae083a3ac84c5b160e67 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:540cf00275e913a9bccc49fe7beba58037696661bde6ae083a3ac84c5b160e67
3
+ size 9565684
image/blobs/sha256/8753e0cfbd424e962ccaf50aaaf02fd06ff2efb8677219657e751a53922efa9f ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8753e0cfbd424e962ccaf50aaaf02fd06ff2efb8677219657e751a53922efa9f
3
+ size 1076048
image/blobs/sha256/b6df468b82a4b2f9ee3ca3a79a6bbe99b5cda02ddc63b7e3c89e7bd08ef41706 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b6df468b82a4b2f9ee3ca3a79a6bbe99b5cda02ddc63b7e3c89e7bd08ef41706
3
+ size 48137518
image/blobs/sha256/c00314a02c644cf2aea5a8ae4e3ee5ee383f4072d3e35248fe20dea57a90c46c ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c00314a02c644cf2aea5a8ae4e3ee5ee383f4072d3e35248fe20dea57a90c46c
3
+ size 1727127131
image/blobs/sha256/c18d0f3c8022bcd5a8059f67fc0e5cfd53ef36a44997335d6cd0aa6b19db140d ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c18d0f3c8022bcd5a8059f67fc0e5cfd53ef36a44997335d6cd0aa6b19db140d
3
+ size 35141895
image/blobs/sha256/c3bc7b373b4523cabcdd9c64ab10d32510ef61bace60873a279ccc4902738989 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c3bc7b373b4523cabcdd9c64ab10d32510ef61bace60873a279ccc4902738989
3
+ size 139055657
image/blobs/sha256/ca0b072b65f8c21199f96e7498f2ccbc252490ce39060800647332751f857287 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ca0b072b65f8c21199f96e7498f2ccbc252490ce39060800647332751f857287
3
+ size 35988690
image/blobs/sha256/cbb77a738c7df827819a8b8bf87682eeca8bdb434d41411a9682c38717f2f187 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cbb77a738c7df827819a8b8bf87682eeca8bdb434d41411a9682c38717f2f187
3
+ size 82239359
image/blobs/sha256/e19b6f1fb65dd2888d9003ef9513a21d129041a328bb8a9a4164d29ef0382b16 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e19b6f1fb65dd2888d9003ef9513a21d129041a328bb8a9a4164d29ef0382b16
3
+ size 138938826
image/blobs/sha256/fdc1ed79ffd24d66f8be3754ec4dc80b1ab0fcc8c8165a6007748c464c94f897 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fdc1ed79ffd24d66f8be3754ec4dc80b1ab0fcc8c8165a6007748c464c94f897
3
+ size 415353