Add files using upload-large-folder tool
Browse files- .gitattributes +14 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/README.md +143 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/__init__.py +0 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/config/context_contract.json +360 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/config/selected_precision_config.json +22 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/__init__.py +0 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/reference.py +145 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_attention.py +117 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_attention_decode.py +140 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decode_compaction_fifo.py +596 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decoder_layer.py +176 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_decoder_layer_decode.py +155 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_determinism.py +167 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_full_model.py +976 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_moe.py +232 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_multichip_decoder.py +1057 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_optimized_decoder.py +503 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_perf.py +611 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_precision_config.py +369 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_reference.py +123 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_rmsnorm.py +99 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tests/test_trace.py +200 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt-model-localplugin.yaml +152 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt-model.yaml +242 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/__init__.py +0 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/functional_decoder.py +1169 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator.py +1637 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/generator_vllm.py +1568 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/model.py +1723 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/multichip_decoder.py +1982 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/optimized_decoder.py +1093 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/precision.py +378 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/tt/weight_mapping.py +184 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle/qwen3_coder_30b_a3b_instruct/tt_qwen3_coder_30b_a3b_instruct.py +35 -0
- code/models/demos/blackhole/qwen3_coder_30b_a3b/vllm_bundle/qwen3_coder_30b_a3b_instruct/vllm_metadata.json +8 -0
- image/blobs/sha256/0926a8eb0e608a5c6888d1cd5594184bdf3ed3aa311dba5b42a547caefdc6f2e +3 -0
- image/blobs/sha256/24cba7375920bef8d4cc4f0ce4294f8f70c65b9a37c4ed6f6c2d63405c76ba3c +3 -0
- image/blobs/sha256/3de1f5eb93e54b4561afa733cf844d1914a9ba260f49e86fe600f344c5dd025c +3 -0
- image/blobs/sha256/530b0e35f44c6f963e06fdaacdbfecb2021d4c55d49fe4c0019e161d94c18de3 +3 -0
- image/blobs/sha256/540cf00275e913a9bccc49fe7beba58037696661bde6ae083a3ac84c5b160e67 +3 -0
- image/blobs/sha256/8753e0cfbd424e962ccaf50aaaf02fd06ff2efb8677219657e751a53922efa9f +3 -0
- image/blobs/sha256/b6df468b82a4b2f9ee3ca3a79a6bbe99b5cda02ddc63b7e3c89e7bd08ef41706 +3 -0
- image/blobs/sha256/c00314a02c644cf2aea5a8ae4e3ee5ee383f4072d3e35248fe20dea57a90c46c +3 -0
- image/blobs/sha256/c18d0f3c8022bcd5a8059f67fc0e5cfd53ef36a44997335d6cd0aa6b19db140d +3 -0
- image/blobs/sha256/c3bc7b373b4523cabcdd9c64ab10d32510ef61bace60873a279ccc4902738989 +3 -0
- image/blobs/sha256/ca0b072b65f8c21199f96e7498f2ccbc252490ce39060800647332751f857287 +3 -0
- image/blobs/sha256/cbb77a738c7df827819a8b8bf87682eeca8bdb434d41411a9682c38717f2f187 +3 -0
- image/blobs/sha256/e19b6f1fb65dd2888d9003ef9513a21d129041a328bb8a9a4164d29ef0382b16 +3 -0
- 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
|