ndaly commited on
Commit
4fc4cfb
·
verified ·
1 Parent(s): da5c67b

Add files using upload-large-folder tool

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. README.md +113 -0
  2. code/models/common/modules/moe/configs/deepseek_ocr.yaml +28 -0
  3. code/models/common/modules/moe/configs/deepseek_v3.yaml +28 -0
  4. code/models/common/modules/moe/configs/deepseek_v3_single_glx.yaml +28 -0
  5. code/models/common/modules/moe/configs/deepseek_v4_flash.yaml +32 -0
  6. code/models/common/modules/moe/configs/deepseek_v4_pro.yaml +32 -0
  7. code/models/common/modules/moe/configs/gemma_4_26b.yaml +35 -0
  8. code/models/common/modules/moe/configs/glm5.yaml +30 -0
  9. code/models/common/modules/moe/configs/glm_47.yaml +31 -0
  10. code/models/common/modules/moe/configs/gpt_oss.yaml +33 -0
  11. code/models/common/modules/moe/configs/kimi_k25.yaml +30 -0
  12. code/models/common/modules/moe/configs/ling_1t.yaml +29 -0
  13. code/models/common/modules/moe/configs/models_table.md +18 -0
  14. code/models/common/modules/moe/configs/qwen35_35b.yaml +26 -0
  15. code/models/common/modules/moe/configs/qwen35_397b.yaml +31 -0
  16. code/models/common/modules/moe/configs/qwen3_235b.yaml +27 -0
  17. code/models/common/modules/moe/configs/qwen3_omni_talker.yaml +29 -0
  18. code/models/common/modules/moe/configs/qwen3_omni_thinker.yaml +27 -0
  19. code/models/common/tests/demos/cleanup_utils.py +113 -0
  20. code/models/common/tests/demos/run_helpers.py +1071 -0
  21. code/models/common/tests/demos/test_cleanup_utils.py +154 -0
  22. code/models/common/tests/demos/test_run_helpers.py +687 -0
  23. code/models/common/tests/host/test_metrics_pytorch_only.py +61 -0
  24. code/models/common/tests/host/test_utility_functions_imports.py +31 -0
  25. code/models/common/tests/llm_runtime/test_config.py +147 -0
  26. code/models/common/tests/llm_runtime/test_decode_runtime.py +1337 -0
  27. code/models/common/tests/llm_runtime/test_execution.py +831 -0
  28. code/models/common/tests/llm_runtime/test_executor_integration.py +1833 -0
  29. code/models/common/tests/llm_runtime/test_lane_group.py +1373 -0
  30. code/models/common/tests/llm_runtime/test_llama3_8b_integration.py +1065 -0
  31. code/models/common/tests/llm_runtime/test_llama3_8b_model_contract.py +349 -0
  32. code/models/common/tests/llm_runtime/test_model_contract.py +972 -0
  33. code/models/common/tests/llm_runtime/test_model_executor.py +291 -0
  34. code/models/common/tests/llm_runtime/test_output_reader.py +221 -0
  35. code/models/common/tests/llm_runtime/test_paged_kv_cache.py +415 -0
  36. code/models/common/tests/llm_runtime/test_prefill_inputs.py +173 -0
  37. code/models/common/tests/llm_runtime/test_prefill_runtime.py +0 -0
  38. code/models/common/tests/llm_runtime/test_program_compiler.py +268 -0
  39. code/models/common/tests/llm_runtime/test_tensor_resources.py +62 -0
  40. code/models/common/tests/llm_runtime/test_trace_compiler.py +576 -0
  41. code/models/common/tests/llm_runtime/test_vllm_adapter.py +1025 -0
  42. code/models/common/tests/llm_runtime/test_warmup.py +1012 -0
  43. code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_demo_contract.py +452 -0
  44. code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_hf_adaptor.py +290 -0
  45. code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_prefill_last_token_contract.py +71 -0
  46. code/models/common/tests/models/llama32_1b/test_batched_prefill_postprocess.py +104 -0
  47. code/models/common/tests/models/llama32_1b/test_demo_warmup.py +134 -0
  48. code/models/common/tests/models/llama32_1b/test_hf_adaptor.py +264 -0
  49. code/models/common/tests/models/llama32_3b/test_batched_prefill_postprocess.py +104 -0
  50. code/models/common/tests/models/llama32_3b/test_demo_warmup.py +218 -0
README.md ADDED
@@ -0,0 +1,113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ tags:
3
+ - blackhole
4
+ - p300x2
5
+ - tt-model-cache
6
+ - tt-model-container
7
+ - vllm-plugin
8
+ license: apache-2.0
9
+ pipeline_tag: automatic-speech-recognition
10
+ base_model:
11
+ - mistralai/Voxtral-Mini-4B-Realtime-2602
12
+ ---
13
+
14
+ # voxtral-mini-4b-realtime-2602-p300x2
15
+
16
+ Voxtral-Mini-4B-Realtime-2602, Mistral AI's realtime speech-to-text model (a 970M-parameter causal audio encoder feeding a 3.4B Ministral text decoder; 13 languages; one text token per 80 ms of audio), served on four Blackhole chips (2 x p300c, 1x4 ring, TP=4) through a TTNN autoport and the Tenstorrent vLLM plugin as an OpenAI-compatible `/v1/audio/transcriptions` endpoint.
17
+
18
+ Runs on **p300x2** (mesh `P300x2`) — 131,072-token context, up to 32 concurrent sequences.
19
+
20
+ Packaged and published with [tt-model-manager](https://github.com/tenstorrent/tt-model-manager) 0.1.0 (manifest schema 5.1).
21
+
22
+ ## At a glance
23
+
24
+ | | |
25
+ | --- | --- |
26
+ | Architecture | 4B realtime ASR (970M causal audio encoder + 3.4B Ministral-3B decoder) |
27
+ | Hardware | p300x2 |
28
+ | Context | 131,072 tokens |
29
+ | License | apache-2.0 |
30
+ | Status | Experimental community bring-up (tt-model-bringup autoport, 11 stages, locally committed) |
31
+
32
+ ## Intended use
33
+
34
+ **Direct use:** Speech-to-text transcription of mono speech audio (wav, flac, mp3 and other formats soundfile decodes; resampled to 16 kHz) through the OpenAI audio transcriptions API, one interactive user or up to 32 concurrent requests, in the 13 languages the checkpoint covers.
35
+
36
+ **Out-of-scope use:** Chat, text generation, translation, speaker diarization, or audio question answering: the checkpoint is transcription-only and the server exposes only that task.
37
+
38
+ ## Quickstart
39
+
40
+ ```bash
41
+ uv tool install tenstorrent # once — the Tenstorrent CLI, `tt`
42
+ tt model pull ndaly/Voxtral-Mini-4B-Realtime-2602-tt-p300x2
43
+ tt serve ndaly/Voxtral-Mini-4B-Realtime-2602-tt-p300x2
44
+ ```
45
+
46
+ `tt model pull` (or `tt-model pull --with-weights`) downloads the Docker image and the [`mistralai/Voxtral-Mini-4B-Realtime-2602`](https://huggingface.co/mistralai/Voxtral-Mini-4B-Realtime-2602) weights at `2769294da9567371363522aac9bbcfdd19447add` (into your HF cache; they are not in the image). `tt serve` (or `tt-model serve`) starts an OpenAI-compatible server on port 20000 (or the next free port, if that one is busy); the first start compiles kernels for your device, which takes several minutes, and the server is ready when it logs `Application startup complete`.
47
+
48
+ Without tt-cli — tt-model alone does the whole job:
49
+
50
+ ```bash
51
+ tt-model pull ndaly/Voxtral-Mini-4B-Realtime-2602-tt-p300x2 --with-weights
52
+ tt-model serve ndaly/Voxtral-Mini-4B-Realtime-2602-tt-p300x2
53
+ ```
54
+
55
+ ### What to expect
56
+ - Hardware: a TT-QuietBox 2 or any host with 2 x p300c (4 Blackhole chips, 32 GB GDDR6 each), Docker, 1G hugepages.
57
+ - Weights: one 8.9 GB safetensors file (plus tokenizer and config files) is downloaded into your Hugging Face cache on first `serve`.
58
+ - First boot compiles kernels and captures the decode trace plus 32 per-slot prefill traces (several minutes); later boots reuse the
59
+ kernel cache under `~/.cache/tt-model/voxtral-mini-4b-realtime-2602-p300x2/`.
60
+ - Measured (stage 10 serving run, 2026-09-23): 37.5 ms served TTFT and ~50 tokens/s per user for one user; 606 tokens/s aggregate
61
+ for 32 concurrent requests. Full evidence: `doc/optimized_vllm/README.md` and `doc/tti_release/RUN_NOTES.md` in the tt-metal autoport tree.
62
+ ### Use with your client
63
+ Point an OpenAI-compatible client at `http://127.0.0.1:20000` (or the port `tt-model serve` printed) and call the audio
64
+ transcriptions API with model id `mistralai/Voxtral-Mini-4B-Realtime-2602` (see Using it below).
65
+
66
+ ## Using it
67
+
68
+ The server speaks the OpenAI API at `http://127.0.0.1:20000/v1` (or whichever port your serve command reported; chat completions, completions, and `/v1/models`). Pass `"model": "mistralai/Voxtral-Mini-4B-Realtime-2602"` — the weights id, not this package's name.
69
+
70
+ The endpoint is `POST /v1/audio/transcriptions` (multipart form, OpenAI shape). With the server on the port `tt-model serve` printed
71
+ (20000 by default):
72
+ ```bash
73
+ curl -s http://127.0.0.1:20000/v1/audio/transcriptions \
74
+ -F model=mistralai/Voxtral-Mini-4B-Realtime-2602 -F file=@clip.wav
75
+ # streaming (server-sent events, text deltas as they are decoded):
76
+ curl -sN http://127.0.0.1:20000/v1/audio/transcriptions \
77
+ -F model=mistralai/Voxtral-Mini-4B-Realtime-2602 -F file=@clip.wav -F stream=true
78
+ ```
79
+ Any OpenAI client works the same way (`client.audio.transcriptions.create(model=..., file=...)`). `tt-model curl` sends a chat
80
+ completion and gets a 404 from this model: use the request above instead.
81
+
82
+ ## Expected performance
83
+
84
+ Measured on p300x2 through the served endpoint (stage 10 of the bring-up, 2026-09-23): single user, one 5.28 s clip (78 completion tokens, greedy, streaming): 37.5 ms served time-to-first-token (the adapter's prefill wall), 20.14 ms per decode step (49.7 tokens/s per user), 1.59 s end to end; 32 concurrent requests: 606 tokens/s aggregate, 3.63 s median end to end, 32/32 complete. The generator's own traced decode is 19.55 ms per token. Accuracy: 5.64 % WER on LibriSpeech test-other (2,939 utterances, greedy, tt-inference-server release run through lmms-eval); transcripts identical to the Hugging Face reference on 8 of 10 qualitative clips (the other 2 differ by near-tie words); top-1 token agreement 0.995 against the bf16 reference at the selected precision (bfp4 + LoFi inner projections, bfp8 LM head and KV cache, 6.8 GB per chip).
85
+
86
+ ## Limitations
87
+
88
+ Transcription only: the checkpoint supports no chat or text completion, so the server mounts `/v1/audio/transcriptions` (and `/v1/models`) but not `/v1/chat/completions`; `tt-model curl` cannot talk to it (see Using it). One TP=4 replica on exactly four chips (p300x2); no data parallelism, no other board validated. At most 32 concurrent sequences; 131,072-token context contract (the checkpoint's advertised context); `block_size` is fixed at 32 by the generator. Requires `--tokenizer-mode mistral` (tekken.json). Chunked prefill and prefix caching are disabled by the TT backend. Validated on utterances up to about 35 s (LibriSpeech) and a 30 s benchmark payload; longer audio runs through the same streaming path but was not measured. The vLLM stack is the tenstorrent/vllm fork's empty-target wheel plus its in-tree plugin with local commits, not the standalone vllm-tt-plugin.
89
+
90
+ ## Risks and safety considerations
91
+
92
+ The bfp4 + LoFi precision policy drops one word onset on 1 of 2,939 LibriSpeech test-other utterances where the bf16 reference does not. Under concurrency, 1 of 2,939 utterances came back empty from a rotation-dependent near-tie onset (the Hugging Face reference also returns empty on the evaluator's exact encoding). The release accuracy gate compared against a placeholder reference because no published LibriSpeech test-other WER exists for this checkpoint. Transcripts are greedy by default; sampled decoding is supported but not evaluated for accuracy.
93
+
94
+ ## Licensing
95
+
96
+ Weights under Mistral AI's Apache-2.0 licence; the TTNN port and serving code in `code/` are Apache-2.0 (tt-metal), as are the vLLM fork and its plugin inside the image.
97
+
98
+ ## Feedback
99
+
100
+ Questions or problems with this package: open a discussion at https://huggingface.co/ndaly/Voxtral-Mini-4B-Realtime-2602-tt-p300x2/discussions — that is what reaches its author. A problem with the `tt` tooling itself: `tt report issue` (collects your environment and opens a prefilled issue against tenstorrent/tt-cli). Product feedback: support@tenstorrent.com.
101
+
102
+ ## Provenance
103
+
104
+ The exact sources the image was built from — `code/` in this repo is byte-identical to the model code inside the image:
105
+
106
+ | component | built from |
107
+ | --- | --- |
108
+ | tt-metal | a local checkout — commit not published *(dirty tree — the image includes uncommitted changes)* |
109
+ | vLLM | `vllm-0.1.dev14195+g8c28fcecb.d20260916.empty-py3-none-any.whl` — a wheel the author built |
110
+ | vllm-tt-plugin | a local checkout — commit not published *(dirty tree — the image includes uncommitted changes)* |
111
+ | `code/` digest | `b26d652115a867ab` (sha256, first 16 hex digits) |
112
+ | built | 2026-09-25T16:20:42+00:00 by tt-model 0.1.0 |
113
+
code/models/common/modules/moe/configs/deepseek_ocr.yaml ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # deepseek_ocr
2
+ mesh_shape: [8, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 1280
6
+ select_experts_k: 6
7
+ num_routed_experts: 64
8
+ num_shared_experts: 2
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # deepseek_ocr = ungrouped softmax top-6, no bias. 64 experts → padded to the op's 256-face in
13
+ # TTMoEGate.forward.
14
+ n_group: 1
15
+ score_func: softmax
16
+ routed_scaling_factor: 1.0
17
+ gate_notes: "DeepSeek-OCR: softmax kept POST-topk (softmax-over-selected, the default softmax_position) — mathematically ≡ the model's softmax-over-all→top-6→renorm (global Z cancels, no bias). Set softmax_position: pre for the literal front softmax (numerically-stable exp over all experts)."
18
+
19
+ compute:
20
+ intermediate_size: 896
21
+ activation_type: SILU
22
+
23
+ reduce:
24
+ shared_expert_scale: 1.0
25
+
26
+ experts:
27
+ expert_mapping: sequential
28
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/deepseek_v3.yaml ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # deepseek_v3
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 7168
6
+ select_experts_k: 8
7
+ num_routed_experts: 256
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8)
12
+ n_group: 8
13
+ score_func: sigmoid
14
+ routed_scaling_factor: 2.5
15
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sigmoid scores for SELECTION only
16
+ # (output weights stay unbiased). Explicit per-model flag (not inferred from score_func).
17
+ score_correction_bias: true
18
+
19
+ compute:
20
+ intermediate_size: 2048
21
+ activation_type: SILU
22
+
23
+ reduce:
24
+ shared_expert_scale: 1.0
25
+
26
+ experts:
27
+ expert_mapping: sequential
28
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/deepseek_v3_single_glx.yaml ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # deepseek_v3
2
+ mesh_shape: [8, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 7168
6
+ select_experts_k: 8
7
+ num_routed_experts: 256
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8)
12
+ n_group: 8
13
+ score_func: sigmoid
14
+ routed_scaling_factor: 2.5
15
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sigmoid scores for SELECTION only
16
+ # (output weights stay unbiased). Explicit per-model flag (not inferred from score_func).
17
+ score_correction_bias: true
18
+
19
+ compute:
20
+ intermediate_size: 2048
21
+ activation_type: SILU
22
+
23
+ reduce:
24
+ shared_expert_scale: 1.0
25
+
26
+ experts:
27
+ expert_mapping: sequential
28
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/deepseek_v4_flash.yaml ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # deepseek_v4_flash
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 4096
6
+ select_experts_k: 6
7
+ num_routed_experts: 256
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # deepseek-v4: ungrouped (1 group) top-6 with the sqrtsoftplus score = sqrt(softplus(logit)) (applied
13
+ # externally via ttnn.sqrt(ttnn.softplus(.)), no in-kernel op) + score_correction_bias for selection,
14
+ # linear renormalization, scaled by 1.5. 256 experts → single 256-face (no combine). Same gate as
15
+ # deepseek_v4_pro but 256 experts / scale 1.5 (pro is 384 / 2.5).
16
+ n_group: 1
17
+ score_func: sqrtsoftplus
18
+ routed_scaling_factor: 1.5
19
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sqrtsoftplus scores for SELECTION
20
+ # only (output weights stay unbiased). Explicit per-model flag (not inferred from score_func).
21
+ score_correction_bias: true
22
+
23
+ compute:
24
+ intermediate_size: 2048
25
+ activation_type: SILU
26
+
27
+ reduce:
28
+ shared_expert_scale: 1.0
29
+
30
+ experts:
31
+ expert_mapping: sequential
32
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/deepseek_v4_pro.yaml ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # deepseek_v4_pro
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 7168
6
+ select_experts_k: 6
7
+ num_routed_experts: 384
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # deepseek-v4: ungrouped top-6 with the sqrtsoftplus score = sqrt(softplus(logit)) (applied externally
13
+ # via ttnn.sqrt(ttnn.softplus(.)), no in-kernel op) + score_correction_bias for selection, linear
14
+ # renormalization, scaled by 2.5. 384 experts → the op's 512-combine path (pad 384→512, two 256-blocks →
15
+ # global top-6), same combine path as kimi_k25.
16
+ n_group: 1
17
+ score_func: sqrtsoftplus
18
+ routed_scaling_factor: 2.5
19
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sqrtsoftplus scores for SELECTION
20
+ # only (output weights stay unbiased). Explicit per-model flag (not inferred from score_func).
21
+ score_correction_bias: true
22
+
23
+ compute:
24
+ intermediate_size: 3072
25
+ activation_type: SILU
26
+
27
+ reduce:
28
+ shared_expert_scale: 1.0
29
+
30
+ experts:
31
+ expert_mapping: sequential
32
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/gemma_4_26b.yaml ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # gemma_4_26b
2
+ mesh_shape: [8, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 2816
6
+ select_experts_k: 8
7
+ num_routed_experts: 128
8
+ num_shared_experts: 0
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # Gemma4 router = linear → softmax(all) → top-8 → renormalize. TTMoEGate covers exactly that
13
+ # (score_func=softmax ≡ softmax-front + renorm, the global Z cancels; routed_scaling_factor=1.0 → no
14
+ # scalar scale). 128 experts → padded to the op's 256-face.
15
+ n_group: 1
16
+ score_func: softmax
17
+ routed_scaling_factor: 1.0
18
+
19
+ # REMINDER — Gemma4's gate has steps TTMoEGate does NOT cover (caller must apply, see modular_gemma4.py
20
+ # Gemma4TextRouter.forward):
21
+ # step 1: RMSNorm(hidden) (with_scale=False) — before this gate
22
+ # step 2: hidden *= self.scale[hidden] * hidden_size**-0.5 — a learned per-dim input scale × 1/√hidden
23
+ # step 6b: top_k_weights *= per_expert_scale[selected] — a learned PER-EXPERT output scale ([num_experts]),
24
+ # NOT a scalar, so it can't be folded into routed_scaling_factor.
25
+ gate_notes: "Gemma4: TTMoEGate does steps 3-6a (linear→softmax→top8→renorm). MISSING upstream: RMSNorm + per-dim input scale (×1/√hidden); MISSING downstream: per-expert output scale per_expert_scale[selected] (not a scalar). SOFTMAX kept POST-topk (softmax-over-selected, the default softmax_position) — mathematically ≡ the model's softmax-over-all→top8→renorm (global Z cancels, no bias); set softmax_position: pre for the literal front softmax (numerically-stable exp over all experts)."
26
+
27
+ compute:
28
+ intermediate_size: 704
29
+ activation_type: GELU
30
+
31
+ reduce:
32
+ shared_expert_scale: 1.0
33
+
34
+ experts:
35
+ expert_mapping: sequential
code/models/common/modules/moe/configs/glm5.yaml ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # glm5
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 6144
6
+ select_experts_k: 8
7
+ num_routed_experts: 256
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # GLM-5 is deepseek-arch but n_group=1 / topk_group=1 → the grouped path is a NO-OP (select 1 of 1 group =
13
+ # all experts), i.e. UNGROUPED top-8 with sigmoid(logit)+score_correction_bias for selection, linear
14
+ # renormalization, scaled by 2.5. 256 experts → single 256-face (no combine).
15
+ n_group: 1
16
+ score_func: sigmoid
17
+ routed_scaling_factor: 2.5
18
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sigmoid scores for SELECTION only.
19
+ score_correction_bias: true
20
+
21
+ compute:
22
+ intermediate_size: 2048
23
+ activation_type: SILU
24
+
25
+ reduce:
26
+ shared_expert_scale: 1.0
27
+
28
+ experts:
29
+ expert_mapping: sequential
30
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/glm_47.yaml ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # glm_47
2
+ mesh_shape: [8, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 5120
6
+ select_experts_k: 8
7
+ num_routed_experts: 160
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # GLM-4.7 is deepseek-arch but n_group=1 / topk_group=1 → the grouped path is a NO-OP (select 1 of 1 group
13
+ # = all experts), i.e. UNGROUPED top-8 with sigmoid(logit)+score_correction_bias for selection, linear
14
+ # renormalization, scaled by 2.5. Same gate as GLM-5 but 160 experts → padded to the op's 256-face
15
+ # (phantom experts get a large-negative logit+bias so they rank last in the sigmoid selection).
16
+ n_group: 1
17
+ score_func: sigmoid
18
+ routed_scaling_factor: 2.5
19
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sigmoid scores for SELECTION only.
20
+ score_correction_bias: true
21
+
22
+ compute:
23
+ intermediate_size: 1536
24
+ activation_type: SILU
25
+
26
+ reduce:
27
+ shared_expert_scale: 1.0
28
+
29
+ experts:
30
+ expert_mapping: sequential
31
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/gpt_oss.yaml ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # gpt_oss
2
+ mesh_shape: [8, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 2880
6
+ select_experts_k: 4
7
+ num_routed_experts: 128
8
+ num_shared_experts: 0
9
+ has_bias: true
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # GPT-OSS router (GptOssTopKRouter): logits = Wx + b → top-4 (by Wx+b) → softmax over the selected (Wx+b).
13
+ # TTMoEGate's score_func=softmax is exactly that (rank by logit, then exp-over-selected), scale 1.0, and the
14
+ # router LINEAR bias b is supplied as the gate projection bias (gate_proj_bias below → torch_gate_proj_bias),
15
+ # so logits = Wx + b and b flows into BOTH selection and the softmax weights (matches the reference).
16
+ # 128 experts → padded to the op's 256-face.
17
+ n_group: 1
18
+ score_func: softmax
19
+ routed_scaling_factor: 1.0
20
+ # router LINEAR bias: GPT-OSS's GptOssTopKRouter does F.linear(x, W, b). TTMoEGate adds b via ttnn.linear's
21
+ # bias (torch_gate_proj_bias) — distinct from the deepseek score-correction bias (which is selection-only).
22
+ # (has_bias=true above is the separate EXPERT-MLP bias.)
23
+ gate_proj_bias: true
24
+
25
+ compute:
26
+ intermediate_size: 2880
27
+ activation_type: SWIGLU
28
+
29
+ reduce:
30
+ shared_expert_scale: 1.0
31
+
32
+ experts:
33
+ expert_mapping: sequential
code/models/common/modules/moe/configs/kimi_k25.yaml ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # kimi_k25
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 7168
6
+ select_experts_k: 8
7
+ num_routed_experts: 384
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # Kimi-K2 is deepseek-arch but config.json has n_group=1 / topk_group=1 → the grouped path is a NO-OP
13
+ # (1 group, select 1 of 1 = all experts), i.e. UNGROUPED top-8 with sigmoid(logit)+score_correction_bias
14
+ # for selection, linear renormalization (norm_topk_prob), scaled by 2.827. 384 experts → 2-block combine.
15
+ n_group: 1
16
+ score_func: sigmoid
17
+ routed_scaling_factor: 2.827
18
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sigmoid scores for SELECTION only.
19
+ score_correction_bias: true
20
+
21
+ compute:
22
+ intermediate_size: 2048
23
+ activation_type: SILU
24
+
25
+ reduce:
26
+ shared_expert_scale: 1.0
27
+
28
+ experts:
29
+ expert_mapping: sequential
30
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/ling_1t.yaml ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ling_1t
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 8192
6
+ select_experts_k: 8
7
+ num_routed_experts: 256
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # ling_1t is deepseek-style grouped routing: sigmoid(logit)+score_correction_bias for group selection,
13
+ # linear renormalization of the selected experts, scaled by routed_scaling_factor.
14
+ n_group: 8
15
+ score_func: sigmoid
16
+ routed_scaling_factor: 2.5
17
+ # noaux_tc score-correction bias (e_score_correction_bias): added to the sigmoid scores for SELECTION only.
18
+ score_correction_bias: true
19
+
20
+ compute:
21
+ intermediate_size: 2048
22
+ activation_type: SILU
23
+
24
+ reduce:
25
+ shared_expert_scale: 1.0
26
+
27
+ experts:
28
+ expert_mapping: sequential
29
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/models_table.md ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ | **Model** | **hidden_size** | **moe_intermediate** | **shared_expert_interm** | **n_routed_experts** | **n_shared_experts** | **K (top-k)** | **activation** | **scoring_func** | **topk_method** | **n_group/topk_group** | **scaling_factor** | **expert_bias** | **router_bias** | **first_k_dense** | **num_layers** | **parallel dense** | **base_arch** |
2
+ |---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
3
+ | **GPT-OSS 120B** | 2880 | 2880 | — | 128 | 0 | 4 | custom GELU-gated | softmax | simple | —/— | 1.0 | Yes | Yes | all MoE | 36 | No | GPT-OSS |
4
+ | **DeepSeek V3** | 7168 | 2048 | 2048 | 256 | 1 | 8 | SiLU/SwiGLU | `sigmoid` | noaux_tc | 8/4 | 2.5 | No | correction | 3 | 61 | No | DS V3 |
5
+ | **DS V4 Flash** | 4096 | 2048 | 2048 | 256 | 1 | 6 | SiLU/SwiGLU (swiglu_limit=10, routed only) | `sqrtsoftplus` (code) | noaux_tc; hash L0–2 (code) | —/— | 1.5 | No | correction (code) | 0 | 43 | No | DS V4 |
6
+ | **DS V4 Pro** | 7168 | 3072 | 3072 | 384 | 1 | 6 | SiLU/SwiGLU (swiglu_limit=10, routed only) | `sqrtsoftplus` (code) | noaux_tc; hash L0–2 (code) | —/— | 2.5 | No | correction (code) | 0 | 61 | No | DS V4 |
7
+ | **GLM-5** | 6144 | 2048 | 2048 | 256 | 1 | 8 | SiLU/SwiGLU | `sigmoid` | noaux_tc | 1/1 | 2.5 | No | correction | 3 | 78 | No | DS V3-like |
8
+ | **Kimi K2.5** | 7168 | 2048 | 2048 | 384 | 1 | 8 | SiLU/SwiGLU | `sigmoid` | noaux_tc | 1/1 | 2.827 | No | correction | 1 | 61 | No | DS V3-like |
9
+ | **Ling 1T** | 8192 | 2048 | 2048 | 256 | 1 | 8 | SiLU/SwiGLU | `sigmoid` | group_limited_topk (code) | 8/4 | 2.5 | No | correction (bias-enabled) | 4 | 80 | No | DS V3-like |
10
+ | **GLM-4.7** | 5120 | 1536 | 1536 | 160 | 1 | 8 | SiLU/SwiGLU | `sigmoid` (code) | noaux_tc-style (code) | 1/1 | 2.5 | No | correction (code) | 3 | 92 | No | DS V3-like |
11
+ | **Qwen3 235B** | 4096 | 1536 | — | 128 | 0 (no field) | 8 | SiLU/SwiGLU | softmax (code) | simple top-k (code) | —/— | — | No | No | all MoE | 94 | No | Qwen3 |
12
+ | **Qwen3.5 397B** | 4096 | 1024 | 1024 | 512 | 1 (inferred) | 10 | SiLU/SwiGLU | softmax (code) | simple top-k (code) | —/— | — | No | No | all MoE | 60 | No | Qwen3.5 |
13
+ | **Qwen3.5 35B** | 2048 | 512 | 512 | 256 | 1 (inferred) | 8 | SiLU/SwiGLU | softmax (code) | simple top-k (code) | —/— | — | No | No | all MoE | 40 | No | Qwen3.5 |
14
+ | **Qwen3-Omni Thinker** | 2048 | 768 | — | 128 | 0 (inferred, size=0) | 8 | SiLU/SwiGLU | softmax (code) | simple top-k (code) | —/— | — | No | No | all MoE | 48 | No | Qwen3-Omni |
15
+ | **Qwen3-Omni Talker** | 1024 | 384 | 768 | 128 | 1 (inferred, size=768) | 6 | SiLU/SwiGLU | softmax (code) | simple top-k (code) | —/— | — | No | No | all MoE | 20 | No | Qwen3-Omni |
16
+ | **Mistral Large 3** | 7168 | 4096 | 4096 (from vLLM) | 128 | 1 | 4 | SiLU/SwiGLU (code) | softmax (code) | simple top-k (code) | 1/1 | 1.0 | No | No | 3 | 61 | No | Mistral |
17
+ | **Gemma 4 26B** | 2816 | 704 | — (parallel dense) | 128 | 0 (parallel dense) | 8 | **GELU/SwiGLU** | softmax | simple+per_expert_scale | —/— | per-expert learned | No | No (has learned scale) | all MoE | 30 | **Yes** | Gemma4 |
18
+ | **DS-OCR** | 1280 | 896 | 1792 | 64 | 2 | 6 | SiLU/SwiGLU | softmax (V2 default) | `greedy` | 1/1 | 1.0 (V2 default) | No | No | 1 | 12 | No | DS V2 |
code/models/common/modules/moe/configs/qwen35_35b.yaml ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # qwen35_35b
2
+ mesh_shape: [16, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 2048
6
+ select_experts_k: 8
7
+ num_routed_experts: 256
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8)
12
+ n_group: 1
13
+ score_func: softmax
14
+ routed_scaling_factor: 1.0
15
+ gate_notes: "Qwen3.5 35B: softmax kept POST-topk (softmax-over-selected, the default softmax_position) — mathematically ≡ the model's softmax-over-all→top-8→renorm (global Z cancels, no bias). Set softmax_position: pre for the literal front softmax (numerically-stable exp over all experts)."
16
+
17
+ compute:
18
+ intermediate_size: 512
19
+ activation_type: SILU
20
+
21
+ reduce:
22
+ shared_expert_scale: 1.0
23
+
24
+ experts:
25
+ expert_mapping: sequential
26
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/qwen35_397b.yaml ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # qwen35_397b
2
+ mesh_shape: [16, 8]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 4096
6
+ select_experts_k: 10
7
+ num_routed_experts: 512
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # Qwen3.5 397B = linear → softmax(all) → top-10 → renormalize, no bias, no scale. UNSUPPORTED by the kernel
13
+ # op (top-10 ∉ {4,6,8}, and 512 experts), so TTMoEGate takes the pure-ttnn FALLBACK (matmul → softmax(all)
14
+ # → topk → linear renorm). softmax_position: pre puts the softmax over ALL experts UP FRONT — matching this
15
+ # model's definition directly and giving a numerically-stable exp (vs exp-over-selected raw logits); the
16
+ # linear renorm over the selected makes it mathematically ≡ softmax-over-selected, scale 1.0.
17
+ n_group: 1
18
+ score_func: softmax
19
+ softmax_position: pre
20
+ routed_scaling_factor: 1.0
21
+
22
+ compute:
23
+ intermediate_size: 1024
24
+ activation_type: SILU
25
+
26
+ reduce:
27
+ shared_expert_scale: 1.0
28
+
29
+ experts:
30
+ expert_mapping: sequential
31
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/qwen3_235b.yaml ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # qwen3_235b
2
+ mesh_shape: [8, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 4096
6
+ select_experts_k: 8
7
+ num_routed_experts: 128
8
+ num_shared_experts: 0
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # Qwen3 MoE = softmax over all experts → top-k → renormalize (≡ softmax over selected, the global Z
13
+ # cancels), no bias, no scaling. 128 experts → padded to the op's 256-face in TTMoEGate.forward.
14
+ n_group: 1
15
+ score_func: softmax
16
+ routed_scaling_factor: 1.0
17
+ gate_notes: "Qwen3 235B: softmax kept POST-topk (softmax-over-selected, the default softmax_position) — mathematically ≡ the model's softmax-over-all→top-8→renorm (global Z cancels, no bias). Set softmax_position: pre for the literal front softmax (numerically-stable exp over all experts)."
18
+
19
+ compute:
20
+ intermediate_size: 1536
21
+ activation_type: SILU
22
+
23
+ reduce:
24
+ shared_expert_scale: 1.0
25
+
26
+ experts:
27
+ expert_mapping: sequential
code/models/common/modules/moe/configs/qwen3_omni_talker.yaml ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # qwen3_omni_talker
2
+ mesh_shape: [16, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 1024
6
+ select_experts_k: 6
7
+ num_routed_experts: 128
8
+ num_shared_experts: 1
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # Same as qwen3_omni_thinker but top-6: Qwen3 MoE = softmax over all → top-6 → renormalize (≡ softmax
13
+ # over selected, Z cancels), no bias, no scaling. 128 experts → padded to the op's 256-face in
14
+ # TTMoEGate.forward.
15
+ n_group: 1
16
+ score_func: softmax
17
+ routed_scaling_factor: 1.0
18
+ gate_notes: "Qwen3 Omni Talker: softmax kept POST-topk (softmax-over-selected, the default softmax_position) — mathematically ≡ the model's softmax-over-all→top-6→renorm (global Z cancels, no bias). Set softmax_position: pre for the literal front softmax (numerically-stable exp over all experts)."
19
+
20
+ compute:
21
+ intermediate_size: 384
22
+ activation_type: SILU
23
+
24
+ reduce:
25
+ shared_expert_scale: 1.0
26
+
27
+ experts:
28
+ expert_mapping: sequential
29
+ shared_expert_ids_to_devices: fully_replicated
code/models/common/modules/moe/configs/qwen3_omni_thinker.yaml ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # qwen3_omni_thinker
2
+ mesh_shape: [16, 4]
3
+ cluster_axis: 0
4
+ batch_per_device: 32
5
+ hidden_size: 2048
6
+ select_experts_k: 8
7
+ num_routed_experts: 128
8
+ num_shared_experts: 0
9
+ has_bias: false
10
+
11
+ # gate (routing front-end): n_group 1 = generalized op (ungrouped) / 8 = deepseek op (grouped top-8).
12
+ # Qwen3 MoE = softmax over all experts → top-8 → renormalize (≡ softmax over selected, the global Z
13
+ # cancels), no bias, no scaling. 128 experts → padded to the op's 256-face in TTMoEGate.forward.
14
+ n_group: 1
15
+ score_func: softmax
16
+ routed_scaling_factor: 1.0
17
+ gate_notes: "Qwen3 Omni Thinker: softmax kept POST-topk (softmax-over-selected, the default softmax_position) — mathematically ≡ the model's softmax-over-all→top-8→renorm (global Z cancels, no bias). Set softmax_position: pre for the literal front softmax (numerically-stable exp over all experts)."
18
+
19
+ compute:
20
+ intermediate_size: 768
21
+ activation_type: SILU
22
+
23
+ reduce:
24
+ shared_expert_scale: 1.0
25
+
26
+ experts:
27
+ expert_mapping: sequential
code/models/common/tests/demos/cleanup_utils.py ADDED
@@ -0,0 +1,113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import gc
5
+
6
+ import ttnn
7
+ from models.common.modules.lazy_weight import LazyWeight
8
+
9
+
10
+ def cleanup_ttnn_value(value):
11
+ if value is None:
12
+ return
13
+
14
+ if isinstance(value, ttnn.Tensor):
15
+ ttnn.deallocate(value)
16
+ return
17
+
18
+ if isinstance(value, dict):
19
+ for nested_value in value.values():
20
+ cleanup_ttnn_value(nested_value)
21
+ return
22
+
23
+ if isinstance(value, (list, tuple, set)):
24
+ for nested_value in value:
25
+ cleanup_ttnn_value(nested_value)
26
+
27
+
28
+ def cleanup_object_graph(obj, seen=None):
29
+ if obj is None:
30
+ return
31
+ if seen is None:
32
+ seen = set()
33
+
34
+ obj_id = id(obj)
35
+ if obj_id in seen:
36
+ return
37
+ seen.add(obj_id)
38
+
39
+ if isinstance(obj, ttnn.Tensor):
40
+ cleanup_ttnn_value(obj)
41
+ return
42
+
43
+ if isinstance(obj, LazyWeight):
44
+ if obj._value is not None:
45
+ cleanup_ttnn_value(obj._value)
46
+ obj._value = None
47
+ return
48
+
49
+ if isinstance(obj, dict):
50
+ for value in obj.values():
51
+ cleanup_object_graph(value, seen)
52
+ return
53
+
54
+ if isinstance(obj, (list, tuple, set)):
55
+ for value in obj:
56
+ cleanup_object_graph(value, seen)
57
+ return
58
+
59
+ state = getattr(obj, "__dict__", None)
60
+ if state is None:
61
+ return
62
+
63
+ for name, value in list(state.items()):
64
+ cleanup_object_graph(value, seen)
65
+ if isinstance(value, ttnn.Tensor):
66
+ setattr(obj, name, None)
67
+
68
+ if hasattr(obj, "_device_weights_loaded"):
69
+ obj._device_weights_loaded = False
70
+
71
+
72
+ def cleanup_model_case(model, mesh_device):
73
+ ttnn.synchronize_device(mesh_device)
74
+ if model is not None:
75
+ cleanup_object_graph(model)
76
+ ttnn.synchronize_device(mesh_device)
77
+ gc.collect()
78
+
79
+
80
+ def cleanup_dp_model_case(group, lanes, models, parent_mesh, submeshes):
81
+ """Release one carved DP case and restore command-queue ownership to its parent."""
82
+
83
+ failures = []
84
+
85
+ def run(action):
86
+ try:
87
+ action()
88
+ except Exception as error:
89
+ failures.append(error)
90
+
91
+ if group is not None:
92
+ run(group.cleanup)
93
+ else:
94
+ for lane in lanes:
95
+ run(lane.cleanup)
96
+
97
+ for model, submesh in models:
98
+ run(lambda model=model, submesh=submesh: cleanup_model_case(model, submesh))
99
+
100
+ # A carved child owns a distinct MeshCommandQueue over its parent's physical
101
+ # devices. Drain every live child through the parent before closing the child
102
+ # handles and returning command-queue ownership to the fixture-owned parent.
103
+ run(parent_mesh.quiesce_devices)
104
+ for submesh in submeshes:
105
+ run(lambda submesh=submesh: ttnn.close_mesh_device(submesh))
106
+
107
+ if failures:
108
+ primary = failures[0]
109
+ add_note = getattr(primary, "add_note", None)
110
+ if add_note is not None:
111
+ for failure in failures[1:]:
112
+ add_note(f"additional DP cleanup failure: {type(failure).__name__}: {failure}")
113
+ raise primary
code/models/common/tests/demos/run_helpers.py ADDED
@@ -0,0 +1,1071 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Reusable teacher-forcing and benchmark run helpers for demos and tests."""
5
+
6
+ from __future__ import annotations
7
+
8
+ import json
9
+ import os
10
+ import time
11
+ from dataclasses import dataclass, field
12
+ from pathlib import Path
13
+
14
+ import torch
15
+ from loguru import logger
16
+
17
+ import ttnn
18
+
19
+ _SAME_SAMPLING_PARAMS = object()
20
+
21
+
22
+ def make_contiguous_page_table(batch_size: int, max_seq_len: int, block_size: int = 32) -> torch.Tensor:
23
+ """Create a contiguous demo page table with one disjoint block range per user."""
24
+ if min(batch_size, max_seq_len, block_size) <= 0:
25
+ raise ValueError("page-table dimensions must be positive")
26
+ blocks_per_user = (max_seq_len + block_size - 1) // block_size
27
+ return torch.arange(batch_size * blocks_per_user, dtype=torch.int32).reshape(batch_size, blocks_per_user)
28
+
29
+
30
+ @dataclass
31
+ class TeacherForceResult:
32
+ """Result from a teacher-forcing evaluation run."""
33
+
34
+ predicted_tokens: list[int]
35
+ predicted_tokens_per_user: list[list[int]]
36
+ reference_top5: torch.Tensor
37
+ prefill_time_s: float = 0.0
38
+ compile_decode_time_s: float = 0.0
39
+ decode_times_s: list[float] = field(default_factory=list)
40
+ batch_size: int = 1
41
+ prefill_len: int = 0
42
+
43
+ def top1_accuracy(self) -> float:
44
+ matches = sum(
45
+ 1 for i, prediction in enumerate(self.predicted_tokens) if self.reference_top5[i, 0].item() == prediction
46
+ )
47
+ return matches / len(self.predicted_tokens)
48
+
49
+ def top5_accuracy(self) -> float:
50
+ matches = sum(
51
+ 1 for i, prediction in enumerate(self.predicted_tokens) if prediction in self.reference_top5[i, :]
52
+ )
53
+ return matches / len(self.predicted_tokens)
54
+
55
+ @property
56
+ def ttft_ms(self) -> float:
57
+ return self.prefill_time_s / self.batch_size * 1000 if self.batch_size else 0.0
58
+
59
+ @property
60
+ def prefill_time_to_token_s(self) -> float:
61
+ return self.prefill_time_s / self.batch_size if self.batch_size else 0.0
62
+
63
+ @property
64
+ def prefill_tok_s(self) -> float:
65
+ return (self.batch_size * self.prefill_len) / self.prefill_time_s if self.prefill_time_s > 0 else 0.0
66
+
67
+ @property
68
+ def decode_tok_s_u(self) -> float:
69
+ return len(self.decode_times_s) / sum(self.decode_times_s) if self.decode_times_s else 0.0
70
+
71
+ @property
72
+ def decode_tok_s(self) -> float:
73
+ return self.decode_tok_s_u * self.batch_size
74
+
75
+
76
+ @dataclass
77
+ class PerfBenchmarkResult:
78
+ """Result from a performance benchmark run."""
79
+
80
+ prefill_time_s: float
81
+ compile_decode_time_s: float
82
+ decode_times_s: list[float]
83
+ batch_size: int
84
+ num_decode_tokens: int
85
+ generated_token_ids: list[list[int]]
86
+ decode_iteration_times_s: list[float] = field(default_factory=list)
87
+ argmax_top2_margins: list[list[float]] | None = None
88
+
89
+ @property
90
+ def ttft_ms(self) -> float:
91
+ """TTTv1-style average time to first token per user."""
92
+ return self.prefill_time_s / self.batch_size * 1000
93
+
94
+ @property
95
+ def tok_s_u(self) -> float:
96
+ """Tokens per second per user during steady-state decode."""
97
+ if not self.decode_times_s:
98
+ return 0.0
99
+ return len(self.decode_times_s) / sum(self.decode_times_s)
100
+
101
+ @property
102
+ def tok_s(self) -> float:
103
+ """Total decode throughput."""
104
+ return self.tok_s_u * self.batch_size
105
+
106
+ @property
107
+ def decode_latency_mean_ms(self) -> float:
108
+ if not self.decode_times_s:
109
+ return 0.0
110
+ return sum(self.decode_times_s) / len(self.decode_times_s) * 1000
111
+
112
+ def meets_target(self, expected: dict, tolerance: float = 0.05) -> dict[str, bool]:
113
+ """Check benchmark metrics against the expected thresholds."""
114
+ return {
115
+ "tok_s_u": self.tok_s_u >= expected["tok_s_u"] * (1 - tolerance),
116
+ "ttft_ms": self.ttft_ms <= expected["ttft_ms"] * (1 + tolerance),
117
+ }
118
+
119
+
120
+ def _compile_prefill_and_decode(
121
+ execution_target,
122
+ *,
123
+ prefill_tokens: torch.Tensor,
124
+ prefill_page_table: torch.Tensor,
125
+ kv_cache=None,
126
+ prompt_lens: torch.Tensor | None = None,
127
+ empty_slots: list[int] | None = None,
128
+ start_pos: torch.Tensor | None = None,
129
+ sampling_params=None,
130
+ prefill_sampling_params=_SAME_SAMPLING_PARAMS,
131
+ decode_tokens: torch.Tensor | None = None,
132
+ decode_start_pos: torch.Tensor | None = None,
133
+ decode_page_table: torch.Tensor | None = None,
134
+ ) -> None:
135
+ """Compile the concrete prefill and decode cases through the public target surface."""
136
+ assert prefill_tokens.dim() == 2, f"prefill_tokens must be [batch_size, seq_len], got {prefill_tokens.dim()}D"
137
+ assert (
138
+ prefill_page_table.dim() == 2
139
+ ), f"prefill_page_table must be [batch_size, max_blocks], got {prefill_page_table.dim()}D"
140
+
141
+ batch_size = prefill_tokens.shape[0]
142
+ if decode_tokens is None:
143
+ decode_tokens = torch.zeros(batch_size, dtype=torch.long, device=prefill_tokens.device)
144
+ if decode_start_pos is None:
145
+ decode_start_pos = torch.full(
146
+ (decode_tokens.shape[0],),
147
+ prefill_tokens.shape[-1],
148
+ dtype=torch.long,
149
+ device=prefill_tokens.device,
150
+ )
151
+ if decode_page_table is None:
152
+ decode_page_table = prefill_page_table
153
+
154
+ if prefill_sampling_params is _SAME_SAMPLING_PARAMS:
155
+ prefill_sampling_params = sampling_params
156
+
157
+ if sampling_params is not None:
158
+ execution_target.compile_decode(
159
+ tokens=decode_tokens,
160
+ start_pos=decode_start_pos,
161
+ page_table=decode_page_table,
162
+ kv_cache=kv_cache,
163
+ sampling_params=sampling_params,
164
+ )
165
+ execution_target.compile_prefill(
166
+ tokens=prefill_tokens,
167
+ page_table=prefill_page_table,
168
+ kv_cache=kv_cache,
169
+ prompt_lens=prompt_lens,
170
+ empty_slots=empty_slots,
171
+ start_pos=start_pos,
172
+ sampling_params=prefill_sampling_params,
173
+ )
174
+ return
175
+
176
+ execution_target.compile_prefill(
177
+ tokens=prefill_tokens,
178
+ page_table=prefill_page_table,
179
+ kv_cache=kv_cache,
180
+ prompt_lens=prompt_lens,
181
+ empty_slots=empty_slots,
182
+ start_pos=start_pos,
183
+ sampling_params=None,
184
+ )
185
+ execution_target.compile_decode(
186
+ tokens=decode_tokens,
187
+ start_pos=decode_start_pos,
188
+ page_table=decode_page_table,
189
+ kv_cache=kv_cache,
190
+ sampling_params=None,
191
+ )
192
+
193
+
194
+ def _profiler_start(profiler, name: str) -> None:
195
+ if profiler is not None:
196
+ profiler.start(name)
197
+
198
+
199
+ def _profiler_end(profiler, name: str) -> None:
200
+ if profiler is not None:
201
+ profiler.end(name)
202
+
203
+
204
+ def run_teacher_forcing(
205
+ executor,
206
+ *,
207
+ prompt_tokens: torch.Tensor,
208
+ reference_tokens: torch.Tensor,
209
+ top5_tokens: torch.Tensor,
210
+ kv_cache: list,
211
+ page_table: torch.Tensor,
212
+ max_batch_size: int = 1,
213
+ profiler=None,
214
+ ) -> TeacherForceResult:
215
+ """Run teacher-forcing accuracy measurement against an execution target."""
216
+ execution_target = executor
217
+ batch_size = prompt_tokens.shape[0]
218
+ assert (
219
+ batch_size == max_batch_size
220
+ ), f"Teacher forcing expects active batch to match max_batch_size, got {batch_size} vs {max_batch_size}"
221
+ prompt_len = prompt_tokens.shape[-1]
222
+ num_target = len(reference_tokens) - prompt_len
223
+ prompt_lens = torch.tensor([prompt_len] * batch_size)
224
+ empty_slots = list(range(batch_size))
225
+
226
+ _compile_prefill_and_decode(
227
+ execution_target,
228
+ prefill_tokens=prompt_tokens,
229
+ prefill_page_table=page_table,
230
+ kv_cache=kv_cache,
231
+ prompt_lens=prompt_lens,
232
+ empty_slots=empty_slots,
233
+ )
234
+
235
+ logger.info(f"Teacher forcing: prefilling {prompt_len} tokens with batch={batch_size}")
236
+ _profiler_start(profiler, "inference_prefill")
237
+ try:
238
+ start_time = time.perf_counter()
239
+ prefill_output = execution_target.prefill_forward(
240
+ prompt_tokens,
241
+ page_table=page_table,
242
+ kv_cache=kv_cache,
243
+ prompt_lens=prompt_lens,
244
+ empty_slots=empty_slots,
245
+ )
246
+ _synchronize_target(execution_target)
247
+ prefill_time_s = time.perf_counter() - start_time
248
+ finally:
249
+ _profiler_end(profiler, "inference_prefill")
250
+ first_tokens = torch.argmax(prefill_output, dim=-1).view(-1).tolist()
251
+ predicted_tokens_per_user = [[int(token)] for token in first_tokens]
252
+
253
+ logger.info(f"Teacher forcing: decoding {num_target - 1} tokens")
254
+ compile_decode_time_s = 0.0
255
+ decode_times_s = []
256
+ _profiler_start(profiler, "inference_decode")
257
+ try:
258
+ for step in range(1, num_target):
259
+ ground_truth_token = reference_tokens[prompt_len + step - 1]
260
+ decode_token = torch.full((batch_size,), ground_truth_token, dtype=torch.long)
261
+ current_pos = torch.full((batch_size,), prompt_len + step - 1, dtype=torch.long)
262
+ start_time = time.perf_counter()
263
+ logits, _ = execution_target.decode_forward(
264
+ decode_token,
265
+ current_pos,
266
+ page_table=page_table,
267
+ kv_cache=kv_cache,
268
+ read_from_device=True,
269
+ )
270
+ elapsed = time.perf_counter() - start_time
271
+ if step == 1:
272
+ compile_decode_time_s = elapsed
273
+ else:
274
+ decode_times_s.append(elapsed)
275
+
276
+ next_tokens = torch.argmax(logits[:, -1, :], dim=-1).view(-1).tolist()
277
+ for user_id, token in enumerate(next_tokens):
278
+ predicted_tokens_per_user[user_id].append(int(token))
279
+ finally:
280
+ _profiler_end(profiler, "inference_decode")
281
+
282
+ return TeacherForceResult(
283
+ predicted_tokens=predicted_tokens_per_user[0],
284
+ predicted_tokens_per_user=predicted_tokens_per_user,
285
+ reference_top5=top5_tokens[:num_target],
286
+ prefill_time_s=prefill_time_s,
287
+ compile_decode_time_s=compile_decode_time_s,
288
+ decode_times_s=decode_times_s,
289
+ batch_size=batch_size,
290
+ prefill_len=prompt_len,
291
+ )
292
+
293
+
294
+ def _split_output(output):
295
+ return output if isinstance(output, tuple) else (output, None)
296
+
297
+
298
+ def _target_mesh_device(execution_target):
299
+ return getattr(execution_target, "mesh_device", None)
300
+
301
+
302
+ def _target_cluster_shape(execution_target):
303
+ cluster_shape = getattr(execution_target, "cluster_shape", None)
304
+ if cluster_shape is not None:
305
+ return list(cluster_shape)
306
+ mesh_device = _target_mesh_device(execution_target)
307
+ return list(mesh_device.shape) if mesh_device is not None else [1, 1]
308
+
309
+
310
+ def _synchronize_target(execution_target):
311
+ mesh_devices = getattr(execution_target, "mesh_devices", None)
312
+ if mesh_devices is not None:
313
+ for mesh_device in mesh_devices:
314
+ if mesh_device is not None:
315
+ ttnn.synchronize_device(mesh_device)
316
+ return
317
+ mesh_device = _target_mesh_device(execution_target)
318
+ if mesh_device is not None:
319
+ ttnn.synchronize_device(mesh_device)
320
+
321
+
322
+ def _concat_host_output(output, cluster_shape):
323
+ output_tensors = [ttnn.to_torch(tensor) for tensor in ttnn.get_device_tensors(output)]
324
+ _, columns = cluster_shape
325
+ mesh_rows = [output_tensors[i : i + columns] for i in range(0, len(output_tensors), columns)]
326
+ return torch.cat([torch.cat(row, dim=-1) for row in mesh_rows], dim=1)
327
+
328
+
329
+ def _process_legacy_sampled_tokens(output, batch_size, cluster_shape):
330
+ torch_output = _concat_host_output(output, cluster_shape)
331
+ if torch_output.ndim >= 4:
332
+ if torch_output.shape[2] >= batch_size:
333
+ return torch_output[0, 0, :batch_size, 0]
334
+ if torch_output.shape[3] >= batch_size:
335
+ return torch_output[0, 0, 0, :batch_size]
336
+ return torch_output.reshape(-1)[:batch_size]
337
+
338
+
339
+ def _to_host(value, *, blocking):
340
+ if value is None:
341
+ return None
342
+ if isinstance(value, torch.Tensor):
343
+ return value.cpu()
344
+ try:
345
+ return value.cpu(blocking=blocking)
346
+ except TypeError:
347
+ return value.cpu()
348
+
349
+
350
+ def _submit_decode_read(execution_target, decode_output):
351
+ read_decode_output = getattr(execution_target, "read_decode_output", None)
352
+ if callable(read_decode_output):
353
+ return read_decode_output(decode_output, async_read=True)
354
+
355
+ output, log_probs = _split_output(decode_output)
356
+ host_output = (_to_host(output, blocking=False), _to_host(log_probs, blocking=False))
357
+ return host_output, [ttnn.record_event(_target_mesh_device(execution_target), 0)]
358
+
359
+
360
+ def _synchronize_read_events(events):
361
+ if events is None:
362
+ return
363
+ if not isinstance(events, (list, tuple, set)):
364
+ events = [events]
365
+ for event in events:
366
+ ttnn.event_synchronize(event)
367
+
368
+
369
+ def _consume_sampled_output(
370
+ execution_target,
371
+ host_output,
372
+ batch_size,
373
+ cluster_shape,
374
+ generated_token_ids,
375
+ *,
376
+ process_host_output,
377
+ ):
378
+ process_decode_output_host = (
379
+ getattr(execution_target, "process_decode_output_host", None) if process_host_output else None
380
+ )
381
+ if callable(process_decode_output_host):
382
+ tokens, _ = process_decode_output_host(host_output, is_tokens=True)
383
+ else:
384
+ tokens, _ = _split_output(host_output)
385
+ if isinstance(tokens, torch.Tensor):
386
+ tokens = tokens.view(-1)[:batch_size].detach().cpu()
387
+ else:
388
+ tokens = _process_legacy_sampled_tokens(tokens, batch_size, cluster_shape)
389
+ tokens = tokens.view(-1)[:batch_size].detach().cpu()
390
+
391
+ for user_id, token in enumerate(tokens.tolist()):
392
+ generated_token_ids[user_id].append(int(token))
393
+
394
+
395
+ def _host_argmax_with_margins(logits: torch.Tensor, batch_size: int) -> tuple[torch.Tensor, torch.Tensor]:
396
+ """Return host argmax and top1-minus-top2 margins for eval diagnostics."""
397
+
398
+ rows = logits[:, -1, :]
399
+ top2 = torch.topk(rows.float(), k=2, dim=-1)
400
+ tokens = top2.indices[:, 0].view(-1)[:batch_size].detach().cpu()
401
+ margins = (top2.values[:, 0] - top2.values[:, 1]).view(-1)[:batch_size].detach().cpu()
402
+ return tokens, margins
403
+
404
+
405
+ def run_perf_benchmark(
406
+ executor,
407
+ *,
408
+ tokens: torch.Tensor,
409
+ kv_cache: list,
410
+ page_table: torch.Tensor,
411
+ num_decode_tokens: int = 128,
412
+ max_batch_size: int = 1,
413
+ prompt_lens: torch.Tensor | None = None,
414
+ start_pos: list[int] | None = None,
415
+ sampling_params=None,
416
+ prefill_sampling_params=_SAME_SAMPLING_PARAMS,
417
+ pipeline_readback: bool = False,
418
+ profiler=None,
419
+ collect_argmax_diagnostics: bool = False,
420
+ ) -> PerfBenchmarkResult:
421
+ """Run the timed prefill and decode loop against a public execution target."""
422
+ execution_target = executor
423
+ mesh_device = _target_mesh_device(execution_target)
424
+ has_public_readback = callable(getattr(execution_target, "read_decode_output", None))
425
+ has_legacy_readback = mesh_device is not None and hasattr(ttnn, "record_event")
426
+ can_pipeline_readback = (
427
+ sampling_params is not None
428
+ and pipeline_readback
429
+ and hasattr(ttnn, "event_synchronize")
430
+ and (has_public_readback or has_legacy_readback)
431
+ )
432
+ if sampling_params is not None and pipeline_readback and not can_pipeline_readback:
433
+ logger.warning("PIPELINE_READBACK requested, but this execution target does not expose async readback")
434
+
435
+ batch_size, prompt_len = tokens.shape
436
+ max_batch_size = max(max_batch_size, batch_size)
437
+ cluster_shape = _target_cluster_shape(execution_target)
438
+ prompt_lens = prompt_lens if prompt_lens is not None else torch.tensor([prompt_len] * batch_size)
439
+ if prefill_sampling_params is _SAME_SAMPLING_PARAMS:
440
+ prefill_sampling_params = sampling_params
441
+ prefill_kwargs = dict(
442
+ # Decode always carries the full lane-capacity table. Prefill only owns
443
+ # the active request rows, and its public contract requires matching
444
+ # token/page-table batch dimensions.
445
+ page_table=page_table[:batch_size],
446
+ kv_cache=kv_cache,
447
+ prompt_lens=prompt_lens,
448
+ empty_slots=list(range(batch_size)),
449
+ start_pos=start_pos,
450
+ )
451
+
452
+ compile_tokens = torch.zeros(max_batch_size, prompt_len, dtype=tokens.dtype)
453
+ compile_tokens[:batch_size] = tokens
454
+ compile_prompt_lens = torch.zeros(max_batch_size, dtype=prompt_lens.dtype)
455
+ compile_prompt_lens[:batch_size] = prompt_lens
456
+ compile_page_table = page_table
457
+ compile_empty_slots = list(range(batch_size))
458
+ if batch_size < max_batch_size:
459
+ # A partial-cardinality diagnostic still decodes on the model's full
460
+ # lane-capacity program, but prefill owns only its active requests.
461
+ compile_tokens = tokens
462
+ compile_prompt_lens = prompt_lens
463
+ compile_page_table = page_table[:batch_size]
464
+ _compile_prefill_and_decode(
465
+ execution_target,
466
+ prefill_tokens=compile_tokens,
467
+ prefill_page_table=compile_page_table,
468
+ kv_cache=kv_cache,
469
+ prompt_lens=compile_prompt_lens,
470
+ empty_slots=compile_empty_slots,
471
+ start_pos=start_pos,
472
+ sampling_params=sampling_params,
473
+ prefill_sampling_params=prefill_sampling_params,
474
+ decode_tokens=torch.zeros(max_batch_size, dtype=torch.long),
475
+ decode_start_pos=torch.full((max_batch_size,), prompt_len, dtype=torch.long),
476
+ decode_page_table=page_table,
477
+ )
478
+
479
+ _profiler_start(profiler, "inference_prefill")
480
+ try:
481
+ start_time = time.perf_counter()
482
+ prefill_output = execution_target.prefill_forward(
483
+ tokens,
484
+ **prefill_kwargs,
485
+ sampling_params=prefill_sampling_params,
486
+ )
487
+ _synchronize_target(execution_target)
488
+ prefill_time = time.perf_counter() - start_time
489
+ finally:
490
+ _profiler_end(profiler, "inference_prefill")
491
+
492
+ first_margins = None
493
+ if collect_argmax_diagnostics and sampling_params is None:
494
+ first_token, first_margins = _host_argmax_with_margins(prefill_output, batch_size)
495
+ else:
496
+ first_token = prefill_output[0] if isinstance(prefill_output, tuple) else torch.argmax(prefill_output, dim=-1)
497
+ first_token = first_token.view(-1)[:batch_size].detach().cpu()
498
+ generated_token_ids = [[int(token)] for token in first_token.tolist()]
499
+ argmax_top2_margins = [[float(margin)] for margin in first_margins.tolist()] if first_margins is not None else None
500
+
501
+ current_tokens = torch.zeros(max_batch_size, dtype=torch.long)
502
+ current_tokens[:batch_size] = first_token
503
+ current_pos = torch.full((max_batch_size,), -1, dtype=torch.long)
504
+ current_pos[:batch_size] = prompt_lens[:batch_size]
505
+
506
+ compile_time = None
507
+ decode_times = []
508
+ decode_iteration_times = []
509
+ sampled_decode_start = None
510
+ pending_host_output = None
511
+ pending_read_events = None
512
+
513
+ _profiler_start(profiler, "inference_decode")
514
+ try:
515
+ for iteration in range(num_decode_tokens):
516
+ start_time = time.perf_counter()
517
+ read_from_device = sampling_params is None or not can_pipeline_readback
518
+ if sampling_params is not None and iteration == 1:
519
+ sampled_decode_start = start_time
520
+
521
+ decode_output = execution_target.decode_forward(
522
+ current_tokens,
523
+ current_pos,
524
+ page_table=page_table,
525
+ kv_cache=kv_cache,
526
+ read_from_device=read_from_device,
527
+ sampling_params=sampling_params,
528
+ reset_batch=iteration == 0,
529
+ )
530
+ output, _ = _split_output(decode_output)
531
+
532
+ completed_host_output = None
533
+ if can_pipeline_readback:
534
+ host_output, read_events = _submit_decode_read(execution_target, decode_output)
535
+ if pending_read_events is not None:
536
+ _synchronize_read_events(pending_read_events)
537
+ completed_host_output = pending_host_output
538
+ pending_host_output = host_output
539
+ pending_read_events = read_events
540
+
541
+ if sampling_params is not None and iteration == 0:
542
+ _synchronize_target(execution_target)
543
+ elapsed = time.perf_counter() - start_time
544
+
545
+ if iteration == 0:
546
+ compile_time = elapsed
547
+ else:
548
+ decode_iteration_times.append(elapsed)
549
+ if sampling_params is None or can_pipeline_readback:
550
+ decode_times.append(elapsed)
551
+
552
+ if completed_host_output is not None:
553
+ _consume_sampled_output(
554
+ execution_target,
555
+ completed_host_output,
556
+ batch_size,
557
+ cluster_shape,
558
+ generated_token_ids,
559
+ process_host_output=True,
560
+ )
561
+
562
+ if sampling_params is None:
563
+ if isinstance(output, torch.Tensor) and output.dim() >= 2:
564
+ if collect_argmax_diagnostics:
565
+ next_token, next_margins = _host_argmax_with_margins(output, batch_size)
566
+ else:
567
+ next_token = torch.argmax(output[:, -1, :], dim=-1)
568
+ else:
569
+ if collect_argmax_diagnostics:
570
+ raise TypeError("argmax margin diagnostics require host decode logits")
571
+ next_token = output
572
+ next_token = next_token.view(-1)[:batch_size].detach().cpu()
573
+ for user_id, token in enumerate(next_token.tolist()):
574
+ generated_token_ids[user_id].append(int(token))
575
+ if argmax_top2_margins is not None:
576
+ argmax_top2_margins[user_id].append(float(next_margins[user_id]))
577
+ current_tokens[:batch_size] = next_token
578
+ elif not can_pipeline_readback:
579
+ _consume_sampled_output(
580
+ execution_target,
581
+ decode_output,
582
+ batch_size,
583
+ cluster_shape,
584
+ generated_token_ids,
585
+ process_host_output=False,
586
+ )
587
+ current_pos[:batch_size] += 1
588
+ finally:
589
+ _profiler_end(profiler, "inference_decode")
590
+
591
+ if sampling_params is not None:
592
+ if pending_read_events is not None:
593
+ _synchronize_read_events(pending_read_events)
594
+ _consume_sampled_output(
595
+ execution_target,
596
+ pending_host_output,
597
+ batch_size,
598
+ cluster_shape,
599
+ generated_token_ids,
600
+ process_host_output=True,
601
+ )
602
+ if sampled_decode_start is not None and not can_pipeline_readback:
603
+ _synchronize_target(execution_target)
604
+ sampled_decode_time = time.perf_counter() - sampled_decode_start
605
+ decode_times = [sampled_decode_time / (num_decode_tokens - 1)] * (num_decode_tokens - 1)
606
+
607
+ return PerfBenchmarkResult(
608
+ prefill_time_s=prefill_time,
609
+ compile_decode_time_s=compile_time or 0.0,
610
+ decode_times_s=decode_times,
611
+ batch_size=batch_size,
612
+ num_decode_tokens=num_decode_tokens,
613
+ generated_token_ids=generated_token_ids,
614
+ decode_iteration_times_s=decode_iteration_times,
615
+ argmax_top2_margins=argmax_top2_margins,
616
+ )
617
+
618
+
619
+ def _add_token_ids(target: set[int], value) -> None:
620
+ if value is None or isinstance(value, bool):
621
+ return
622
+ if isinstance(value, int):
623
+ target.add(int(value))
624
+ return
625
+ if isinstance(value, (list, tuple, set)):
626
+ for item in value:
627
+ _add_token_ids(target, item)
628
+
629
+
630
+ def _stop_token_ids(tokenizer) -> set[int]:
631
+ stop_ids: set[int] = set()
632
+ _add_token_ids(stop_ids, getattr(tokenizer, "eos_token_id", None))
633
+ _add_token_ids(stop_ids, getattr(tokenizer, "stop_tokens", None))
634
+ convert_tokens_to_ids = getattr(tokenizer, "convert_tokens_to_ids", None)
635
+ if callable(convert_tokens_to_ids):
636
+ eot_id = convert_tokens_to_ids("<|eot_id|>")
637
+ if isinstance(eot_id, int) and eot_id >= 0:
638
+ stop_ids.add(eot_id)
639
+ return stop_ids
640
+
641
+
642
+ def _truncate_at_stop(token_ids: list[int], stop_ids: set[int]) -> list[int]:
643
+ for index, token_id in enumerate(token_ids):
644
+ if token_id in stop_ids:
645
+ return token_ids[:index]
646
+ return token_ids
647
+
648
+
649
+ def assert_no_special_tokens(
650
+ generated_token_ids: list[list[int]],
651
+ tokenizer,
652
+ *,
653
+ case_name: str = "",
654
+ is_ci_env: bool | None = None,
655
+ ) -> None:
656
+ """Warn locally; fail under CI or ``TT_DEMO_STRICT_SPECIAL_TOKENS=1``."""
657
+ if is_ci_env is None:
658
+ is_ci_env = os.environ.get("CI") == "true" or os.environ.get("TT_DEMO_STRICT_SPECIAL_TOKENS") == "1"
659
+
660
+ special_ids = set(getattr(tokenizer, "all_special_ids", []) or [])
661
+ stop_ids = _stop_token_ids(tokenizer)
662
+ offending_users = 0
663
+ for token_ids in generated_token_ids:
664
+ output_before_stop = _truncate_at_stop(list(token_ids), stop_ids)
665
+ if any(token_id in special_ids for token_id in output_before_stop):
666
+ offending_users += 1
667
+
668
+ if offending_users == 0:
669
+ return
670
+
671
+ prefix = f"[{case_name}] " if case_name else ""
672
+ message = f"{prefix}model produced special tokens ({offending_users}/{len(generated_token_ids)} users)"
673
+ logger.warning(message)
674
+ if is_ci_env:
675
+ raise AssertionError(message)
676
+
677
+
678
+ def load_eval_repeat_prompts_batch32() -> list[str]:
679
+ """The 32 numeric sequence-continuation prompts TTTv1's ci-eval-32 uses (parity)."""
680
+ path = Path("models/tt_transformers/demo/sample_prompts/eval_repeat_prompts_batch32.json")
681
+ with open(path) as f:
682
+ data = json.load(f)
683
+ return [entry["prompt"] for entry in data]
684
+
685
+
686
+ def rotate_prompts(all_prompts: list[str], repeat: int) -> list[str]:
687
+ """Rotate the prompt->slot assignment by ``repeat``: slot j holds prompt (j+repeat)%N."""
688
+ n = len(all_prompts)
689
+ return [all_prompts[(j + repeat) % n] for j in range(n)]
690
+
691
+
692
+ def eval_page_table_for_repeat(page_table: torch.Tensor, repeat: int, *, mode: str) -> torch.Tensor:
693
+ """Select physical KV allocation for one rotated eval repeat.
694
+
695
+ ``slot-stable`` is the TTTv1 acceptance geometry: page-table row ``j``
696
+ remains at decode slot ``j`` while prompts rotate. ``prompt-stable`` is a
697
+ diagnostic A/B: rows rotate with prompts, so a prompt retains its original
698
+ physical KV blocks even as it moves to another decode slot.
699
+ """
700
+
701
+ if mode == "slot-stable":
702
+ return page_table
703
+ if mode == "prompt-stable":
704
+ return torch.roll(page_table, shifts=-repeat, dims=0)
705
+ raise ValueError(f"unsupported eval page-table mode {mode!r}; use 'slot-stable' or 'prompt-stable'")
706
+
707
+
708
+ def eval_decode_trace_mode(mode: str) -> str:
709
+ """Resolve the eval decode execution A/B without changing its default gate."""
710
+
711
+ if mode == "traced":
712
+ return "decode_only"
713
+ if mode == "eager":
714
+ return "none"
715
+ raise ValueError(f"unsupported eval decode mode {mode!r}; use 'traced' or 'eager'")
716
+
717
+
718
+ def require_canonical_eval_modes_in_ci(environ) -> None:
719
+ """Prevent diagnostic A/B knobs from replacing a canonical CI gate."""
720
+
721
+ if environ.get("CI") != "true":
722
+ return
723
+ noncanonical = []
724
+ if environ.get("EVAL_DECODE_MODE", "traced") != "traced":
725
+ noncanonical.append("EVAL_DECODE_MODE")
726
+ if environ.get("EVAL_PAGE_TABLE_MODE", "slot-stable") != "slot-stable":
727
+ noncanonical.append("EVAL_PAGE_TABLE_MODE")
728
+ noncanonical.extend(name for name in ("EVAL_IDENTICAL_PROMPT_INDEX", "EVAL_ACTIVE_BATCH_SIZE") if name in environ)
729
+ if noncanonical:
730
+ raise RuntimeError("diagnostic eval modes cannot replace the canonical CI gate: " + ", ".join(noncanonical))
731
+
732
+
733
+ def truncate_at_stop(ids: list[int], stop_ids: set[int]) -> list[int]:
734
+ """Prefix of ``ids`` up to (excluding) the first id in ``stop_ids``."""
735
+ out: list[int] = []
736
+ for t in ids:
737
+ if t in stop_ids:
738
+ break
739
+ out.append(t)
740
+ return out
741
+
742
+
743
+ def hf_stop_ids(tokenizer, hf_model_id: str | None = None) -> set[int]:
744
+ """Best-effort stop-token id set for an HF ``AutoTokenizer``.
745
+
746
+ Raw HF tokenizers have no ``.stop_tokens`` (that only exists on the TTTv1 wrapped
747
+ tokenizer). Build the set from ``eos_token_id`` (int|list|None), and — when an
748
+ ``hf_model_id`` is supplied — also fold in the model's ``generation_config`` eos ids,
749
+ since chat models (e.g. Llama-3 Instruct) often carry extra eot ids there rather than
750
+ on ``eos_token_id``. Missing/empty -> empty set (truncation simply runs full length).
751
+ """
752
+ stop: set[int] = set()
753
+
754
+ def _add(value) -> None:
755
+ if value is None:
756
+ return
757
+ if isinstance(value, bool): # guard: bool is an int subclass
758
+ return
759
+ if isinstance(value, int):
760
+ stop.add(int(value))
761
+ elif isinstance(value, (list, tuple, set)):
762
+ for e in value:
763
+ _add(e)
764
+
765
+ _add(getattr(tokenizer, "eos_token_id", None))
766
+ # tt_transformers ModelArgs.tokenizer augments the HF tokenizer with ``stop_tokens``
767
+ # (eos + any extra eot ids); raw HF AutoTokenizers don't have it (getattr -> None).
768
+ _add(getattr(tokenizer, "stop_tokens", None))
769
+ if hf_model_id is not None:
770
+ try:
771
+ from transformers import GenerationConfig
772
+
773
+ gen_cfg = GenerationConfig.from_pretrained(hf_model_id)
774
+ _add(getattr(gen_cfg, "eos_token_id", None))
775
+ except Exception as e: # generation_config absent / unreadable — eos_token_id is enough
776
+ logger.debug(f"ci-eval-32: could not read generation_config eos ids for {hf_model_id}: {e}")
777
+ return stop
778
+
779
+
780
+ def decode_eval_output(tokenizer, token_ids: list[int], stop_ids: set[int]) -> str:
781
+ """Return the continuation text used by TTTv1's ci-eval-32 comparison.
782
+
783
+ Different valid BPE segmentations can decode to the same text. The TTTv1
784
+ reference stores ``tokenizer.decode(...)`` results and compares those
785
+ strings, so comparing token-id lists here would make the port stricter than
786
+ its source workload and report false slot-rotation failures.
787
+ """
788
+ decoded = tokenizer.decode(truncate_at_stop(token_ids, stop_ids))
789
+ if not isinstance(decoded, str):
790
+ raise TypeError(f"tokenizer.decode must return str, got {type(decoded).__name__}")
791
+ return decoded
792
+
793
+
794
+ def assert_cross_cardinality_consistency(
795
+ outputs_by_cardinality: dict[int, dict[str, str]],
796
+ *,
797
+ expected_cardinalities: tuple[int, ...] = (1, 2, 4, 32),
798
+ ) -> None:
799
+ """Require each fixed request's decoded output to be invariant as batch cardinality grows."""
800
+ if tuple(outputs_by_cardinality) != expected_cardinalities:
801
+ raise AssertionError(
802
+ f"cross-cardinality experiment expected {expected_cardinalities}, " f"got {tuple(outputs_by_cardinality)}"
803
+ )
804
+ reference: dict[str, tuple[int, str]] = {}
805
+ for cardinality, outputs in outputs_by_cardinality.items():
806
+ if len(outputs) != cardinality:
807
+ raise AssertionError(f"cardinality {cardinality} returned {len(outputs)} request outputs")
808
+ for request_id, output in outputs.items():
809
+ if request_id in reference:
810
+ reference_cardinality, reference_output = reference[request_id]
811
+ if output != reference_output:
812
+ raise AssertionError(
813
+ f"request {request_id!r} differs at cardinality {reference_cardinality}->{cardinality}: "
814
+ f"{reference_output[:120]!r} != {output[:120]!r}"
815
+ )
816
+ else:
817
+ reference[request_id] = (cardinality, output)
818
+
819
+
820
+ def assert_cross_batch_consistency(
821
+ per_repeat_outputs: list[list[str]],
822
+ *,
823
+ per_repeat_token_ids: list[list[list[int]]] | None = None,
824
+ per_repeat_prompt_lens: list[list[int]] | None = None,
825
+ per_repeat_argmax_margins: list[list[list[float]] | None] | None = None,
826
+ ) -> None:
827
+ """Assert decoded prompt-position invariance across repeats.
828
+
829
+ ``per_repeat_outputs[b][u]`` = decoded continuation for slot ``u`` of repeat ``b``.
830
+ With slot j of repeat b holding prompt (j+b)%N (see ``rotate_prompts``), the same prompt
831
+ sits at slot (offset+1)%N of repeat b and slot offset of repeat b+1 — so those two
832
+ outputs must be identical if no per-user state leaks.
833
+ """
834
+ num_batches = len(per_repeat_outputs)
835
+ assert num_batches >= 2, "cross-batch consistency needs >=2 repeats"
836
+ n = len(per_repeat_outputs[0])
837
+ failed, total = 0, 0
838
+ first_failure = None
839
+ for b in range(num_batches - 1):
840
+ cur, nxt = per_repeat_outputs[b], per_repeat_outputs[b + 1]
841
+ for offset in range(n):
842
+ total += 1
843
+ if cur[(offset + 1) % n] != nxt[offset]:
844
+ failed += 1
845
+ if first_failure is None:
846
+ current_slot = (offset + 1) % n
847
+ prompt_index = (offset + b + 1) % n
848
+ token_detail = ""
849
+ if per_repeat_token_ids is not None:
850
+ current_tokens = per_repeat_token_ids[b][current_slot]
851
+ next_tokens = per_repeat_token_ids[b + 1][offset]
852
+ common_tokens = 0
853
+ for current_token, next_token in zip(current_tokens, next_tokens):
854
+ if current_token != next_token:
855
+ break
856
+ common_tokens += 1
857
+ current_token = current_tokens[common_tokens] if common_tokens < len(current_tokens) else None
858
+ next_token = next_tokens[common_tokens] if common_tokens < len(next_tokens) else None
859
+ token_detail = (
860
+ f"; first token divergence at generation step {common_tokens} "
861
+ f"({current_token!r} != {next_token!r})"
862
+ )
863
+ if per_repeat_argmax_margins is not None:
864
+ current_repeat_margins = per_repeat_argmax_margins[b]
865
+ next_repeat_margins = per_repeat_argmax_margins[b + 1]
866
+ if current_repeat_margins is not None and next_repeat_margins is not None:
867
+ current_margin = current_repeat_margins[current_slot][common_tokens]
868
+ next_margin = next_repeat_margins[offset][common_tokens]
869
+ token_detail += f"; top2 margins {current_margin:.6g} and {next_margin:.6g}"
870
+ length_detail = ""
871
+ if per_repeat_prompt_lens is not None:
872
+ current_length = per_repeat_prompt_lens[b][current_slot]
873
+ next_length = per_repeat_prompt_lens[b + 1][offset]
874
+ length_detail = f"; prompt lengths {current_length} and {next_length}"
875
+ first_failure = (
876
+ b,
877
+ offset,
878
+ current_slot,
879
+ prompt_index,
880
+ cur[current_slot],
881
+ nxt[offset],
882
+ token_detail,
883
+ length_detail,
884
+ )
885
+ assert failed == 0, (
886
+ f"ci-eval-32: {failed}/{total} cross-batch consistency checks failed "
887
+ f"(first at repeat {first_failure[0]} slot {first_failure[2]} -> "
888
+ f"repeat {first_failure[0] + 1} slot {first_failure[1]}, prompt index {first_failure[3]}"
889
+ f"{first_failure[6]}{first_failure[7]}; "
890
+ f"decoded outputs {first_failure[4][:80]!r} != {first_failure[5][:80]!r})"
891
+ )
892
+
893
+
894
+ def assert_within_batch_slot_consistency(
895
+ decoded_outputs: list[str],
896
+ *,
897
+ token_ids: list[list[int]],
898
+ argmax_margins: list[list[float]] | None,
899
+ prompt_index: int,
900
+ ) -> None:
901
+ """Prove one fixed request is invariant across logical slots in one run."""
902
+
903
+ reference = decoded_outputs[0]
904
+ for slot, decoded in enumerate(decoded_outputs[1:], start=1):
905
+ if decoded == reference:
906
+ continue
907
+ reference_tokens = token_ids[0]
908
+ slot_tokens = token_ids[slot]
909
+ common_tokens = 0
910
+ for reference_token, slot_token in zip(reference_tokens, slot_tokens):
911
+ if reference_token != slot_token:
912
+ break
913
+ common_tokens += 1
914
+ reference_token = reference_tokens[common_tokens] if common_tokens < len(reference_tokens) else None
915
+ slot_token = slot_tokens[common_tokens] if common_tokens < len(slot_tokens) else None
916
+ margin_detail = ""
917
+ if argmax_margins is not None:
918
+ margin_detail = (
919
+ f"; top2 margins {argmax_margins[0][common_tokens]:.6g} "
920
+ f"and {argmax_margins[slot][common_tokens]:.6g}"
921
+ )
922
+ raise AssertionError(
923
+ f"ci-eval-32 identical-request diagnostic: prompt index {prompt_index} differs between "
924
+ f"logical slots 0 and {slot} at generation step {common_tokens} "
925
+ f"({reference_token!r} != {slot_token!r}){margin_detail}; "
926
+ f"decoded outputs {reference[:80]!r} != {decoded[:80]!r}"
927
+ )
928
+
929
+
930
+ def run_eval_repeat_batch32(
931
+ *,
932
+ make_executor,
933
+ allocate_kv_cache,
934
+ page_table: torch.Tensor,
935
+ prompts: list[str],
936
+ tokenizer,
937
+ tokenize_fn,
938
+ num_decode_tokens: int,
939
+ max_batch_size: int,
940
+ sampling_params=None,
941
+ repeat_batches: int = 3,
942
+ hf_model_id: str | None = None,
943
+ first_repeat_profiler=None,
944
+ page_table_mode: str = "slot-stable",
945
+ identical_prompt_index: int | None = None,
946
+ active_batch_size: int | None = None,
947
+ ) -> PerfBenchmarkResult:
948
+ """Drive the ci-eval-32 determinism case, building a fresh traced executor per repeat.
949
+
950
+ Each repeat builds its own traced executor (``make_executor()``) and its own zeroed KV
951
+ cache (``allocate_kv_cache(executor)``), so the rotated batches are fully independent —
952
+ no shared device or host state can leak across repeats. The executor is cleaned up after
953
+ each repeat. (The model is bit-deterministic across repeats either way; fresh-per-repeat
954
+ is simply the cleanest independence guarantee for a determinism test, and the trace
955
+ recapture cost is negligible at batch-32.)
956
+
957
+ Args:
958
+ make_executor: Zero-arg callable returning a fresh traced executor
959
+ (``run_perf_benchmark`` requires traced). Called once per repeat.
960
+ allocate_kv_cache: Callable(executor) -> fresh zeroed kv_cache bound on that executor.
961
+ page_table: Fixed contiguous page table (shared across repeats).
962
+ prompts: The N (=max_batch_size) prompts to rotate (TTTv1 ci-eval-32 numeric prompts;
963
+ see the module note above re: degenerate-output sensitivity on small models).
964
+ tokenizer: HF tokenizer (for stop / special ids).
965
+ tokenize_fn: Callable(list[str]) -> (tokens, prompt_lens).
966
+ num_decode_tokens: Decode steps per repeat.
967
+ max_batch_size: Padded batch (== len(prompts) for this fixed-32 case).
968
+ sampling_params: None -> host argmax (deterministic, mesh-agnostic default).
969
+ repeat_batches: Number of rotated repeats (TTTv1 uses 3).
970
+ hf_model_id: Optional, to enrich stop ids from generation_config.
971
+ first_repeat_profiler: Optional profiler passed only to the first repeat, allowing a caller
972
+ to emit perf telemetry without changing the three-repeat determinism geometry.
973
+ page_table_mode: ``slot-stable`` preserves TTTv1 acceptance geometry.
974
+ ``prompt-stable`` keeps each prompt on the same physical KV blocks
975
+ as a diagnostic A/B while it rotates through decode slots.
976
+ identical_prompt_index: Diagnostic-only fixed request replicated into
977
+ every logical slot. Slot consistency is checked after the first
978
+ batch, before any repeat-lifecycle comparison.
979
+ active_batch_size: Diagnostic-only number of leading logical slots to
980
+ activate while retaining the model's full lane capacity. This is
981
+ restricted to the identical-request probe so inactive lanes cannot
982
+ complicate prompt-rotation semantics.
983
+ """
984
+ assert (
985
+ len(prompts) == max_batch_size
986
+ ), f"ci-eval-32 expects len(prompts)==max_batch_size; got {len(prompts)} vs {max_batch_size}"
987
+ if active_batch_size is not None:
988
+ if identical_prompt_index is None:
989
+ raise ValueError("active_batch_size requires identical_prompt_index")
990
+ if not 1 <= active_batch_size <= max_batch_size:
991
+ raise ValueError(f"active_batch_size must be in [1, {max_batch_size}]")
992
+ if identical_prompt_index is not None:
993
+ if not 0 <= identical_prompt_index < len(prompts):
994
+ raise ValueError(f"identical_prompt_index must be in [0, {len(prompts) - 1}]")
995
+ prompts = [prompts[identical_prompt_index]] * (active_batch_size or len(prompts))
996
+ stop_ids = hf_stop_ids(tokenizer, hf_model_id)
997
+ special_ids = set(getattr(tokenizer, "all_special_ids", []) or [])
998
+ # Garbage guard targets only special tokens that are NOT recognized stops: a legitimate
999
+ # stop is removed by truncation, so anything special left in the body is degenerate output.
1000
+ garbage_ids = special_ids - stop_ids
1001
+ logger.info(
1002
+ f"ci-eval-32: repeat_batches={repeat_batches}, N={len(prompts)}, "
1003
+ f"stop_ids={sorted(stop_ids)}, |special_ids|={len(special_ids)}, sampling_params={sampling_params}, "
1004
+ f"page_table_mode={page_table_mode}, identical_prompt_index={identical_prompt_index}, "
1005
+ f"active_batch_size={len(prompts)}"
1006
+ )
1007
+
1008
+ per_repeat: list[list[str]] = []
1009
+ per_repeat_token_ids: list[list[list[int]]] = []
1010
+ per_repeat_prompt_lens: list[list[int]] = []
1011
+ per_repeat_argmax_margins: list[list[list[float]] | None] = []
1012
+ first_result = None
1013
+ for i in range(repeat_batches):
1014
+ traced_executor = make_executor()
1015
+ try:
1016
+ kv_cache = allocate_kv_cache(traced_executor)
1017
+ rotated = rotate_prompts(prompts, i)
1018
+ tokens, prompt_lens = tokenize_fn(rotated)
1019
+ repeat_page_table = eval_page_table_for_repeat(page_table, i, mode=page_table_mode)
1020
+ result = run_perf_benchmark(
1021
+ traced_executor,
1022
+ tokens=tokens,
1023
+ kv_cache=kv_cache,
1024
+ page_table=repeat_page_table,
1025
+ num_decode_tokens=num_decode_tokens,
1026
+ max_batch_size=max_batch_size,
1027
+ prompt_lens=prompt_lens,
1028
+ sampling_params=sampling_params,
1029
+ profiler=first_repeat_profiler if i == 0 else None,
1030
+ collect_argmax_diagnostics=sampling_params is None,
1031
+ )
1032
+ finally:
1033
+ traced_executor.cleanup()
1034
+ truncated = [truncate_at_stop(ids, stop_ids) for ids in result.generated_token_ids]
1035
+ for u, ids in enumerate(truncated):
1036
+ bad = set(ids) & garbage_ids
1037
+ assert not bad, f"ci-eval-32: user {u} produced special token(s) {sorted(bad)} mid-stream"
1038
+ decoded = [decode_eval_output(tokenizer, ids, stop_ids) for ids in result.generated_token_ids]
1039
+ per_repeat.append(decoded)
1040
+ per_repeat_token_ids.append(result.generated_token_ids)
1041
+ per_repeat_prompt_lens.append([int(length) for length in prompt_lens])
1042
+ per_repeat_argmax_margins.append(getattr(result, "argmax_top2_margins", None))
1043
+ if identical_prompt_index is not None:
1044
+ assert_within_batch_slot_consistency(
1045
+ decoded,
1046
+ token_ids=result.generated_token_ids,
1047
+ argmax_margins=getattr(result, "argmax_top2_margins", None),
1048
+ prompt_index=identical_prompt_index,
1049
+ )
1050
+ if i == 0:
1051
+ first_result = result
1052
+ logger.info(
1053
+ f"ci-eval-32 repeat {i}: truncated token lengths = {[len(t) for t in truncated]}, "
1054
+ f"decoded character lengths = {[len(text) for text in decoded]}"
1055
+ )
1056
+
1057
+ if identical_prompt_index is None:
1058
+ assert_cross_batch_consistency(
1059
+ per_repeat,
1060
+ per_repeat_token_ids=per_repeat_token_ids,
1061
+ per_repeat_prompt_lens=per_repeat_prompt_lens,
1062
+ per_repeat_argmax_margins=per_repeat_argmax_margins,
1063
+ )
1064
+ logger.info(f"ci-eval-32: all {(repeat_batches - 1) * len(prompts)} cross-batch consistency checks passed")
1065
+ else:
1066
+ logger.info(
1067
+ f"ci-eval-32 identical-request diagnostic: prompt index {identical_prompt_index} "
1068
+ f"is invariant across all {len(prompts)} logical slots"
1069
+ )
1070
+ assert first_result is not None
1071
+ return first_result
code/models/common/tests/demos/test_cleanup_utils.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import ast
5
+ from pathlib import Path
6
+ from types import SimpleNamespace
7
+
8
+ import pytest
9
+
10
+ from models.common.tests.demos import cleanup_utils
11
+
12
+
13
+ def test_cleanup_dp_model_case_orders_owners_before_parent_and_children(monkeypatch):
14
+ calls = []
15
+ parent = SimpleNamespace(quiesce_devices=lambda: calls.append("parent-quiesce"))
16
+ submeshes = [object(), object()]
17
+ group = SimpleNamespace(cleanup=lambda: calls.append("group-cleanup"))
18
+ models = [("model-0", submeshes[0]), ("model-1", submeshes[1])]
19
+
20
+ monkeypatch.setattr(
21
+ cleanup_utils,
22
+ "cleanup_model_case",
23
+ lambda model, submesh: calls.append(("model-cleanup", model, submesh)),
24
+ )
25
+ monkeypatch.setattr(
26
+ cleanup_utils.ttnn,
27
+ "close_mesh_device",
28
+ lambda submesh: calls.append(("child-close", submesh)),
29
+ )
30
+
31
+ cleanup_utils.cleanup_dp_model_case(group, [], models, parent, submeshes)
32
+
33
+ assert calls == [
34
+ "group-cleanup",
35
+ ("model-cleanup", "model-0", submeshes[0]),
36
+ ("model-cleanup", "model-1", submeshes[1]),
37
+ "parent-quiesce",
38
+ ("child-close", submeshes[0]),
39
+ ("child-close", submeshes[1]),
40
+ ]
41
+
42
+
43
+ def test_cleanup_dp_model_case_closes_every_carved_child_after_partial_build(monkeypatch):
44
+ calls = []
45
+ submeshes = [object(), object(), object(), object()]
46
+ lanes = [SimpleNamespace(cleanup=lambda: calls.append("lane-cleanup"))]
47
+ parent = SimpleNamespace(quiesce_devices=lambda: calls.append("parent-quiesce"))
48
+
49
+ monkeypatch.setattr(
50
+ cleanup_utils,
51
+ "cleanup_model_case",
52
+ lambda model, submesh: calls.append(("model-cleanup", model, submesh)),
53
+ )
54
+ monkeypatch.setattr(
55
+ cleanup_utils.ttnn,
56
+ "close_mesh_device",
57
+ lambda submesh: calls.append(("child-close", submesh)),
58
+ )
59
+
60
+ cleanup_utils.cleanup_dp_model_case(None, lanes, [("model-0", submeshes[0])], parent, submeshes)
61
+
62
+ assert calls[:3] == ["lane-cleanup", ("model-cleanup", "model-0", submeshes[0]), "parent-quiesce"]
63
+ assert calls[3:] == [("child-close", submesh) for submesh in submeshes]
64
+
65
+
66
+ def test_cleanup_dp_model_case_attempts_all_children_after_cleanup_failure(monkeypatch, expect_error):
67
+ closed = []
68
+ submeshes = [object(), object(), object()]
69
+
70
+ def fail_group_cleanup():
71
+ raise RuntimeError("group cleanup failed")
72
+
73
+ monkeypatch.setattr(cleanup_utils, "cleanup_model_case", lambda *_: None)
74
+ monkeypatch.setattr(cleanup_utils.ttnn, "close_mesh_device", closed.append)
75
+
76
+ with expect_error(RuntimeError, "group cleanup failed"):
77
+ cleanup_utils.cleanup_dp_model_case(
78
+ SimpleNamespace(cleanup=fail_group_cleanup),
79
+ [],
80
+ [],
81
+ SimpleNamespace(quiesce_devices=lambda: None),
82
+ submeshes,
83
+ )
84
+
85
+ assert closed == submeshes
86
+
87
+
88
+ def test_two_sequential_dp_profiles_release_children_between_carves(monkeypatch):
89
+ class Child:
90
+ def __init__(self, generation):
91
+ self.generation = generation
92
+ self.closed = False
93
+
94
+ class Parent:
95
+ def __init__(self):
96
+ self.children = []
97
+ self.quiesce_count = 0
98
+
99
+ def quiesce_devices(self):
100
+ self.quiesce_count += 1
101
+
102
+ def carve(self, generation):
103
+ assert all(child.closed for child in self.children)
104
+ children = [Child(generation), Child(generation)]
105
+ self.children.extend(children)
106
+ return children
107
+
108
+ parent = Parent()
109
+ monkeypatch.setattr(cleanup_utils.ttnn, "close_mesh_device", lambda child: setattr(child, "closed", True))
110
+
111
+ for generation in ("performance", "accuracy"):
112
+ parent.quiesce_devices()
113
+ children = parent.carve(generation)
114
+ cleanup_utils.cleanup_dp_model_case(None, [], [], parent, children)
115
+
116
+ assert parent.quiesce_count == 4
117
+ assert all(child.closed for child in parent.children)
118
+
119
+
120
+ @pytest.mark.parametrize(
121
+ ("demo_path", "create_name"),
122
+ [
123
+ ("models/common/tests/demos/qwen2_7b/demo.py", "_create_dp_submeshes"),
124
+ ("models/common/tests/demos/qwen25_7b/demo.py", "_create_dp_submeshes"),
125
+ ("models/common/tests/demos/deepseek_r1_distill_qwen_14b/demo.py", "_create_dp_submeshes"),
126
+ ("models/common/tests/demos/llama32_1b/demo.py", "create_dp_submeshes"),
127
+ ("models/common/tests/demos/llama32_3b/demo.py", "create_dp_submeshes"),
128
+ ],
129
+ )
130
+ def test_dp_demos_quiesce_before_carving_and_use_shared_teardown(demo_path, create_name):
131
+ tree = ast.parse(Path(demo_path).read_text(encoding="utf-8"), filename=demo_path)
132
+ function = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == "_run_dp_smoke")
133
+ calls = [node for node in ast.walk(function) if isinstance(node, ast.Call)]
134
+ create_call = next(node for node in calls if isinstance(node.func, ast.Name) and node.func.id == create_name)
135
+ parent_quiesce = next(
136
+ node
137
+ for node in calls
138
+ if isinstance(node.func, ast.Attribute)
139
+ and isinstance(node.func.value, ast.Name)
140
+ and node.func.value.id == "mesh_device"
141
+ and node.func.attr == "quiesce_devices"
142
+ )
143
+ teardown_call = next(
144
+ node for node in calls if isinstance(node.func, ast.Name) and node.func.id == "cleanup_dp_model_case"
145
+ )
146
+
147
+ assert parent_quiesce.lineno < create_call.lineno < teardown_call.lineno
148
+ assert [ast.unparse(argument) for argument in teardown_call.args] == [
149
+ "group",
150
+ "lanes",
151
+ "models",
152
+ "mesh_device",
153
+ "submeshes",
154
+ ]
code/models/common/tests/demos/test_run_helpers.py ADDED
@@ -0,0 +1,687 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+
6
+ import pytest
7
+ import torch
8
+
9
+ from models.common.tests.demos import run_helpers
10
+ from models.common.tests.demos.run_helpers import run_perf_benchmark, run_teacher_forcing
11
+
12
+
13
+ class FakeProfiler:
14
+ def __init__(self):
15
+ self.events = []
16
+
17
+ def start(self, name):
18
+ self.events.append(("start", name))
19
+
20
+ def end(self, name):
21
+ self.events.append(("end", name))
22
+
23
+
24
+ def _logits(token_ids, vocab_size=8):
25
+ output = torch.zeros(len(token_ids), 1, vocab_size)
26
+ for row, token_id in enumerate(token_ids):
27
+ output[row, 0, token_id] = 1
28
+ return output
29
+
30
+
31
+ class FakeExecutionTarget:
32
+ def __init__(self, *, compile_prefill_output, prefill_output, decode_outputs):
33
+ self.compile_prefill_output = compile_prefill_output
34
+ self.prefill_output = prefill_output
35
+ self.decode_outputs = list(decode_outputs)
36
+ self.calls = []
37
+
38
+ def _record(self, method_name, arguments):
39
+ self.calls.append(
40
+ (
41
+ method_name,
42
+ {name: value for name, value in arguments.items() if name != "self"},
43
+ )
44
+ )
45
+
46
+ @property
47
+ def _engine(self):
48
+ raise AssertionError("execution helpers must not inspect a private wrapped engine")
49
+
50
+ def compile_prefill(
51
+ self,
52
+ *,
53
+ tokens, # ↓ Core request
54
+ page_table,
55
+ prompt_lens=None, # ↓ Sequence metadata
56
+ start_pos=None,
57
+ empty_slots=None, # ↓ Lane routing
58
+ kv_cache=None, # ↓ Borrowed resources
59
+ sampling_params=None, # ↓ Sampling
60
+ execution=None, # ↓ Internal dispatch
61
+ ):
62
+ self._record("compile_prefill", locals())
63
+ return self.compile_prefill_output
64
+
65
+ def compile_decode(
66
+ self,
67
+ *,
68
+ tokens, # ↓ Core request
69
+ start_pos,
70
+ page_table,
71
+ kv_cache=None, # ↓ Borrowed resources
72
+ sampling_params=None, # ↓ Sampling
73
+ reset_batch=False, # ↓ State transition
74
+ execution=None, # ↓ Internal dispatch
75
+ ):
76
+ self._record("compile_decode", locals())
77
+
78
+ def prefill_forward(
79
+ self,
80
+ tokens,
81
+ page_table,
82
+ *,
83
+ prompt_lens=None, # ↓ Sequence metadata
84
+ start_pos=None,
85
+ empty_slots=None, # ↓ Lane routing
86
+ kv_cache=None, # ↓ Borrowed resources
87
+ sampling_params=None, # ↓ Sampling
88
+ execution=None, # ↓ Internal dispatch
89
+ ):
90
+ self._record("prefill_forward", locals())
91
+ return self.prefill_output
92
+
93
+ def decode_forward(
94
+ self,
95
+ tokens,
96
+ start_pos,
97
+ page_table,
98
+ *,
99
+ kv_cache=None, # ↓ Borrowed resources
100
+ sampling_params=None, # ↓ Sampling
101
+ reset_batch=False, # ↓ State transition
102
+ read_from_device=True, # ↓ Output policy
103
+ execution=None, # ↓ Internal dispatch
104
+ ):
105
+ self._record("decode_forward", locals())
106
+ return self.decode_outputs.pop(0)
107
+
108
+
109
+ def test_compile_only_helper_ignores_compiled_prefill_programs():
110
+ target = FakeExecutionTarget(
111
+ compile_prefill_output=(object(),),
112
+ prefill_output=None,
113
+ decode_outputs=[],
114
+ )
115
+
116
+ run_helpers._compile_prefill_and_decode(
117
+ target,
118
+ prefill_tokens=torch.tensor([[1, 2], [3, 4]]),
119
+ prefill_page_table=torch.zeros(2, 1, dtype=torch.int32),
120
+ )
121
+
122
+ assert [name for name, _ in target.calls] == ["compile_prefill", "compile_decode"]
123
+ assert torch.equal(target.calls[1][1]["tokens"], torch.zeros(2, dtype=torch.long))
124
+
125
+
126
+ def test_teacher_forcing_uses_public_target_surface_and_preserves_user_order():
127
+ target = FakeExecutionTarget(
128
+ compile_prefill_output=_logits([3, 4]),
129
+ prefill_output=_logits([3, 4]),
130
+ decode_outputs=[(_logits([5, 1]), None)],
131
+ )
132
+ top5_tokens = torch.tensor([[3, 0, 1, 2, 4], [5, 0, 1, 2, 3]])
133
+
134
+ result = run_teacher_forcing(
135
+ executor=target,
136
+ prompt_tokens=torch.tensor([[1, 2], [1, 2]]),
137
+ reference_tokens=torch.tensor([1, 2, 6, 7]),
138
+ top5_tokens=top5_tokens,
139
+ kv_cache=[],
140
+ page_table=torch.zeros(2, 1, dtype=torch.int32),
141
+ max_batch_size=2,
142
+ )
143
+
144
+ assert result.predicted_tokens_per_user == [[3, 5], [4, 1]]
145
+ assert result.top1_accuracy() == 1.0
146
+ assert result.top5_accuracy() == 1.0
147
+ assert [name for name, _ in target.calls] == [
148
+ "compile_prefill",
149
+ "compile_decode",
150
+ "prefill_forward",
151
+ "decode_forward",
152
+ ]
153
+
154
+
155
+ def test_teacher_forcing_times_prefill_excludes_first_decode_and_brackets_profiler(monkeypatch):
156
+ target = FakeExecutionTarget(
157
+ compile_prefill_output=_logits([3, 4]),
158
+ prefill_output=_logits([3, 4]),
159
+ decode_outputs=[(_logits([5, 1]), None), (_logits([6, 2]), None)],
160
+ )
161
+ profiler = FakeProfiler()
162
+ times = iter([0.0, 0.2, 1.0, 1.1, 2.0, 2.25])
163
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
164
+
165
+ result = run_teacher_forcing(
166
+ executor=target,
167
+ prompt_tokens=torch.tensor([[1, 2], [1, 2]]),
168
+ reference_tokens=torch.tensor([1, 2, 8, 9, 10]),
169
+ top5_tokens=torch.tensor([[3, 0, 1, 2, 4], [5, 0, 1, 2, 3], [6, 0, 1, 2, 3]]),
170
+ kv_cache=[],
171
+ page_table=torch.zeros(2, 1, dtype=torch.int32),
172
+ max_batch_size=2,
173
+ profiler=profiler,
174
+ )
175
+
176
+ assert result.predicted_tokens_per_user == [[3, 5, 6], [4, 1, 2]]
177
+ assert result.prefill_time_s == pytest.approx(0.2)
178
+ assert result.compile_decode_time_s == pytest.approx(0.1)
179
+ assert result.decode_times_s == pytest.approx([0.25])
180
+ assert result.ttft_ms == pytest.approx(100.0)
181
+ assert result.prefill_tok_s == pytest.approx(20.0)
182
+ assert result.decode_tok_s_u == pytest.approx(4.0)
183
+ assert result.decode_tok_s == pytest.approx(8.0)
184
+ assert profiler.events == [
185
+ ("start", "inference_prefill"),
186
+ ("end", "inference_prefill"),
187
+ ("start", "inference_decode"),
188
+ ("end", "inference_decode"),
189
+ ]
190
+
191
+
192
+ def test_perf_benchmark_host_argmax_path_preserves_timing_and_tokens(monkeypatch):
193
+ target = FakeExecutionTarget(
194
+ compile_prefill_output=_logits([2]),
195
+ prefill_output=_logits([2]),
196
+ decode_outputs=[(_logits([3]), None), (_logits([4]), None), (_logits([5]), None)],
197
+ )
198
+ times = iter([0.0, 0.1, 1.0, 1.2, 2.0, 2.25, 3.0, 3.3])
199
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
200
+
201
+ result = run_perf_benchmark(
202
+ target,
203
+ tokens=torch.tensor([[1, 2]]),
204
+ kv_cache=[],
205
+ page_table=torch.zeros(1, 1, dtype=torch.int32),
206
+ num_decode_tokens=3,
207
+ )
208
+
209
+ assert result.prefill_time_s == pytest.approx(0.1)
210
+ assert result.compile_decode_time_s == pytest.approx(0.2)
211
+ assert result.decode_times_s == pytest.approx([0.25, 0.3])
212
+ assert result.decode_iteration_times_s == pytest.approx([0.25, 0.3])
213
+ assert result.generated_token_ids == [[2, 3, 4, 5]]
214
+
215
+
216
+ def test_perf_benchmark_slices_prefill_page_rows_but_keeps_full_decode_capacity(monkeypatch):
217
+ target = FakeExecutionTarget(
218
+ compile_prefill_output=_logits([2, 2, 2, 2]),
219
+ prefill_output=_logits([2, 2]),
220
+ decode_outputs=[(_logits([3, 3, 0, 0]), None)],
221
+ )
222
+ times = iter([0.0, 0.1, 1.0, 1.1])
223
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
224
+ page_table = torch.arange(8, dtype=torch.int32).reshape(4, 2)
225
+
226
+ run_perf_benchmark(
227
+ target,
228
+ tokens=torch.tensor([[1, 2], [1, 2]]),
229
+ kv_cache=[],
230
+ page_table=page_table,
231
+ num_decode_tokens=1,
232
+ max_batch_size=4,
233
+ )
234
+
235
+ calls = {name: arguments for name, arguments in target.calls}
236
+ assert calls["compile_prefill"]["page_table"].shape[0] == 2
237
+ torch.testing.assert_close(calls["prefill_forward"]["page_table"], page_table[:2])
238
+ torch.testing.assert_close(calls["decode_forward"]["page_table"], page_table)
239
+
240
+
241
+ def test_perf_benchmark_brackets_profiler_without_changing_host_argmax(monkeypatch):
242
+ target = FakeExecutionTarget(
243
+ compile_prefill_output=_logits([2]),
244
+ prefill_output=_logits([2]),
245
+ decode_outputs=[(_logits([3]), None), (_logits([4]), None)],
246
+ )
247
+ profiler = FakeProfiler()
248
+ times = iter([0.0, 0.1, 1.0, 1.2, 2.0, 2.25])
249
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
250
+
251
+ result = run_perf_benchmark(
252
+ target,
253
+ tokens=torch.tensor([[1, 2]]),
254
+ kv_cache=[],
255
+ page_table=torch.zeros(1, 1, dtype=torch.int32),
256
+ num_decode_tokens=2,
257
+ profiler=profiler,
258
+ )
259
+
260
+ assert result.generated_token_ids == [[2, 3, 4]]
261
+ assert result.prefill_time_s == pytest.approx(0.1)
262
+ assert result.compile_decode_time_s == pytest.approx(0.2)
263
+ assert result.decode_times_s == pytest.approx([0.25])
264
+ assert profiler.events == [
265
+ ("start", "inference_prefill"),
266
+ ("end", "inference_prefill"),
267
+ ("start", "inference_decode"),
268
+ ("end", "inference_decode"),
269
+ ]
270
+
271
+
272
+ class PublicReadbackTarget(FakeExecutionTarget):
273
+ def __init__(self):
274
+ super().__init__(
275
+ compile_prefill_output=(torch.tensor([2]), None),
276
+ prefill_output=(torch.tensor([2]), None),
277
+ decode_outputs=[
278
+ (torch.tensor([3]), None),
279
+ (torch.tensor([4]), None),
280
+ (torch.tensor([5]), None),
281
+ ],
282
+ )
283
+ self.mesh_device = SimpleNamespace(shape=(1, 1))
284
+ self.next_event = 0
285
+ self.pending_events = []
286
+
287
+ def read_decode_output(self, tt_out, *, async_read=False):
288
+ assert async_read
289
+ event = self.next_event
290
+ self.next_event += 1
291
+ self.pending_events.append(event)
292
+ self.calls.append(("read_decode_output", {"event": event}))
293
+ return tt_out, [event]
294
+
295
+ def process_decode_output_host(self, tt_out, *, is_tokens=False):
296
+ assert is_tokens
297
+ run_helpers.ttnn.event_synchronize(self.pending_events.pop(0))
298
+ self.calls.append(("process_decode_output_host", {}))
299
+ return tt_out
300
+
301
+
302
+ def test_perf_benchmark_uses_public_async_readback_without_trace_introspection(monkeypatch):
303
+ target = PublicReadbackTarget()
304
+ profiler = FakeProfiler()
305
+ synchronized_events = []
306
+ times = iter([0.0, 0.1, 1.0, 1.1, 2.0, 2.1, 3.0, 3.1, 3.4])
307
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
308
+ monkeypatch.setattr(run_helpers.ttnn, "synchronize_device", lambda mesh: None, raising=False)
309
+ monkeypatch.setattr(run_helpers.ttnn, "event_synchronize", synchronized_events.append, raising=False)
310
+
311
+ result = run_perf_benchmark(
312
+ target,
313
+ tokens=torch.tensor([[1, 2]]),
314
+ kv_cache=[],
315
+ page_table=torch.zeros(1, 1, dtype=torch.int32),
316
+ num_decode_tokens=3,
317
+ sampling_params=object(),
318
+ pipeline_readback=True,
319
+ profiler=profiler,
320
+ )
321
+
322
+ assert [name for name, _ in target.calls[:2]] == ["compile_decode", "compile_prefill"]
323
+ # The benchmark owns the device-paced wait; the public normalizer may
324
+ # defensively observe the already-completed event again while retiring it.
325
+ assert synchronized_events == [0, 0, 1, 1, 2, 2]
326
+ assert result.generated_token_ids == [[2, 3, 4, 5]]
327
+ assert len(result.decode_times_s) == 2
328
+ assert result.decode_iteration_times_s == pytest.approx([0.1, 0.1])
329
+ assert profiler.events == [
330
+ ("start", "inference_prefill"),
331
+ ("end", "inference_prefill"),
332
+ ("start", "inference_decode"),
333
+ ("end", "inference_decode"),
334
+ ]
335
+
336
+
337
+ def test_perf_benchmark_does_not_reprocess_blocking_sampled_output(monkeypatch):
338
+ target = PublicReadbackTarget()
339
+ times = iter([0.0, 0.1, 1.0, 1.1])
340
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
341
+ monkeypatch.setattr(run_helpers.ttnn, "synchronize_device", lambda mesh: None, raising=False)
342
+
343
+ result = run_perf_benchmark(
344
+ target,
345
+ tokens=torch.tensor([[1, 2]]),
346
+ kv_cache=[],
347
+ page_table=torch.zeros(1, 1, dtype=torch.int32),
348
+ num_decode_tokens=1,
349
+ sampling_params=object(),
350
+ pipeline_readback=False,
351
+ )
352
+
353
+ assert result.generated_token_ids == [[2, 3]]
354
+ assert "process_decode_output_host" not in [name for name, _ in target.calls]
355
+
356
+
357
+ def test_perf_benchmark_can_use_host_prefill_with_sampled_decode(monkeypatch):
358
+ target = FakeExecutionTarget(
359
+ compile_prefill_output=_logits([2]),
360
+ prefill_output=_logits([2]),
361
+ decode_outputs=[(torch.tensor([3]), None)],
362
+ )
363
+ times = iter([0.0, 0.1, 1.0, 1.1])
364
+ monkeypatch.setattr(run_helpers.time, "perf_counter", lambda: next(times))
365
+
366
+ result = run_perf_benchmark(
367
+ target,
368
+ tokens=torch.tensor([[1, 2]]),
369
+ kv_cache=[],
370
+ page_table=torch.zeros(1, 1, dtype=torch.int32),
371
+ num_decode_tokens=1,
372
+ sampling_params=object(),
373
+ prefill_sampling_params=None,
374
+ )
375
+
376
+ calls = {name: arguments for name, arguments in target.calls}
377
+ assert calls["compile_prefill"]["sampling_params"] is None
378
+ assert calls["prefill_forward"]["sampling_params"] is None
379
+ assert calls["compile_decode"]["sampling_params"] is not None
380
+ assert calls["decode_forward"]["sampling_params"] is not None
381
+ assert result.generated_token_ids == [[2, 3]]
382
+
383
+
384
+ def test_composite_target_synchronizes_each_lane_mesh(monkeypatch):
385
+ target = SimpleNamespace(mesh_device="parent", mesh_devices=("lane-0", "lane-1"))
386
+ synchronized = []
387
+ monkeypatch.setattr(
388
+ run_helpers.ttnn,
389
+ "synchronize_device",
390
+ lambda mesh_device: synchronized.append(mesh_device),
391
+ raising=False,
392
+ )
393
+
394
+ run_helpers._synchronize_target(target)
395
+
396
+ assert synchronized == ["lane-0", "lane-1"]
397
+
398
+
399
+ def test_special_token_guard_truncates_at_stop_tokens_warns_locally_and_fails_in_ci(monkeypatch, expect_error):
400
+ tokenizer = SimpleNamespace(
401
+ all_special_ids=[0, 1, 2, 99],
402
+ eos_token_id=2,
403
+ stop_tokens=[3],
404
+ convert_tokens_to_ids=lambda token: 99 if token == "<|eot_id|>" else -1,
405
+ )
406
+ generated_token_ids = [
407
+ [5, 2, 0],
408
+ [5, 3, 0],
409
+ [6, 99, 1],
410
+ [7, 0, 8],
411
+ [9, 1, 2],
412
+ ]
413
+ warnings = []
414
+ monkeypatch.setattr(run_helpers.logger, "warning", warnings.append)
415
+
416
+ run_helpers.assert_no_special_tokens(generated_token_ids, tokenizer, case_name="batch-1", is_ci_env=False)
417
+
418
+ assert warnings == ["[batch-1] model produced special tokens (2/5 users)"]
419
+ with expect_error(AssertionError, "2/5 users"):
420
+ run_helpers.assert_no_special_tokens(generated_token_ids, tokenizer, case_name="batch-1", is_ci_env=True)
421
+
422
+
423
+ def test_special_token_guard_without_stop_tokens_keeps_generated_tail_visible(expect_error):
424
+ tokenizer = SimpleNamespace(
425
+ all_special_ids=[0],
426
+ eos_token_id=2,
427
+ convert_tokens_to_ids=lambda _token: -1,
428
+ )
429
+
430
+ with expect_error(AssertionError, "1/1 users"):
431
+ run_helpers.assert_no_special_tokens([[5, 3, 0]], tokenizer, is_ci_env=True)
432
+
433
+
434
+ def test_eval_repeat_compares_decoded_text_like_tttv1_despite_different_bpe_segmentations():
435
+ pieces = {10: "12", 11: "3", 12: "1", 13: "23", 99: "<eos>", 77: "ignored"}
436
+ tokenizer = SimpleNamespace(decode=lambda token_ids: "".join(pieces[token] for token in token_ids))
437
+
438
+ first_segmentation = run_helpers.decode_eval_output(tokenizer, [10, 11, 99, 77], {99})
439
+ second_segmentation = run_helpers.decode_eval_output(tokenizer, [12, 13], {99})
440
+
441
+ assert first_segmentation == second_segmentation == "123"
442
+ # Repeat 1 rotates prompt 1 into slot 0 and prompt 0 into slot 1.
443
+ run_helpers.assert_cross_batch_consistency(
444
+ [
445
+ [first_segmentation, "456"],
446
+ ["456", second_segmentation],
447
+ ]
448
+ )
449
+
450
+
451
+ def test_eval_page_table_ab_preserves_slots_or_prompt_physical_blocks(expect_error):
452
+ page_table = torch.tensor([[0, 1], [10, 11], [20, 21]], dtype=torch.int32)
453
+
454
+ assert run_helpers.eval_page_table_for_repeat(page_table, 2, mode="slot-stable") is page_table
455
+ torch.testing.assert_close(
456
+ run_helpers.eval_page_table_for_repeat(page_table, 1, mode="prompt-stable"),
457
+ torch.tensor([[10, 11], [20, 21], [0, 1]], dtype=torch.int32),
458
+ )
459
+ with expect_error(ValueError, "slot-stable.*prompt-stable"):
460
+ run_helpers.eval_page_table_for_repeat(page_table, 0, mode="unsupported")
461
+
462
+
463
+ def test_eval_decode_ab_defaults_to_trace_and_can_isolate_eager_execution(expect_error):
464
+ assert run_helpers.eval_decode_trace_mode("traced") == "decode_only"
465
+ assert run_helpers.eval_decode_trace_mode("eager") == "none"
466
+ with expect_error(ValueError, "traced.*eager"):
467
+ run_helpers.eval_decode_trace_mode("unsupported")
468
+
469
+
470
+ @pytest.mark.parametrize(
471
+ "override",
472
+ [
473
+ {"EVAL_DECODE_MODE": "eager"},
474
+ {"EVAL_PAGE_TABLE_MODE": "prompt-stable"},
475
+ {"EVAL_IDENTICAL_PROMPT_INDEX": "8"},
476
+ {"EVAL_ACTIVE_BATCH_SIZE": "24"},
477
+ ],
478
+ )
479
+ def test_ci_rejects_diagnostic_eval_modes(override, expect_error):
480
+ with expect_error(RuntimeError, "diagnostic eval modes cannot replace the canonical CI gate"):
481
+ run_helpers.require_canonical_eval_modes_in_ci({"CI": "true", **override})
482
+
483
+
484
+ def test_non_ci_diagnostics_and_canonical_ci_are_allowed():
485
+ run_helpers.require_canonical_eval_modes_in_ci({"EVAL_DECODE_MODE": "eager", "EVAL_IDENTICAL_PROMPT_INDEX": "8"})
486
+ run_helpers.require_canonical_eval_modes_in_ci(
487
+ {"CI": "true", "EVAL_DECODE_MODE": "traced", "EVAL_PAGE_TABLE_MODE": "slot-stable"}
488
+ )
489
+
490
+
491
+ def test_host_argmax_diagnostic_reports_top1_minus_top2_margin():
492
+ logits = torch.tensor(
493
+ [
494
+ [[0.0, 4.0, 3.5]],
495
+ [[9.0, 1.0, 8.875]],
496
+ ]
497
+ )
498
+
499
+ tokens, margins = run_helpers._host_argmax_with_margins(logits, 2)
500
+
501
+ torch.testing.assert_close(tokens, torch.tensor([1, 0]))
502
+ torch.testing.assert_close(margins, torch.tensor([0.5, 0.125]))
503
+
504
+
505
+ def test_eval_repeat_driver_canonicalizes_each_rotated_output_through_tokenizer(monkeypatch):
506
+ pieces = {10: "12", 11: "3", 12: "1", 13: "23", 20: "456"}
507
+ tokenizer = SimpleNamespace(
508
+ decode=lambda token_ids: "".join(pieces[token] for token in token_ids),
509
+ eos_token_id=None,
510
+ stop_tokens=[],
511
+ all_special_ids=[],
512
+ )
513
+ generated = iter(
514
+ [
515
+ [[10, 11], [20]],
516
+ [[20], [12, 13]],
517
+ ]
518
+ )
519
+ monkeypatch.setattr(
520
+ run_helpers,
521
+ "run_perf_benchmark",
522
+ lambda *_args, **_kwargs: SimpleNamespace(generated_token_ids=next(generated)),
523
+ )
524
+
525
+ class FakeEvalExecutor:
526
+ def cleanup(self):
527
+ pass
528
+
529
+ result = run_helpers.run_eval_repeat_batch32(
530
+ make_executor=FakeEvalExecutor,
531
+ allocate_kv_cache=lambda _executor: [],
532
+ page_table=torch.zeros(2, 1, dtype=torch.int32),
533
+ prompts=["prompt-0", "prompt-1"],
534
+ tokenizer=tokenizer,
535
+ tokenize_fn=lambda prompts: (torch.zeros(len(prompts), 1, dtype=torch.long), torch.ones(len(prompts))),
536
+ num_decode_tokens=2,
537
+ max_batch_size=2,
538
+ repeat_batches=2,
539
+ )
540
+
541
+ assert result.generated_token_ids == [[10, 11], [20]]
542
+
543
+
544
+ def test_eval_repeat_still_rejects_genuine_decoded_text_divergence(expect_error):
545
+ with expect_error(AssertionError, "1/2 cross-batch consistency checks failed"):
546
+ run_helpers.assert_cross_batch_consistency(
547
+ [
548
+ ["123", "456"],
549
+ ["different", "123"],
550
+ ]
551
+ )
552
+
553
+
554
+ def test_eval_repeat_failure_localizes_prompt_slots_and_first_token_divergence(expect_error):
555
+ with expect_error(
556
+ AssertionError,
557
+ r"repeat 0 slot 1 -> repeat 1 slot 0, prompt index 1; "
558
+ r"first token divergence at generation step 1 \(12 != 99\); "
559
+ r"top2 margins 0.125 and 0.25; prompt lengths 80 and 80",
560
+ ):
561
+ run_helpers.assert_cross_batch_consistency(
562
+ [
563
+ ["alpha", "beta-left"],
564
+ ["beta-right", "alpha"],
565
+ ],
566
+ per_repeat_token_ids=[
567
+ [[1], [11, 12, 13]],
568
+ [[11, 99, 13], [1]],
569
+ ],
570
+ per_repeat_prompt_lens=[
571
+ [64, 80],
572
+ [80, 64],
573
+ ],
574
+ per_repeat_argmax_margins=[
575
+ [[1.0], [0.5, 0.125, 0.75]],
576
+ [[0.5, 0.25, 0.75], [1.0]],
577
+ ],
578
+ )
579
+
580
+
581
+ def test_identical_request_diagnostic_localizes_logical_slot_divergence(expect_error):
582
+ with expect_error(
583
+ AssertionError,
584
+ r"prompt index 8 differs between logical slots 0 and 1 at generation step 2 "
585
+ r"\(13 != 99\); top2 margins 0.125 and 0.25",
586
+ ):
587
+ run_helpers.assert_within_batch_slot_consistency(
588
+ ["same prefix left", "same prefix right"],
589
+ token_ids=[[11, 12, 13], [11, 12, 99]],
590
+ argmax_margins=[[1.0, 0.5, 0.125], [1.0, 0.5, 0.25]],
591
+ prompt_index=8,
592
+ )
593
+
594
+
595
+ def test_identical_request_diagnostic_accepts_slot_invariant_decoded_text():
596
+ run_helpers.assert_within_batch_slot_consistency(
597
+ ["same", "same", "same"],
598
+ token_ids=[[1], [1], [1]],
599
+ argmax_margins=[[0.5], [0.5], [0.5]],
600
+ prompt_index=8,
601
+ )
602
+
603
+
604
+ def test_identical_request_driver_uses_one_fixed_prompt_and_needs_no_repeat(monkeypatch):
605
+ seen_prompts = []
606
+ tokenizer = SimpleNamespace(
607
+ decode=lambda token_ids: "same",
608
+ eos_token_id=None,
609
+ stop_tokens=[],
610
+ all_special_ids=[],
611
+ )
612
+ monkeypatch.setattr(
613
+ run_helpers,
614
+ "run_perf_benchmark",
615
+ lambda *_args, **_kwargs: SimpleNamespace(
616
+ generated_token_ids=[[1]],
617
+ argmax_top2_margins=[[0.5]],
618
+ ),
619
+ )
620
+
621
+ class FakeEvalExecutor:
622
+ def cleanup(self):
623
+ pass
624
+
625
+ run_helpers.run_eval_repeat_batch32(
626
+ make_executor=FakeEvalExecutor,
627
+ allocate_kv_cache=lambda _executor: [],
628
+ page_table=torch.zeros(2, 1, dtype=torch.int32),
629
+ prompts=["prompt-0", "prompt-1"],
630
+ tokenizer=tokenizer,
631
+ tokenize_fn=lambda prompts: (
632
+ seen_prompts.append(prompts) or torch.zeros(len(prompts), 1, dtype=torch.long),
633
+ torch.ones(len(prompts)),
634
+ ),
635
+ num_decode_tokens=1,
636
+ max_batch_size=2,
637
+ repeat_batches=1,
638
+ identical_prompt_index=1,
639
+ active_batch_size=1,
640
+ )
641
+
642
+ assert seen_prompts == [["prompt-1"]]
643
+
644
+
645
+ def test_active_batch_diagnostic_requires_identical_request(expect_error):
646
+ with expect_error(ValueError, "active_batch_size requires identical_prompt_index"):
647
+ run_helpers.run_eval_repeat_batch32(
648
+ make_executor=lambda: None,
649
+ allocate_kv_cache=lambda _executor: [],
650
+ page_table=torch.zeros(2, 1, dtype=torch.int32),
651
+ prompts=["prompt-0", "prompt-1"],
652
+ tokenizer=SimpleNamespace(eos_token_id=None, stop_tokens=[], all_special_ids=[]),
653
+ tokenize_fn=lambda prompts: (torch.zeros(len(prompts), 1), torch.ones(len(prompts))),
654
+ num_decode_tokens=1,
655
+ max_batch_size=2,
656
+ repeat_batches=1,
657
+ active_batch_size=1,
658
+ )
659
+
660
+
661
+ def test_cross_cardinality_consistency_accepts_fixed_request_prefixes():
662
+ run_helpers.assert_cross_cardinality_consistency(
663
+ {
664
+ 1: {"r0": "a"},
665
+ 2: {"r0": "a", "r1": "b"},
666
+ 4: {"r0": "a", "r1": "b", "r2": "c", "r3": "d"},
667
+ 32: {f"r{i}": chr(97 + i) for i in range(32)},
668
+ }
669
+ )
670
+
671
+
672
+ def test_cross_cardinality_consistency_reports_request_and_cardinalities(expect_error):
673
+ with expect_error(AssertionError, "request 'r0' differs at cardinality 1->2"):
674
+ run_helpers.assert_cross_cardinality_consistency(
675
+ {1: {"r0": "same"}, 2: {"r0": "different", "r1": "x"}},
676
+ expected_cardinalities=(1, 2),
677
+ )
678
+
679
+
680
+ def test_loop_policy_is_not_exported_from_production_executor():
681
+ try:
682
+ from models.common.llm_runtime import execution as production_executor
683
+ except AttributeError as exc:
684
+ pytest.skip(f"production executor import requires full ttnn runtime: {exc}")
685
+
686
+ for name in ("TeacherForceResult", "PerfBenchmarkResult", "run_teacher_forcing", "run_perf_benchmark"):
687
+ assert not hasattr(production_executor, name)
code/models/common/tests/host/test_metrics_pytorch_only.py ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """PyTorch-only metric tests that do not require TTNN hardware."""
5
+
6
+ import torch
7
+
8
+ from models.common.metrics import comp_allclose, compute_max_abs_error, compute_mean_abs_error, compute_pcc
9
+
10
+
11
+ def test_metrics_pytorch_only():
12
+ """Verify fallback metric implementations using pure PyTorch tensors."""
13
+
14
+ torch_a = torch.tensor([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]], dtype=torch.float32)
15
+ torch_b = torch.tensor([[1.0, 2.0, 3.5], [4.0, 5.0, 6.0]], dtype=torch.float32)
16
+
17
+ max_error = compute_max_abs_error(torch_a, torch_b)
18
+ mean_error = compute_mean_abs_error(torch_a, torch_b)
19
+ pcc = compute_pcc(torch_a, torch_b)
20
+
21
+ expected_max = 0.5
22
+ expected_mean = (torch_a - torch_b).abs().mean().item()
23
+
24
+ assert abs(max_error - expected_max) < 1e-6, f"Expected {expected_max}, got {max_error}"
25
+ assert abs(mean_error - expected_mean) < 1e-6, "Mean error mismatch"
26
+ assert 0.99 <= pcc <= 1.0, f"PCC should be high (~1.0) for similar tensors, got {pcc}"
27
+
28
+
29
+ def test_comp_allclose_pytorch_only():
30
+ """PyTorch-only tests for comp_allclose covering pass/fail and edge cases."""
31
+
32
+ a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32)
33
+ b = a.clone()
34
+ passed, msg = comp_allclose(a, b)
35
+ assert passed, f"Expected pass for exact equality. Got: {msg}"
36
+ assert "Max ATOL Delta" in msg and "Max RTOL Delta" in msg
37
+
38
+ a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32)
39
+ b = torch.tensor([1.0 + 1e-7, 2.0 - 1e-7, 3.0], dtype=torch.float32)
40
+ passed, msg = comp_allclose(a, b)
41
+ assert passed, f"Expected pass within default tolerance. Got: {msg}"
42
+
43
+ a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32)
44
+ b = torch.tensor([1.0, 2.1, 3.0], dtype=torch.float32)
45
+ passed, msg = comp_allclose(a, b)
46
+ assert not passed and "Allclose check failed" in msg
47
+
48
+ a = torch.tensor([float("nan"), 1.0, 2.0], dtype=torch.float32)
49
+ b = torch.tensor([float("nan"), 1.0, 2.0], dtype=torch.float32)
50
+ passed, _ = comp_allclose(a, b)
51
+ assert passed, "Both NaNs at same positions should pass"
52
+
53
+ a = torch.tensor([float("inf"), -float("inf"), 1.0])
54
+ b = torch.tensor([float("inf"), -float("inf"), 1.0])
55
+ passed, _ = comp_allclose(a, b)
56
+ assert passed, "Same sign infinities should pass"
57
+
58
+ a = torch.tensor([float("inf"), -float("inf"), 1.0])
59
+ b = torch.tensor([float("inf"), float("inf"), 1.0])
60
+ passed, msg = comp_allclose(a, b)
61
+ assert not passed and "Allclose check failed" in msg
code/models/common/tests/host/test_utility_functions_imports.py ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Import-dependency tests for models.common.utility_functions."""
5
+
6
+ import ast
7
+ from pathlib import Path
8
+
9
+
10
+ def test_pytest_is_not_imported_at_module_scope():
11
+ utility_functions_path = Path(__file__).parents[2] / "utility_functions.py"
12
+ syntax_tree = ast.parse(utility_functions_path.read_text())
13
+
14
+ top_level_pytest_imports = [
15
+ node
16
+ for node in syntax_tree.body
17
+ if (isinstance(node, ast.Import) and any(alias.name == "pytest" for alias in node.names))
18
+ or (isinstance(node, ast.ImportFrom) and node.module == "pytest")
19
+ ]
20
+
21
+ assert not top_level_pytest_imports, "pytest must remain an optional test-only dependency"
22
+
23
+
24
+ def test_ti_skip_imports_pytest_lazily():
25
+ utility_functions_path = Path(__file__).parents[2] / "utility_functions.py"
26
+ syntax_tree = ast.parse(utility_functions_path.read_text())
27
+ ti_skip = next(node for node in syntax_tree.body if isinstance(node, ast.FunctionDef) and node.name == "ti_skip")
28
+
29
+ assert any(
30
+ isinstance(node, ast.Import) and any(alias.name == "pytest" for alias in node.names) for node in ti_skip.body
31
+ ), "ti_skip must import pytest when the test helper is used"
code/models/common/tests/llm_runtime/test_config.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from dataclasses import FrozenInstanceError, fields
5
+
6
+ import pytest
7
+
8
+ import ttnn
9
+ from models.common.llm_runtime import config as runtime_config
10
+ from models.common.llm_runtime.config import PagedKVCacheConfig, PageTableLayout, TraceConfig, WarmupConfig
11
+ from models.common.models.llama3_8b.executor import Llama3ExecutorConfig
12
+
13
+
14
+ class _TraceConfigSubclass(TraceConfig):
15
+ pass
16
+
17
+
18
+ def _paged_config(**overrides):
19
+ kwargs = {
20
+ "block_size": 32,
21
+ "max_num_blocks": 1024,
22
+ "dtype": ttnn.bfloat8_b,
23
+ }
24
+ kwargs.update(overrides)
25
+ return PagedKVCacheConfig(**kwargs)
26
+
27
+
28
+ def test_executor_config_has_exact_static_policy_owners_and_is_frozen(expect_error):
29
+ config = Llama3ExecutorConfig(
30
+ trace=TraceConfig(mode="all"),
31
+ warmup=WarmupConfig(),
32
+ paged_kv_cache=_paged_config(),
33
+ device_sampling_enabled=True,
34
+ )
35
+
36
+ assert [field.name for field in fields(config)] == [
37
+ "trace",
38
+ "warmup",
39
+ "paged_kv_cache",
40
+ "device_sampling_enabled",
41
+ "allow_batched_prefill_with_device_sampling_for_diagnostics",
42
+ ]
43
+ assert not config.allow_batched_prefill_with_device_sampling_for_diagnostics
44
+ forbidden = {
45
+ "model",
46
+ "mesh_device",
47
+ "hf_model",
48
+ "tokenizer",
49
+ "dtype",
50
+ "n_layers",
51
+ "sampling_config",
52
+ "sampling_output_dtype",
53
+ }
54
+ assert forbidden.isdisjoint(field.name for field in fields(config))
55
+ assert not hasattr(runtime_config, "LLMGraphCompilerConfig")
56
+ assert not hasattr(runtime_config, "LLMExecutorConfig")
57
+ assert not hasattr(runtime_config, "Sampling1DConfig")
58
+ with expect_error(FrozenInstanceError, ""):
59
+ config.device_sampling_enabled = False
60
+
61
+
62
+ @pytest.mark.parametrize(
63
+ ("field_name", "invalid_value"),
64
+ [
65
+ ("trace", WarmupConfig()),
66
+ ("trace", _TraceConfigSubclass()),
67
+ ("warmup", TraceConfig()),
68
+ ("paged_kv_cache", WarmupConfig()),
69
+ ],
70
+ )
71
+ def test_executor_config_rejects_non_exact_nested_config_types(field_name, invalid_value, expect_error):
72
+ values = {
73
+ "trace": TraceConfig(),
74
+ "warmup": WarmupConfig(),
75
+ "paged_kv_cache": _paged_config(),
76
+ "device_sampling_enabled": False,
77
+ }
78
+ values[field_name] = invalid_value
79
+
80
+ with expect_error(TypeError, rf"{field_name} must be exactly"):
81
+ Llama3ExecutorConfig(**values)
82
+
83
+
84
+ @pytest.mark.parametrize(
85
+ ("mode", "prefill", "decode"),
86
+ [("none", False, False), ("decode_only", False, True), ("all", True, True)],
87
+ )
88
+ def test_trace_config_selects_static_coverage(mode, prefill, decode, expect_error):
89
+ config = TraceConfig(mode=mode)
90
+
91
+ assert config.prefill_enabled is prefill
92
+ assert config.decode_enabled is decode
93
+ with expect_error(FrozenInstanceError, ""):
94
+ config.mode = "none"
95
+
96
+
97
+ def test_trace_config_rejects_unknown_mode(expect_error):
98
+ with expect_error(ValueError, "Unsupported trace mode"):
99
+ TraceConfig(mode="prefill_only")
100
+
101
+
102
+ def test_warmup_config_keeps_model_derived_defaults_and_is_deeply_immutable(expect_error):
103
+ config = WarmupConfig()
104
+
105
+ assert config.prefill_seq_lens is None
106
+ assert config.prefill_batch_sizes == (1, 2, 4, 8, 16, 32)
107
+ assert config.include_decode_top_k is False
108
+ with expect_error(TypeError, "must be a tuple"):
109
+ WarmupConfig(prefill_batch_sizes=[1, 2])
110
+
111
+
112
+ def test_paged_kv_config_has_plan_fields_and_resolved_capacity(expect_error):
113
+ unresolved = _paged_config()
114
+ resolved = _paged_config(num_blocks=512)
115
+
116
+ assert [field.name for field in fields(unresolved)] == [
117
+ "block_size",
118
+ "max_num_blocks",
119
+ "dtype",
120
+ "memory_config",
121
+ "num_blocks",
122
+ ]
123
+ assert unresolved.memory_config == ttnn.DRAM_MEMORY_CONFIG
124
+ assert not unresolved.is_resolved()
125
+ assert resolved.is_resolved()
126
+ with expect_error(FrozenInstanceError, ""):
127
+ resolved.num_blocks = 256
128
+
129
+
130
+ def test_paged_kv_config_rejects_invalid_capacity(expect_error):
131
+ with expect_error(ValueError, "exceeds max_num_blocks"):
132
+ _paged_config(num_blocks=1025)
133
+ with expect_error(ValueError, "block_size"):
134
+ _paged_config(block_size=0)
135
+
136
+
137
+ def test_page_table_layout_is_resolved_without_warmup_policy():
138
+ layout = PageTableLayout.resolve(
139
+ block_size=32,
140
+ model_max_sequence_length=4096,
141
+ physical_num_blocks=100,
142
+ max_prefill_chunk_size=2048,
143
+ )
144
+
145
+ assert layout.raw_capacity_width == 100
146
+ assert layout.decode_width == 104
147
+ assert layout.prefill_width == 168
code/models/common/tests/llm_runtime/test_decode_runtime.py ADDED
@@ -0,0 +1,1337 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from __future__ import annotations
5
+
6
+ import dataclasses
7
+ import inspect
8
+ from types import SimpleNamespace
9
+ from typing import Any
10
+
11
+ import pytest
12
+ import torch
13
+
14
+ import models.common.llm_runtime.decode as decode_module
15
+ import ttnn
16
+ from models.common.llm_runtime.config import PageTableLayout
17
+ from models.common.llm_runtime.decode import (
18
+ DecodeDeviceInputs,
19
+ DecodePersistentInputs,
20
+ DecodeProgramSignature,
21
+ DecodeRuntime,
22
+ DecodeRuntimeConfig,
23
+ DecodeTraceSignature,
24
+ InvocationResult,
25
+ )
26
+ from models.common.llm_runtime.output_reader import OutputReader, PendingRead
27
+ from models.common.modules.sampling.params import PreparedSamplingParams
28
+ from models.common.sampling.sampling_params import SamplingParams
29
+
30
+
31
+ class FakeMesh:
32
+ shape = (1, 1)
33
+
34
+
35
+ class FakeSampling:
36
+ def __init__(self, seed_buffer=None):
37
+ self.config = SimpleNamespace(allow_force_argmax=True, max_batch_size=2, max_top_k=32, seeds=seed_buffer)
38
+
39
+ def decode_forward(
40
+ self,
41
+ logits,
42
+ *,
43
+ k=None,
44
+ p=None,
45
+ temp=None,
46
+ seeds=None,
47
+ tt_out_tok=None,
48
+ enable_log_probs=False,
49
+ ):
50
+ return logits, None
51
+
52
+
53
+ class FakeRope:
54
+ def get_rot_idxs(self, positions, *, on_host):
55
+ assert on_host
56
+ return ("rotary", positions.clone())
57
+
58
+ def get_rot_mats(self, rotary_indices):
59
+ return ("cos", "sin")
60
+
61
+
62
+ class FakeModel:
63
+ def __init__(self, seed_buffer=None):
64
+ self.config = SimpleNamespace(max_batch_size=2)
65
+ self.sampling = FakeSampling(seed_buffer)
66
+ self.rope_setup = FakeRope()
67
+ self.vocab_size = 8
68
+ self.num_devices = 1
69
+
70
+ def iter_executor_named_modules(self):
71
+ return iter(())
72
+
73
+ def increment_positions(self, positions, rotary_indices):
74
+ return None
75
+
76
+
77
+ class FakeLazySeedBuffer:
78
+ def __init__(self):
79
+ self.source = torch.arange(2, dtype=torch.int64)
80
+ self._value = object()
81
+ self.updates = []
82
+
83
+ def update(self, source):
84
+ self.source = source
85
+ self.updates.append(source.clone())
86
+
87
+ def get_device_buffer(self):
88
+ return self._value
89
+
90
+
91
+ def make_runtime(*, sampling=True, force_greedy_top_k=False, seed_buffer=None):
92
+ mesh = FakeMesh()
93
+ model = FakeModel(seed_buffer)
94
+ config = DecodeRuntimeConfig.resolve(
95
+ model=model,
96
+ output_reader=OutputReader(mesh),
97
+ lane_capacity=2,
98
+ page_table_layout=page_table_layout(),
99
+ device_sampling_enabled=sampling,
100
+ force_greedy_top_k=force_greedy_top_k,
101
+ )
102
+ return DecodeRuntime(config)
103
+
104
+
105
+ def page_table_layout(*, raw_width=8, block_size=32):
106
+ return PageTableLayout(
107
+ block_size=block_size,
108
+ raw_capacity_width=raw_width,
109
+ prefill_width=((raw_width + 7) // 8) * 8,
110
+ decode_width=((raw_width + 7) // 8) * 8,
111
+ )
112
+
113
+
114
+ def greedy_sampling():
115
+ return SamplingParams(temperature=[0.0, 0.0], top_k=[1, 1], top_p=[1.0, 1.0])
116
+
117
+
118
+ def test_prepare_uses_resolved_sampler_capacity_and_neutral_inactive_rows():
119
+ prepared = prepare(
120
+ make_runtime(),
121
+ positions=(0, -1),
122
+ sampling_params=SamplingParams(temperature=1.0, top_k=32, top_p=0.08),
123
+ ).prepared_sampling
124
+
125
+ assert isinstance(prepared, PreparedSamplingParams)
126
+ assert prepared.batch_size == 2
127
+ assert prepared.active_rows == 1
128
+ assert prepared.active_mask == (True, False)
129
+ assert prepared.top_k == (32, 1)
130
+ assert prepared.top_p == pytest.approx((0.08, 0.0))
131
+ assert prepared.temperature == (1.0, 1.0)
132
+ assert prepared.row_paths == ("topk", "inactive")
133
+
134
+
135
+ def test_prepare_accepts_vector_tensor_fields_for_full_lane():
136
+ prepared = prepare(
137
+ make_runtime(),
138
+ positions=(0, 0),
139
+ sampling_params=SamplingParams(
140
+ temperature=torch.ones(2),
141
+ top_k=torch.full((2,), 32, dtype=torch.int32),
142
+ top_p=torch.full((2,), 0.08),
143
+ ),
144
+ ).prepared_sampling
145
+
146
+ assert isinstance(prepared, PreparedSamplingParams)
147
+ assert prepared.active_mask == (True, True)
148
+ assert prepared.top_k == (32, 32)
149
+ assert prepared.top_p == pytest.approx((0.08, 0.08))
150
+ assert prepared.temperature == (1.0, 1.0)
151
+
152
+
153
+ def prepare(
154
+ runtime,
155
+ *,
156
+ positions=(0, -1),
157
+ page_table=None,
158
+ sampling_params=None,
159
+ prompt_tokens=None,
160
+ output_tokens=None,
161
+ slot_remap=None,
162
+ reset=False,
163
+ ):
164
+ if page_table is None:
165
+ page_table = torch.tensor([[3, 4, 5], [6, 7, 8]], dtype=torch.int32)
166
+ return runtime.prepare(
167
+ torch.tensor([11, 0]),
168
+ torch.tensor(positions),
169
+ page_table,
170
+ sampling_params=sampling_params,
171
+ prompt_tokens=prompt_tokens,
172
+ output_tokens=output_tokens,
173
+ slot_remap=slot_remap,
174
+ reset_batch=reset,
175
+ )
176
+
177
+
178
+ def test_prepare_places_only_start_pos_active_rows_and_neutralizes_gap_sentinels():
179
+ runtime = make_runtime()
180
+ sampling = SamplingParams(
181
+ temperature=[0.8, 0.8],
182
+ top_k=[999, 7],
183
+ top_p=[0.9, 0.8],
184
+ seed=[-1, -1],
185
+ enable_log_probs=[False, False],
186
+ num_logprobs=[-2, -2],
187
+ )
188
+
189
+ prepared = prepare(
190
+ runtime,
191
+ positions=(-1, 4),
192
+ sampling_params=sampling,
193
+ prompt_tokens=torch.tensor([[10, -1], [20, 21]]),
194
+ output_tokens=[[30, -1], [40, 41]],
195
+ slot_remap=torch.tensor([1, 0]),
196
+ reset=True,
197
+ ).prepared_sampling
198
+
199
+ assert prepared is not None
200
+ assert prepared.active_mask == (False, True)
201
+ assert prepared.row_paths == ("inactive", "topk")
202
+ assert prepared.top_k == (1, 7)
203
+ assert prepared.top_p == pytest.approx((0.0, 0.8))
204
+ assert prepared.seeds == (None, None)
205
+ assert prepared.enable_log_probs == (False, False)
206
+ assert prepared.num_logprobs == (0, 0)
207
+ assert prepared.prompt_tokens.tolist() == [[-1, -1], [20, 21]]
208
+ assert prepared.output_tokens == [[-1, -1], [40, 41]]
209
+ assert prepared.slot_remap.tolist() == [1, 0]
210
+
211
+
212
+ def seeded_sampling(seed0, seed1=None):
213
+ return SamplingParams(
214
+ temperature=[0.8, 0.8],
215
+ top_k=[32, 32],
216
+ top_p=[0.95, 0.95],
217
+ seed=[seed0, seed1],
218
+ )
219
+
220
+
221
+ def stochastic_sampling(seed=None):
222
+ return SamplingParams(
223
+ temperature=[0.8, 0.8],
224
+ top_k=[32, 32],
225
+ top_p=[0.95, 0.95],
226
+ seed=seed,
227
+ )
228
+
229
+
230
+ def make_four_slot_seed_runtime():
231
+ seed_buffer = FakeLazySeedBuffer()
232
+ seed_buffer.source = torch.arange(4, dtype=torch.int64)
233
+ model = FakeModel(seed_buffer)
234
+ model.config.max_batch_size = 4
235
+ model.sampling.config.max_batch_size = 4
236
+ runtime = DecodeRuntime(
237
+ DecodeRuntimeConfig.resolve(
238
+ model=model,
239
+ output_reader=OutputReader(FakeMesh()),
240
+ lane_capacity=4,
241
+ page_table_layout=page_table_layout(),
242
+ device_sampling_enabled=True,
243
+ )
244
+ )
245
+ return runtime, seed_buffer
246
+
247
+
248
+ def four_slot_sampled_warmup(runtime):
249
+ return runtime.prepare(
250
+ torch.zeros(4, dtype=torch.long),
251
+ torch.zeros(4, dtype=torch.long),
252
+ torch.zeros((4, 8), dtype=torch.int32),
253
+ sampling_params=SamplingParams(
254
+ temperature=torch.ones(4),
255
+ top_k=torch.full((4,), 32, dtype=torch.int32),
256
+ top_p=torch.full((4,), 0.08),
257
+ seed=[11, 22, 33, 44],
258
+ ),
259
+ )
260
+
261
+
262
+ def _stub_compile_only_decode(runtime, monkeypatch, run_body):
263
+ monkeypatch.setattr(runtime, "_prepare_inputs_host", lambda prepared: "host")
264
+ monkeypatch.setattr(
265
+ runtime,
266
+ "_stage_inputs_and_kpt",
267
+ lambda host, prepared: (DecodeDeviceInputs(None, None, None, None), None),
268
+ )
269
+ monkeypatch.setattr(runtime, "_run_body", run_body)
270
+
271
+
272
+ def test_compile_only_sampled_decode_temporarily_admits_and_resets_fallback_seed_slots(monkeypatch, expect_error):
273
+ runtime, seed_buffer = make_four_slot_seed_runtime()
274
+ defaults = seed_buffer.source.clone()
275
+ prepared = four_slot_sampled_warmup(runtime)
276
+ during_compile = []
277
+
278
+ def run_body(*args, **kwargs):
279
+ during_compile.append(runtime._seed_state.snapshot())
280
+ assert not kwargs["count_tokens"]
281
+ assert not kwargs["advance_seeds"]
282
+ return object()
283
+
284
+ _stub_compile_only_decode(runtime, monkeypatch, run_body)
285
+
286
+ runtime.invoke(prepared, count_tokens=False)
287
+
288
+ assert during_compile[0].active_slots == (0, 1, 2, 3)
289
+ assert not during_compile[0].buffer_is_default
290
+ reset = runtime._seed_state.snapshot()
291
+ assert reset.active_slots == ()
292
+ assert reset.buffer_is_default
293
+ assert torch.equal(seed_buffer.source, defaults)
294
+ with expect_error(RuntimeError, "reset_batch=True"):
295
+ runtime._refresh_sampling_seeds(prepared)
296
+
297
+
298
+ def test_compile_only_sampled_decode_resets_fallback_seed_slots_after_failure(monkeypatch, expect_error):
299
+ runtime, seed_buffer = make_four_slot_seed_runtime()
300
+ defaults = seed_buffer.source.clone()
301
+ prepared = four_slot_sampled_warmup(runtime)
302
+
303
+ def run_body(*args, **kwargs):
304
+ assert runtime._seed_state.snapshot().active_slots == (0, 1, 2, 3)
305
+ raise RuntimeError("compile boom")
306
+
307
+ _stub_compile_only_decode(runtime, monkeypatch, run_body)
308
+ monkeypatch.setattr(runtime, "_release_or_retain_transient", lambda owned: [])
309
+
310
+ with expect_error(RuntimeError, "compile boom"):
311
+ runtime.invoke(prepared, count_tokens=False)
312
+
313
+ reset = runtime._seed_state.snapshot()
314
+ assert reset.active_slots == ()
315
+ assert reset.buffer_is_default
316
+ assert torch.equal(seed_buffer.source, defaults)
317
+
318
+
319
+ def test_runtime_seed_same_request_and_absolute_position_are_cardinality_independent():
320
+ first_buffer = FakeLazySeedBuffer()
321
+ first = make_runtime(seed_buffer=first_buffer)
322
+ first_prepared = prepare(first, positions=(41, -1), sampling_params=seeded_sampling(1234), reset=True)
323
+ first._refresh_sampling_seeds(first_prepared)
324
+
325
+ remapped_buffer = FakeLazySeedBuffer()
326
+ remapped = make_runtime(seed_buffer=remapped_buffer)
327
+ remapped_prepared = prepare(
328
+ remapped,
329
+ positions=(-1, 41),
330
+ sampling_params=seeded_sampling(None, 1234),
331
+ reset=True,
332
+ )
333
+ remapped._refresh_sampling_seeds(remapped_prepared)
334
+
335
+ assert int(first_buffer.updates[-1][0]) == int(remapped_buffer.updates[-1][1])
336
+ assert first._seed_state.snapshot().active == (True, False)
337
+ assert remapped._seed_state.snapshot().active == (False, True)
338
+
339
+
340
+ @pytest.mark.parametrize("seed", [1234, torch.tensor(1234)])
341
+ def test_runtime_scalar_seed_belongs_to_one_request_and_is_not_broadcast(seed):
342
+ seed_buffer = FakeLazySeedBuffer()
343
+ runtime = make_runtime(seed_buffer=seed_buffer)
344
+ sampling_params = SamplingParams(
345
+ temperature=[0.8, 0.8],
346
+ top_k=[32, 32],
347
+ top_p=[0.95, 0.95],
348
+ seed=seed,
349
+ )
350
+
351
+ prepared = prepare(runtime, positions=(17, 17), sampling_params=sampling_params, reset=True)
352
+ runtime._refresh_sampling_seeds(prepared)
353
+
354
+ assert prepared.prepared_sampling.seeds == (1234, None)
355
+ snapshot = runtime._seed_state.snapshot()
356
+ assert snapshot.request_seeds == (1234, None)
357
+ assert snapshot.active == (True, True)
358
+ assert snapshot.current_device_seeds[0] is not None
359
+ assert snapshot.current_device_seeds[1] is not None
360
+
361
+
362
+ @pytest.mark.parametrize("seed", [[111, 222], torch.tensor([111, 222])])
363
+ def test_runtime_vector_seed_remains_slot_indexed(seed):
364
+ runtime = make_runtime(seed_buffer=FakeLazySeedBuffer())
365
+ sampling_params = SamplingParams(
366
+ temperature=[0.8, 0.8],
367
+ top_k=[32, 32],
368
+ top_p=[0.95, 0.95],
369
+ seed=seed,
370
+ )
371
+
372
+ prepared = prepare(runtime, positions=(5, 5), sampling_params=sampling_params, reset=True)
373
+
374
+ assert prepared.prepared_sampling.seeds == (111, 222)
375
+
376
+
377
+ def test_runtime_simultaneous_equal_request_seeds_share_vllm_stream():
378
+ runtime = make_runtime(seed_buffer=FakeLazySeedBuffer())
379
+ prepared = prepare(
380
+ runtime,
381
+ positions=(5, 5),
382
+ sampling_params=seeded_sampling(77, 77),
383
+ reset=True,
384
+ )
385
+
386
+ runtime._refresh_sampling_seeds(prepared)
387
+
388
+ snapshot = runtime._seed_state.snapshot()
389
+ assert runtime._seed_manager.salt_duplicate_seeds is False
390
+ assert snapshot.request_seeds == (77, 77)
391
+ assert snapshot.salts == (0, 0)
392
+ assert snapshot.current_device_seeds[0] == snapshot.current_device_seeds[1]
393
+
394
+
395
+ def test_runtime_vllm_uniform_seed_is_deterministic_across_concurrent_slots():
396
+ seed_buffer = FakeLazySeedBuffer()
397
+ seed_buffer.source = torch.arange(32, dtype=torch.int64)
398
+ model = FakeModel(seed_buffer)
399
+ model.config.max_batch_size = 32
400
+ model.sampling.config.max_batch_size = 32
401
+ runtime = DecodeRuntime(
402
+ DecodeRuntimeConfig.resolve(
403
+ model=model,
404
+ output_reader=OutputReader(FakeMesh()),
405
+ lane_capacity=32,
406
+ page_table_layout=page_table_layout(),
407
+ device_sampling_enabled=True,
408
+ )
409
+ )
410
+ prepared = runtime.prepare(
411
+ torch.zeros(32, dtype=torch.long),
412
+ torch.zeros(32, dtype=torch.long),
413
+ torch.zeros((32, 8), dtype=torch.int32),
414
+ sampling_params=SamplingParams(
415
+ temperature=torch.ones(32),
416
+ top_k=torch.full((32,), 32, dtype=torch.int32),
417
+ top_p=torch.full((32,), 0.95),
418
+ seed=[1234] * 32,
419
+ ),
420
+ reset_batch=True,
421
+ )
422
+
423
+ runtime._refresh_sampling_seeds(prepared)
424
+
425
+ snapshot = runtime._seed_state.snapshot()
426
+ assert runtime._seed_manager.salt_duplicate_seeds is False
427
+ assert snapshot.salts == (0,) * 32
428
+ assert len(set(snapshot.current_device_seeds)) == 1
429
+
430
+
431
+ def test_runtime_slot_remap_moves_complete_seed_stream_before_refresh():
432
+ seed_buffer = FakeLazySeedBuffer()
433
+ runtime = make_runtime(seed_buffer=seed_buffer)
434
+ initial = prepare(
435
+ runtime,
436
+ positions=(5, -1),
437
+ sampling_params=seeded_sampling(77),
438
+ reset=True,
439
+ )
440
+ runtime._refresh_sampling_seeds(initial)
441
+ original_device_seed = int(seed_buffer.updates[-1][0])
442
+ original_state = runtime._seed_state.snapshot()
443
+
444
+ moved = prepare(
445
+ runtime,
446
+ positions=(-1, 5),
447
+ sampling_params=seeded_sampling(None, 77),
448
+ slot_remap=torch.tensor([0, 0]),
449
+ reset=False,
450
+ )
451
+ runtime._refresh_sampling_seeds(moved)
452
+
453
+ state = runtime._seed_state.snapshot()
454
+ assert state.active == (False, True)
455
+ assert state.request_seeds == (None, 77)
456
+ assert state.token_counters[1] == original_state.token_counters[0]
457
+ assert int(seed_buffer.updates[-1][1]) == original_device_seed
458
+
459
+
460
+ def test_runtime_seed_changes_with_request_seed_and_decode_position():
461
+ seed_buffer = FakeLazySeedBuffer()
462
+ runtime = make_runtime(seed_buffer=seed_buffer)
463
+
464
+ runtime._refresh_sampling_seeds(
465
+ prepare(runtime, positions=(7, -1), sampling_params=seeded_sampling(101), reset=True)
466
+ )
467
+ position_7 = int(seed_buffer.updates[-1][0])
468
+ runtime._refresh_sampling_seeds(
469
+ prepare(runtime, positions=(8, -1), sampling_params=seeded_sampling(101), reset=False)
470
+ )
471
+ position_8 = int(seed_buffer.updates[-1][0])
472
+ runtime._refresh_sampling_seeds(
473
+ prepare(runtime, positions=(7, -1), sampling_params=seeded_sampling(202), reset=True)
474
+ )
475
+ different_request = int(seed_buffer.updates[-1][0])
476
+
477
+ assert len({position_7, position_8, different_request}) == 3
478
+
479
+
480
+ def test_runtime_explicit_seed_absolute_position_is_stable_across_reset_boundaries():
481
+ seed_buffer = FakeLazySeedBuffer()
482
+ runtime = make_runtime(seed_buffer=seed_buffer)
483
+ request = seeded_sampling(909)
484
+
485
+ runtime._refresh_sampling_seeds(prepare(runtime, positions=(19, -1), sampling_params=request, reset=True))
486
+ original = seed_buffer.updates[-1].clone()
487
+ runtime._refresh_sampling_seeds(prepare(runtime, positions=(20, -1), sampling_params=request, reset=False))
488
+ continued = seed_buffer.updates[-1].clone()
489
+ runtime._refresh_sampling_seeds(prepare(runtime, positions=(19, -1), sampling_params=request, reset=True))
490
+ restarted = seed_buffer.updates[-1].clone()
491
+ runtime._refresh_sampling_seeds(prepare(runtime, positions=(20, -1), sampling_params=request, reset=True))
492
+ resumed = seed_buffer.updates[-1].clone()
493
+
494
+ assert torch.equal(original, restarted)
495
+ assert torch.equal(continued, resumed)
496
+
497
+
498
+ def test_runtime_seed_refreshes_before_eager_model_invocation(monkeypatch):
499
+ events = []
500
+ seed_buffer = FakeLazySeedBuffer()
501
+ original_update = seed_buffer.update
502
+
503
+ def record_update(source):
504
+ events.append("seed")
505
+ original_update(source)
506
+
507
+ seed_buffer.update = record_update
508
+ runtime = make_runtime(seed_buffer=seed_buffer)
509
+ prepared = prepare(runtime, positions=(3, -1), sampling_params=seeded_sampling(77), reset=True)
510
+ monkeypatch.setattr(runtime, "_prepare_inputs_host", lambda prepared: object())
511
+ monkeypatch.setattr(
512
+ runtime,
513
+ "_stage_inputs_and_kpt",
514
+ lambda host, prepared: (DecodeDeviceInputs(None, None, None, None), None),
515
+ )
516
+
517
+ def run_body(*args, **kwargs):
518
+ events.append("invoke")
519
+ return object()
520
+
521
+ monkeypatch.setattr(runtime, "_run_body", run_body)
522
+
523
+ runtime.invoke(prepared)
524
+
525
+ assert events[-1] == "invoke"
526
+ assert events[:-1]
527
+ assert set(events[:-1]) == {"seed"}
528
+
529
+
530
+ def test_runtime_trace_captures_stable_seed_handle_and_refreshes_before_replay(monkeypatch):
531
+ events = []
532
+ seed_buffer = FakeLazySeedBuffer()
533
+ original_update = seed_buffer.update
534
+
535
+ def record_update(source):
536
+ events.append("seed")
537
+ original_update(source)
538
+
539
+ seed_buffer.update = record_update
540
+ runtime = make_runtime(seed_buffer=seed_buffer)
541
+ prepared = prepare(runtime, positions=(12, -1), sampling_params=seeded_sampling(88), reset=True)
542
+ monkeypatch.setattr(runtime, "_prepare_inputs_host", lambda prepared: object())
543
+ monkeypatch.setattr(
544
+ runtime,
545
+ "_stage_inputs_and_kpt",
546
+ lambda host, prepared: (DecodeDeviceInputs(None, None, None, None), None),
547
+ )
548
+
549
+ persistent = runtime.capture_plan(prepared).prepare_inputs()
550
+ assert persistent.seed_buffer is seed_buffer.get_device_buffer()
551
+ sampling = prepared.prepared_sampling
552
+ assert sampling is not None
553
+ persistent = dataclasses.replace(
554
+ persistent,
555
+ kpt_signature=[(sampling.top_k, sampling.top_p, sampling.temperature)],
556
+ )
557
+ runtime.refresh_trace(persistent, prepared, SimpleNamespace(full=False, page_table=False))
558
+ events.append("replay")
559
+
560
+ assert events[-1] == "replay"
561
+ assert events[:-1]
562
+ assert set(events[:-1]) == {"seed"}
563
+
564
+
565
+ def test_runtime_seed_handling_does_not_mutate_sampling_params():
566
+ seed_buffer = FakeLazySeedBuffer()
567
+ runtime = make_runtime(seed_buffer=seed_buffer)
568
+ sampling_params = seeded_sampling(11, 22)
569
+ before = dataclasses.asdict(sampling_params)
570
+
571
+ prepared = prepare(runtime, positions=(4, 9), sampling_params=sampling_params, reset=True)
572
+ runtime._refresh_sampling_seeds(prepared)
573
+
574
+ assert dataclasses.asdict(sampling_params) == before
575
+
576
+
577
+ def test_runtime_unseeded_stream_varies_and_same_absolute_position_is_idempotent():
578
+ seed_buffer = FakeLazySeedBuffer()
579
+ defaults = seed_buffer.source.clone()
580
+ runtime = make_runtime(seed_buffer=seed_buffer)
581
+
582
+ unseeded = prepare(runtime, positions=(5, -1), sampling_params=stochastic_sampling(), reset=True)
583
+ runtime._refresh_sampling_seeds(unseeded)
584
+ first = seed_buffer.updates[-1].clone()
585
+ first_state = runtime._seed_state.snapshot()
586
+ assert int(first[0]) != int(defaults[0])
587
+ assert int(first[1]) == int(defaults[1])
588
+
589
+ runtime._refresh_sampling_seeds(dataclasses.replace(unseeded, reset_batch=False))
590
+ repeated = seed_buffer.updates[-1].clone()
591
+ assert torch.equal(repeated, first)
592
+
593
+ advanced = prepare(
594
+ runtime,
595
+ positions=(6, -1),
596
+ sampling_params=stochastic_sampling(),
597
+ reset=False,
598
+ )
599
+ runtime._refresh_sampling_seeds(advanced)
600
+ snapshot = runtime._seed_state.snapshot()
601
+ assert snapshot.request_seeds == (None, None)
602
+ assert snapshot.active == (True, False)
603
+ assert snapshot.token_counters[0] == 2
604
+ assert snapshot.unseeded_rng_states[0] != first_state.unseeded_rng_states[0]
605
+
606
+
607
+ def test_runtime_seed_change_requires_reset_and_preserves_unseeded_survivor(expect_error):
608
+ seed_buffer = FakeLazySeedBuffer()
609
+ runtime = make_runtime(seed_buffer=seed_buffer)
610
+ unseeded = prepare(
611
+ runtime,
612
+ positions=(10, 10),
613
+ sampling_params=stochastic_sampling(seed=[None, None]),
614
+ reset=True,
615
+ )
616
+ runtime._refresh_sampling_seeds(unseeded)
617
+ initial = runtime._seed_state.snapshot()
618
+
619
+ mixed = prepare(
620
+ runtime,
621
+ positions=(11, 11),
622
+ sampling_params=stochastic_sampling(seed=[None, 42]),
623
+ reset=False,
624
+ )
625
+ with expect_error(RuntimeError, "reset_batch=True"):
626
+ runtime._refresh_sampling_seeds(mixed)
627
+
628
+ runtime._refresh_sampling_seeds(dataclasses.replace(mixed, reset_batch=True))
629
+ admitted = runtime._seed_state.snapshot()
630
+ assert admitted.request_seeds == (None, 42)
631
+ assert admitted.active == (True, True)
632
+ assert admitted.token_counters[0] == initial.token_counters[0] + 1
633
+
634
+ continued = dataclasses.replace(mixed, start_pos=torch.tensor([12, 12]))
635
+ runtime._refresh_sampling_seeds(continued)
636
+ continued_state = runtime._seed_state.snapshot()
637
+ assert continued_state.token_counters[0] == admitted.token_counters[0] + 1
638
+ assert continued_state.request_seeds == (None, 42)
639
+
640
+
641
+ def test_runtime_inactive_peer_is_cleaned_up_without_readmitting_survivor():
642
+ seed_buffer = FakeLazySeedBuffer()
643
+ defaults = seed_buffer.source.clone()
644
+ runtime = make_runtime(seed_buffer=seed_buffer)
645
+ initial = prepare(
646
+ runtime,
647
+ positions=(20, 20),
648
+ sampling_params=stochastic_sampling(seed=[None, 42]),
649
+ reset=True,
650
+ )
651
+ runtime._refresh_sampling_seeds(initial)
652
+ initial_state = runtime._seed_state.snapshot()
653
+
654
+ peer_left = prepare(
655
+ runtime,
656
+ positions=(21, -1),
657
+ sampling_params=stochastic_sampling(seed=[None, 42]),
658
+ reset=False,
659
+ )
660
+ runtime._refresh_sampling_seeds(peer_left)
661
+
662
+ state = runtime._seed_state.snapshot()
663
+ assert state.active == (True, False)
664
+ assert state.request_seeds == (None, None)
665
+ assert state.token_counters[0] == initial_state.token_counters[0] + 1
666
+ assert int(seed_buffer.updates[-1][1]) == int(defaults[1])
667
+
668
+
669
+ @pytest.mark.parametrize(
670
+ ("method", "expected"),
671
+ (
672
+ (
673
+ DecodeRuntime.prepare,
674
+ (
675
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
676
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
677
+ ("start_pos", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
678
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
679
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
680
+ ("prompt_tokens", inspect.Parameter.KEYWORD_ONLY, None),
681
+ ("output_tokens", inspect.Parameter.KEYWORD_ONLY, None),
682
+ ("slot_remap", inspect.Parameter.KEYWORD_ONLY, None),
683
+ ("reset_batch", inspect.Parameter.KEYWORD_ONLY, False),
684
+ ),
685
+ ),
686
+ (
687
+ DecodeRuntime.read_decode_output,
688
+ (
689
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
690
+ ("tt_out", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
691
+ ("async_read", inspect.Parameter.KEYWORD_ONLY, False),
692
+ ),
693
+ ),
694
+ (
695
+ DecodeRuntime.process_decode_output_host,
696
+ (
697
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
698
+ ("tt_out", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
699
+ ("is_tokens", inspect.Parameter.KEYWORD_ONLY, False),
700
+ ),
701
+ ),
702
+ ),
703
+ )
704
+ def test_runtime_api_signatures_are_exact(method, expected):
705
+ parameters = inspect.signature(method).parameters.values()
706
+
707
+ assert tuple((parameter.name, parameter.kind, parameter.default) for parameter in parameters) == expected
708
+
709
+
710
+ def test_runtime_api_preserves_positional_prefixes_and_rejects_extra_fields(expect_error):
711
+ cases = (
712
+ (DecodeRuntime.prepare, ("tokens", "start_pos", "page_table"), "sampling_params"),
713
+ (DecodeRuntime.read_decode_output, ("tt_out",), "async_read"),
714
+ (DecodeRuntime.process_decode_output_host, ("tt_out",), "is_tokens"),
715
+ )
716
+
717
+ for method, positional_prefix, keyword_name in cases:
718
+ signature = inspect.signature(method)
719
+ signature.bind(None, *positional_prefix, **{keyword_name: True})
720
+ with expect_error(TypeError, "too many positional arguments"):
721
+ signature.bind(None, *positional_prefix, True)
722
+ with expect_error(TypeError, "unexpected keyword argument 'unknown'"):
723
+ signature.bind(None, *positional_prefix, unknown=True)
724
+
725
+ assert inspect.get_annotations(DecodeRuntime.prepare, eval_str=True)["sampling_params"] is Any
726
+
727
+
728
+ def test_config_resolves_canonical_static_capabilities_and_is_frozen(expect_error):
729
+ mesh = FakeMesh()
730
+ model = FakeModel()
731
+ config = DecodeRuntimeConfig.resolve(
732
+ model=model,
733
+ output_reader=OutputReader(mesh),
734
+ lane_capacity=2,
735
+ page_table_layout=page_table_layout(),
736
+ device_sampling_enabled=True,
737
+ )
738
+
739
+ assert config.cluster_shape == (1, 1)
740
+ assert config.num_devices == 1
741
+ assert config.vocab_size == 8
742
+ assert config.allow_force_argmax
743
+ assert config.max_device_top_k == 32
744
+ assert config.sampling_batch_size == 2
745
+ assert config.sampling_state_controller is None
746
+ assert config.sampling_state is None
747
+ assert config.position_feedback_capable
748
+ with expect_error(dataclasses.FrozenInstanceError, "cannot assign to field"):
749
+ config.lane_capacity = 1
750
+ with expect_error(TypeError, "DecodeRuntimeConfig"):
751
+ DecodeRuntime(config=None)
752
+
753
+
754
+ def test_config_rejects_inconsistent_collaborators_and_dimensions(expect_error):
755
+ mesh = FakeMesh()
756
+ model = FakeModel()
757
+ model.config.mesh_device = mesh
758
+ with expect_error(ValueError, "same mesh_device"):
759
+ DecodeRuntimeConfig.resolve(
760
+ model=model,
761
+ output_reader=OutputReader(FakeMesh()),
762
+ lane_capacity=2,
763
+ page_table_layout=page_table_layout(),
764
+ device_sampling_enabled=True,
765
+ )
766
+ with expect_error(ValueError, "positive integer"):
767
+ DecodeRuntimeConfig.resolve(
768
+ model=model,
769
+ output_reader=OutputReader(mesh),
770
+ lane_capacity=True,
771
+ page_table_layout=page_table_layout(),
772
+ device_sampling_enabled=True,
773
+ )
774
+ with expect_error(TypeError, "PageTableLayout"):
775
+ DecodeRuntimeConfig.resolve(
776
+ model=model,
777
+ output_reader=OutputReader(mesh),
778
+ lane_capacity=2,
779
+ page_table_layout=SimpleNamespace(raw_capacity_width=8, decode_width=8, block_size=32),
780
+ device_sampling_enabled=True,
781
+ )
782
+
783
+
784
+ def test_layout_replacement_is_immutable_and_bounded(expect_error):
785
+ runtime = make_runtime()
786
+ original = runtime.config
787
+ replacement = page_table_layout(raw_width=4)
788
+
789
+ runtime.configure_page_table_layout(replacement)
790
+
791
+ assert runtime.config is not original
792
+ assert original.page_table_layout.raw_capacity_width == 8
793
+ assert runtime.config.page_table_layout is replacement
794
+ with expect_error(ValueError, "block_size"):
795
+ runtime.configure_page_table_layout(page_table_layout(raw_width=4, block_size=16))
796
+ with expect_error(ValueError, "ceiling"):
797
+ runtime.configure_page_table_layout(page_table_layout(raw_width=9))
798
+ with expect_error(ValueError, "decode width"):
799
+ runtime.configure_page_table_layout(PageTableLayout(32, 4, 8, 1024))
800
+
801
+
802
+ def test_sampling_admission_follows_resolved_configuration(expect_error):
803
+ runtime = make_runtime(sampling=False)
804
+ with expect_error(ValueError, "device sampling is disabled"):
805
+ prepare(runtime, sampling_params=greedy_sampling())
806
+
807
+
808
+ def test_feedback_and_sampling_path_follow_resolved_capabilities():
809
+ argmax_runtime = make_runtime()
810
+ argmax_prepared = prepare(argmax_runtime, sampling_params=greedy_sampling())
811
+ assert argmax_prepared.device_feedback
812
+ assert argmax_prepared.sampling_path == "argmax"
813
+
814
+ model = FakeModel()
815
+ model.increment_positions = None
816
+ mesh = FakeMesh()
817
+ no_feedback = DecodeRuntime(
818
+ DecodeRuntimeConfig.resolve(
819
+ model=model,
820
+ output_reader=OutputReader(mesh),
821
+ lane_capacity=2,
822
+ page_table_layout=page_table_layout(),
823
+ device_sampling_enabled=True,
824
+ force_greedy_top_k=True,
825
+ )
826
+ )
827
+ prepared = prepare(no_feedback, sampling_params=greedy_sampling())
828
+
829
+ assert not prepared.device_feedback
830
+ assert prepared.sampling_path == "topk"
831
+
832
+
833
+ def test_single_and_multi_device_logits_conversion(monkeypatch):
834
+ single = make_runtime()
835
+ logits = torch.arange(16, dtype=torch.float32).reshape(1, 1, 2, 8)
836
+ monkeypatch.setattr(ttnn, "to_torch", lambda value: logits)
837
+ converted, _ = single._normalize_host_output("single-device", is_tokens=False)
838
+ assert converted.shape == (2, 1, 8)
839
+
840
+ mesh = SimpleNamespace(shape=(1, 2))
841
+ model = FakeModel()
842
+ model.num_devices = 2
843
+ multi = DecodeRuntime(
844
+ DecodeRuntimeConfig.resolve(
845
+ model=model,
846
+ output_reader=OutputReader(mesh),
847
+ lane_capacity=2,
848
+ page_table_layout=page_table_layout(),
849
+ device_sampling_enabled=True,
850
+ )
851
+ )
852
+ calls = []
853
+ monkeypatch.setattr(
854
+ decode_module,
855
+ "_concat_host_output",
856
+ lambda value, shape: calls.append((value, shape)) or logits,
857
+ )
858
+ converted, _ = multi._normalize_host_output("multi-device", is_tokens=False)
859
+ assert converted.shape == (2, 1, 8)
860
+ assert calls == [("multi-device", (1, 2))]
861
+
862
+
863
+ def test_signatures_expose_ordered_material_and_separate_types():
864
+ runtime = make_runtime()
865
+ prepared = prepare(runtime, sampling_params=greedy_sampling())
866
+
867
+ program = runtime.program_signature(prepared)
868
+ trace = runtime.trace_signature(prepared)
869
+
870
+ assert isinstance(program, DecodeProgramSignature)
871
+ assert isinstance(trace, DecodeTraceSignature)
872
+ assert program.key_material() == (
873
+ ("operation", "decode"),
874
+ ("batch_size", 2),
875
+ ("page_table_width", 8),
876
+ ("sampling_path", "argmax"),
877
+ ("device_feedback", True),
878
+ )
879
+ assert trace.key_material() == program.key_material()
880
+ assert runtime.program_signature(prepare(runtime)).sampling_path == "logits"
881
+
882
+
883
+ def test_signature_tracks_native_penalty_and_sampled_logprob_program_modes():
884
+ runtime = make_runtime()
885
+ sampling = SamplingParams(
886
+ temperature=[0.0, 0.0],
887
+ top_k=[1, 1],
888
+ top_p=[1.0, 1.0],
889
+ presence_penalty=[0.5, 0.0],
890
+ enable_log_probs=[True, False],
891
+ num_logprobs=[0, -2],
892
+ )
893
+
894
+ prepared = prepare(runtime, sampling_params=sampling)
895
+ native = prepared.prepared_sampling
896
+ signature = runtime.program_signature(prepared)
897
+
898
+ assert native is not None
899
+ assert native.logprob_modes == ("sampled_token", "none")
900
+ assert native.penalties_enabled
901
+ assert native.log_probs_enabled
902
+ assert prepared.sampling_path == "topk"
903
+ assert signature.penalties_enabled
904
+ assert signature.logprobs_enabled
905
+ assert signature.key_material()[-2:] == (
906
+ ("penalties_enabled", True),
907
+ ("logprobs_enabled", True),
908
+ )
909
+
910
+
911
+ def test_configured_topk_policy_is_not_collapsed_to_argmax_by_greedy_temperature():
912
+ runtime = make_runtime(force_greedy_top_k=True)
913
+ sampling = SamplingParams(temperature=[0.0, 0.0], top_k=[32, 32], top_p=[0.08, 0.08])
914
+
915
+ prepared = prepare(runtime, sampling_params=sampling)
916
+
917
+ assert prepared.sampling_path == "topk"
918
+
919
+
920
+ def test_unconfigured_topk_values_use_argmax_for_greedy_temperature():
921
+ runtime = make_runtime()
922
+ sampling = SamplingParams(temperature=[0.0, 0.0], top_k=[32, 32], top_p=[0.08, 0.08])
923
+
924
+ assert prepare(runtime, sampling_params=sampling).sampling_path == "argmax"
925
+
926
+
927
+ def test_sampling_params_are_prepared_once_and_reused_for_kpt(monkeypatch):
928
+ runtime = make_runtime(force_greedy_top_k=True)
929
+ calls = []
930
+ formatter = decode_module.prepare_sampling_params
931
+
932
+ def formatter_spy(*args, **kwargs):
933
+ calls.append((args, kwargs))
934
+ return formatter(*args, **kwargs)
935
+
936
+ monkeypatch.setattr(
937
+ decode_module,
938
+ "prepare_sampling_params",
939
+ formatter_spy,
940
+ )
941
+ prepared = prepare(runtime, sampling_params=greedy_sampling())
942
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh: "mapper")
943
+ monkeypatch.setattr(ttnn, "from_torch", lambda value, **kwargs: value)
944
+
945
+ runtime._make_host_kpt(prepared)
946
+ assert len(calls) == 1
947
+
948
+
949
+ def test_tile_padded_sampler_preserves_one_semantic_lane_and_neutralizes_inactive_rows(monkeypatch):
950
+ mesh = FakeMesh()
951
+ model = FakeModel()
952
+ model.config.max_batch_size = 1
953
+ model.sampling.config.max_batch_size = 32
954
+ runtime = DecodeRuntime(
955
+ DecodeRuntimeConfig.resolve(
956
+ model=model,
957
+ output_reader=OutputReader(mesh),
958
+ lane_capacity=1,
959
+ page_table_layout=page_table_layout(),
960
+ device_sampling_enabled=True,
961
+ force_greedy_top_k=True,
962
+ )
963
+ )
964
+ prepared = runtime.prepare(
965
+ torch.tensor([11]),
966
+ torch.tensor([0]),
967
+ torch.tensor([[3, 4, 5]], dtype=torch.int32),
968
+ sampling_params=SamplingParams(
969
+ temperature=0.7,
970
+ top_k=32,
971
+ top_p=0.08,
972
+ presence_penalty=0.25,
973
+ frequency_penalty=0.5,
974
+ repetition_penalty=1.2,
975
+ seed=17,
976
+ enable_log_probs=True,
977
+ num_logprobs=1,
978
+ ),
979
+ )
980
+ sampling = prepared.prepared_sampling
981
+
982
+ assert runtime.config.lane_capacity == 1
983
+ assert runtime.config.sampling_batch_size == 32
984
+ assert sampling.batch_size == 32
985
+ assert sampling.active_rows == 1
986
+ assert sampling.active_mask == (True,) + (False,) * 31
987
+ assert sampling.top_k == (32,) + (1,) * 31
988
+ assert sampling.top_p == pytest.approx((0.08,) + (0.0,) * 31)
989
+ assert sampling.temperature == pytest.approx((1.0 / 0.7,) + (1.0,) * 31)
990
+ assert sampling.seeds == (17,) + (None,) * 31
991
+ assert sampling.presence_penalty == pytest.approx((0.25,) + (0.0,) * 31)
992
+ assert sampling.frequency_penalty == pytest.approx((0.5,) + (0.0,) * 31)
993
+ assert sampling.repetition_penalty == pytest.approx((1.2,) + (1.0,) * 31)
994
+ assert sampling.enable_log_probs == (True,) + (False,) * 31
995
+ assert sampling.num_logprobs == (1,) + (0,) * 31
996
+ assert sampling.logprob_modes[1:] == ("none",) * 31
997
+
998
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh_device: "mapper")
999
+ monkeypatch.setattr(ttnn, "from_torch", lambda value, **kwargs: value)
1000
+ k, p, temperature = runtime._make_host_kpt(prepared)
1001
+
1002
+ assert tuple(k.shape) == tuple(p.shape) == tuple(temperature.shape) == (32,)
1003
+ assert k.tolist() == [32] + [1] * 31
1004
+ assert p.tolist() == pytest.approx([0.08] + [0.0] * 31)
1005
+ assert temperature.tolist() == pytest.approx([1.0 / 0.7] + [1.0] * 31)
1006
+
1007
+
1008
+ def test_normalization_preserves_feedback_lookahead_and_inactive_convention():
1009
+ runtime = make_runtime()
1010
+ page_table = torch.tensor([[10, 11, 99], [20, 21, 98]], dtype=torch.int64)
1011
+
1012
+ prepared = prepare(
1013
+ runtime,
1014
+ positions=(31, -1),
1015
+ page_table=page_table,
1016
+ sampling_params=greedy_sampling(),
1017
+ )
1018
+
1019
+ assert prepared.page_table.dtype == torch.int32
1020
+ assert prepared.page_table.shape == (2, 8)
1021
+ assert prepared.page_table[0].tolist() == [10, 11, 0, 0, 0, 0, 0, 0]
1022
+ assert prepared.page_table[1].tolist() == [0, 0, 0, 0, 0, 0, 0, 0]
1023
+
1024
+
1025
+ def test_normalization_reuses_equal_source_with_same_copy_counts():
1026
+ runtime = make_runtime()
1027
+ first = prepare(runtime, positions=(0, -1))
1028
+ second = prepare(runtime, positions=(1, -1))
1029
+
1030
+ assert second.page_table is first.page_table
1031
+
1032
+
1033
+ def test_normalization_cache_detects_in_place_source_mutation():
1034
+ runtime = make_runtime()
1035
+ source = torch.tensor([[3, 4], [6, 7]], dtype=torch.int32)
1036
+ first = prepare(runtime, positions=(0, -1), page_table=source)
1037
+ source[0, 0] = 9
1038
+ second = prepare(runtime, positions=(0, -1), page_table=source)
1039
+
1040
+ assert second.page_table is not first.page_table
1041
+ assert second.page_table[0, 0].item() == 9
1042
+
1043
+
1044
+ def test_normalization_cache_misses_when_copy_counts_or_feedback_change():
1045
+ runtime = make_runtime()
1046
+ source = torch.tensor([[3, 4], [6, 7]], dtype=torch.int32)
1047
+ one_block = prepare(runtime, positions=(0, -1), page_table=source)
1048
+ two_blocks = prepare(runtime, positions=(32, -1), page_table=source)
1049
+ no_feedback = prepare(runtime, positions=(31, -1), page_table=source)
1050
+ with_feedback = prepare(
1051
+ runtime,
1052
+ positions=(31, -1),
1053
+ page_table=source,
1054
+ sampling_params=greedy_sampling(),
1055
+ )
1056
+
1057
+ assert two_blocks.page_table is not one_block.page_table
1058
+ assert no_feedback.page_table[0, 1].item() == 0
1059
+ assert with_feedback.page_table is not no_feedback.page_table
1060
+ assert with_feedback.page_table[0, 1].item() == 4
1061
+
1062
+
1063
+ def test_fixed_capacity_and_page_table_capacity_are_validated(expect_error):
1064
+ runtime = make_runtime()
1065
+ with expect_error(ValueError, "must equal lane capacity"):
1066
+ runtime.prepare(torch.tensor([1]), torch.tensor([0]), torch.tensor([[1]]))
1067
+ with expect_error(ValueError, "batches must match"):
1068
+ runtime.prepare(
1069
+ torch.tensor([1, 2]),
1070
+ torch.tensor([0]),
1071
+ torch.tensor([[1], [2]]),
1072
+ )
1073
+ with expect_error(ValueError, "paged-KV capacity"):
1074
+ prepare(runtime, positions=(8 * 32, -1))
1075
+ with expect_error(ValueError, "too narrow"):
1076
+ prepare(
1077
+ runtime,
1078
+ positions=(64, -1),
1079
+ page_table=torch.tensor([[1, 2], [0, 0]], dtype=torch.int32),
1080
+ )
1081
+
1082
+
1083
+ def test_preparation_tracks_first_used_page_change_reset_and_ignores_unused_tail():
1084
+ runtime = make_runtime()
1085
+ first = prepare(runtime, positions=(0, -1), reset=True)
1086
+ assert first.page_table_changed
1087
+ assert first.reset_batch
1088
+
1089
+ runtime.note_submitted(first)
1090
+ same_semantics = prepare(
1091
+ runtime,
1092
+ positions=(0, -1),
1093
+ page_table=torch.tensor([[3, 90, 91], [88, 87, 86]], dtype=torch.int32),
1094
+ )
1095
+ assert not same_semantics.page_table_changed
1096
+ assert not same_semantics.reset_batch
1097
+
1098
+ changed = prepare(
1099
+ runtime,
1100
+ positions=(0, -1),
1101
+ page_table=torch.tensor([[4, 90, 91], [88, 87, 86]], dtype=torch.int32),
1102
+ )
1103
+ assert changed.page_table_changed
1104
+
1105
+
1106
+ def test_submission_state_tracks_last_table_despite_stale_prepare_change_hint():
1107
+ runtime = make_runtime()
1108
+ baseline = prepare(
1109
+ runtime,
1110
+ positions=(0, -1),
1111
+ page_table=torch.tensor([[3], [0]], dtype=torch.int32),
1112
+ )
1113
+ runtime.note_submitted(baseline)
1114
+
1115
+ changed = prepare(
1116
+ runtime,
1117
+ positions=(0, -1),
1118
+ page_table=torch.tensor([[4], [0]], dtype=torch.int32),
1119
+ )
1120
+ back_to_baseline = prepare(
1121
+ runtime,
1122
+ positions=(0, -1),
1123
+ page_table=torch.tensor([[3], [0]], dtype=torch.int32),
1124
+ )
1125
+ assert changed.page_table_changed
1126
+ assert not back_to_baseline.page_table_changed
1127
+
1128
+ runtime.note_submitted(changed)
1129
+ runtime.note_submitted(back_to_baseline)
1130
+
1131
+ assert not prepare(runtime, page_table=back_to_baseline.page_table).page_table_changed
1132
+ assert prepare(
1133
+ runtime,
1134
+ positions=(0, -1),
1135
+ page_table=torch.tensor([[4], [0]], dtype=torch.int32),
1136
+ ).page_table_changed
1137
+
1138
+
1139
+ def test_capture_plan_describes_full_step_refresh_and_typed_persistent_inputs(monkeypatch):
1140
+ runtime = make_runtime()
1141
+ prepared = prepare(runtime, sampling_params=greedy_sampling())
1142
+ device = DecodeDeviceInputs("tokens", "positions", "rotary", "page_table")
1143
+ monkeypatch.setattr(runtime, "_prepare_inputs_host", lambda request: "host")
1144
+ monkeypatch.setattr(runtime, "_stage_inputs_and_kpt", lambda host, request: (device, "kpt"))
1145
+
1146
+ def run_body(
1147
+ inputs,
1148
+ prepared,
1149
+ kpt,
1150
+ *,
1151
+ device_feedback,
1152
+ count_tokens=True,
1153
+ advance_seeds=True,
1154
+ ):
1155
+ return "captured"
1156
+
1157
+ monkeypatch.setattr(runtime, "_run_body", run_body)
1158
+
1159
+ plan = runtime.capture_plan(prepared)
1160
+ persistent = plan.prepare_inputs()
1161
+
1162
+ assert persistent.device_inputs is device
1163
+ assert persistent.kpt == "kpt"
1164
+ sampling = prepared.prepared_sampling
1165
+ assert sampling is not None
1166
+ assert persistent.kpt_signature == [(sampling.top_k, sampling.top_p, sampling.temperature)]
1167
+ assert plan.capture(persistent) == "captured"
1168
+ assert plan.refresh_policy.every_replay == ("sampling",)
1169
+ assert plan.refresh_policy.full_on_batch_reset
1170
+ assert plan.refresh_policy.full_on_graph_switch
1171
+ assert plan.refresh_policy.full_without_device_feedback
1172
+ assert plan.refresh_policy.refresh_page_table_on_change
1173
+
1174
+
1175
+ def test_trace_refresh_skips_unchanged_sampling_values(monkeypatch):
1176
+ runtime = make_runtime(force_greedy_top_k=True)
1177
+ prepared = prepare(runtime, sampling_params=greedy_sampling())
1178
+ sampling = prepared.prepared_sampling
1179
+ assert sampling is not None
1180
+ persistent = DecodePersistentInputs(
1181
+ device_inputs=DecodeDeviceInputs("tokens", "positions", "rotary", "page_table"),
1182
+ kpt="kpt",
1183
+ kpt_signature=[(sampling.top_k, sampling.top_p, sampling.temperature)],
1184
+ )
1185
+
1186
+ def fail_refresh_kpt(device_kpt, prepared):
1187
+ pytest.fail("unchanged KPT was refreshed")
1188
+
1189
+ monkeypatch.setattr(runtime, "_refresh_kpt", fail_refresh_kpt)
1190
+
1191
+ runtime.refresh_trace(
1192
+ persistent,
1193
+ prepared,
1194
+ SimpleNamespace(full=False, page_table=False),
1195
+ )
1196
+
1197
+
1198
+ def test_eager_invoke_returns_owned_result_and_advances_submission_state(monkeypatch):
1199
+ runtime = make_runtime()
1200
+ prepared = prepare(runtime)
1201
+ device = DecodeDeviceInputs("tokens", "positions", "rotary", "page_table")
1202
+ calls = []
1203
+ monkeypatch.setattr(runtime, "_prepare_inputs_host", lambda request: "host")
1204
+ monkeypatch.setattr(runtime, "_stage_inputs_and_kpt", lambda host, request: (device, None))
1205
+ monkeypatch.setattr(
1206
+ runtime,
1207
+ "_run_body",
1208
+ lambda inputs, prepared, kpt, *, device_feedback, **kwargs: calls.append(device_feedback) or ("raw", None),
1209
+ )
1210
+
1211
+ result = runtime.invoke(prepared)
1212
+
1213
+ assert isinstance(result, InvocationResult)
1214
+ assert result.value == ("raw", None)
1215
+ assert result.owned == (("raw", None), (device, None))
1216
+ assert not result.is_tokens
1217
+ assert calls == [False]
1218
+ assert not prepare(runtime, page_table=prepared.page_table).page_table_changed
1219
+
1220
+
1221
+ def test_blocking_consume_normalizes_logits_and_releases_owned_values(monkeypatch):
1222
+ runtime = make_runtime()
1223
+ host_logits = torch.arange(16, dtype=torch.float32).reshape(1, 1, 2, 8)
1224
+ released = []
1225
+ monkeypatch.setattr(runtime.config.output_reader, "read", lambda value, *, blocking: (host_logits, "probs"))
1226
+ monkeypatch.setattr(runtime, "_release_or_retain_transient", lambda value: released.append(value) or [])
1227
+ result = InvocationResult(value="raw", owned="owned", is_tokens=False)
1228
+
1229
+ logits, log_probs = runtime.consume(result)
1230
+
1231
+ assert logits.shape == (2, 1, 8)
1232
+ assert log_probs == "probs"
1233
+ assert released == ["owned"]
1234
+
1235
+
1236
+ def test_sampled_token_logprobs_are_flattened_to_lane_row_order():
1237
+ runtime = make_runtime()
1238
+ host_tokens = torch.tensor([[[[7], [8]]]], dtype=torch.int32)
1239
+ host_log_probs = torch.tensor([[[[-0.25, -0.75]]]], dtype=torch.bfloat16)
1240
+
1241
+ tokens, log_probs = runtime._normalize_host_output(
1242
+ (host_tokens, host_log_probs),
1243
+ is_tokens=True,
1244
+ )
1245
+
1246
+ assert tokens.tolist() == [7, 8]
1247
+ assert tokens.dtype == torch.int64
1248
+ assert log_probs.tolist() == pytest.approx([-0.25, -0.75])
1249
+ assert log_probs.dtype == torch.float32
1250
+
1251
+
1252
+ def test_raw_blocking_and_async_leases_release_exact_records(monkeypatch):
1253
+ runtime = make_runtime()
1254
+ deallocated = []
1255
+ monkeypatch.setattr(
1256
+ decode_module,
1257
+ "best_effort_deallocate_owned_tensors",
1258
+ lambda values, completed: deallocated.append(values) or [],
1259
+ )
1260
+
1261
+ first = InvocationResult(value=object(), owned="first-owned", is_tokens=False)
1262
+ assert runtime.consume(first, read_from_device=False) is first.value
1263
+ monkeypatch.setattr(runtime.config.output_reader, "read", lambda value, *, blocking: "first-host")
1264
+ assert runtime.read_decode_output(first.value) == "first-host"
1265
+
1266
+ second = InvocationResult(value=object(), owned="second-owned", is_tokens=True)
1267
+ runtime.consume(second, read_from_device=False)
1268
+ host_tokens = torch.tensor([[[[7], [8]]]], dtype=torch.int32)
1269
+ pending = PendingRead(value=(host_tokens, None), events=("event",), sequence=4, _owner=object())
1270
+ monkeypatch.setattr(runtime.config.output_reader, "submit", lambda value: pending)
1271
+ monkeypatch.setattr(runtime.config.output_reader, "complete", lambda value: pending.value)
1272
+
1273
+ host, events = runtime.read_decode_output(second.value, async_read=True)
1274
+ assert host is pending.value
1275
+ assert events == ["event"]
1276
+ tokens, log_probs = runtime.process_decode_output_host(host, is_tokens=True)
1277
+ assert tokens.tolist() == [7, 8]
1278
+ assert tokens.dtype == torch.int64
1279
+ assert log_probs is None
1280
+ assert deallocated == [
1281
+ (first.value, "first-owned"),
1282
+ (second.value, "second-owned"),
1283
+ ]
1284
+
1285
+
1286
+ def test_async_trace_lease_never_releases_borrowed_trace_output(monkeypatch):
1287
+ runtime = make_runtime()
1288
+ deallocated = []
1289
+ raw = object()
1290
+ host_tokens = torch.tensor([[[[7], [8]]]], dtype=torch.int32)
1291
+ pending = PendingRead(value=(host_tokens, None), events=("event",), sequence=4, _owner=object())
1292
+ monkeypatch.setattr(
1293
+ decode_module,
1294
+ "best_effort_deallocate_owned_tensors",
1295
+ lambda values, completed: deallocated.append(values) or [],
1296
+ )
1297
+ monkeypatch.setattr(runtime.config.output_reader, "submit", lambda value: pending)
1298
+ monkeypatch.setattr(runtime.config.output_reader, "complete", lambda value: pending.value)
1299
+
1300
+ result = InvocationResult(value=raw, owned=None, is_tokens=True)
1301
+ assert runtime.consume(result, read_from_device=False) is raw
1302
+ host, events = runtime.read_decode_output(raw, async_read=True)
1303
+ assert events == ["event"]
1304
+ tokens, log_probs = runtime.process_decode_output_host(host, is_tokens=True)
1305
+
1306
+ assert tokens.tolist() == [7, 8]
1307
+ assert log_probs is None
1308
+ assert deallocated == []
1309
+
1310
+
1311
+ def test_failed_transient_release_blocks_use_and_cleanup_retries(monkeypatch, expect_error):
1312
+ runtime = make_runtime()
1313
+
1314
+ class FakeTensor:
1315
+ pass
1316
+
1317
+ tensor = FakeTensor()
1318
+ attempts = []
1319
+ monkeypatch.setattr(decode_module.ttnn, "Tensor", FakeTensor)
1320
+
1321
+ def deallocate(value):
1322
+ attempts.append(value)
1323
+ if len(attempts) == 1:
1324
+ raise RuntimeError("release failed")
1325
+
1326
+ monkeypatch.setattr(decode_module.ttnn, "deallocate", deallocate)
1327
+
1328
+ failures = runtime._release_or_retain_transient(tensor)
1329
+ assert [str(error) for error in failures] == ["release failed"]
1330
+ assert runtime.transient_orphan_count == 1
1331
+ with expect_error(RuntimeError, "unreleased transient"):
1332
+ prepare(runtime)
1333
+
1334
+ runtime.cleanup_transients()
1335
+ assert attempts == [tensor, tensor]
1336
+ assert runtime.transient_orphan_count == 0
1337
+ assert prepare(runtime).sampling_path == "logits"
code/models/common/tests/llm_runtime/test_execution.py ADDED
@@ -0,0 +1,831 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import inspect
5
+ from dataclasses import dataclass
6
+ from types import SimpleNamespace
7
+
8
+ import torch
9
+
10
+ import models.common.llm_runtime.execution as execution_module
11
+ import ttnn
12
+ from models.common.llm_runtime.decode import DecodeRuntime
13
+ from models.common.llm_runtime.decode import InvocationResult as DecodeInvocationResult
14
+ from models.common.llm_runtime.execution import EagerExecutor, TracedExecutor
15
+ from models.common.llm_runtime.prefill.result_collector import InvocationResult as PrefillInvocationResult
16
+ from models.common.llm_runtime.prefill.runtime import PrefillRuntime
17
+ from models.common.llm_runtime.program_compiler import ProgramCompiler
18
+ from models.common.llm_runtime.trace_compiler import TraceCompiler
19
+
20
+
21
+ @dataclass(frozen=True)
22
+ class _Signature:
23
+ operation: str
24
+ variant: int
25
+
26
+ @property
27
+ def key_material(self):
28
+ return (("operation", self.operation), ("variant", self.variant))
29
+
30
+
31
+ def _runtime(runtime_type, **methods):
32
+ runtime = object.__new__(runtime_type)
33
+ for name, method in methods.items():
34
+ setattr(runtime, name, method)
35
+ return runtime
36
+
37
+
38
+ def _compiler(monkeypatch):
39
+ monkeypatch.setattr(ttnn, "synchronize_device", lambda mesh: None)
40
+ return ProgramCompiler("mesh", lambda: object())
41
+
42
+
43
+ def _trace_compiler(program_compiler, *, mode="all"):
44
+ return TraceCompiler(program_compiler)
45
+
46
+
47
+ def _prepared_prefill(*, trace_eligible=True, signatures=None, name="regular"):
48
+ if signatures is None:
49
+ signatures = (_Signature("prefill", 1),)
50
+ return SimpleNamespace(
51
+ name=name,
52
+ program_signatures=signatures,
53
+ trace_eligible=trace_eligible,
54
+ trace_signature=_Signature("prefill-trace", 1) if trace_eligible else None,
55
+ )
56
+
57
+
58
+ def _prepared_decode(*, variant=1):
59
+ return SimpleNamespace(
60
+ variant=variant,
61
+ device_feedback=True,
62
+ reset_batch=False,
63
+ page_table_changed=False,
64
+ sampling_params=None,
65
+ )
66
+
67
+
68
+ def test_execution_strategies_use_exact_identity_composition_without_type_frameworks(monkeypatch):
69
+ prefill = _runtime(PrefillRuntime)
70
+ decode = _runtime(DecodeRuntime)
71
+ program_compiler = _compiler(monkeypatch)
72
+ eager = EagerExecutor(prefill=prefill, decode=decode, program_compiler=program_compiler)
73
+ trace_compiler = _trace_compiler(program_compiler)
74
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
75
+
76
+ assert eager.prefill is prefill
77
+ assert eager.decode is decode
78
+ assert eager.program_compiler is program_compiler
79
+ assert traced.eager_executor is eager
80
+ assert traced.trace_compiler is trace_compiler
81
+ assert EagerExecutor not in TracedExecutor.__mro__
82
+ assert EagerExecutor.__bases__ == (object,)
83
+ assert TracedExecutor.__bases__ == (object,)
84
+
85
+ source = inspect.getsource(execution_module)
86
+ assert "Protocol" not in source
87
+ assert "ABC" not in source
88
+ assert "LightweightModule" not in source
89
+ assert not hasattr(execution_module, "EagerExecutorConfig")
90
+ assert not hasattr(execution_module, "TracedExecutorConfig")
91
+ assert not hasattr(EagerExecutor, "cleanup")
92
+ assert not hasattr(TracedExecutor, "cleanup")
93
+
94
+
95
+ def test_execution_request_signatures_are_exact_and_aligned():
96
+ required = inspect.Parameter.empty
97
+ positional = inspect.Parameter.POSITIONAL_OR_KEYWORD
98
+ keyword_only = inspect.Parameter.KEYWORD_ONLY
99
+ prefill_contract = [
100
+ ("self", positional, required),
101
+ ("tokens", keyword_only, required),
102
+ ("page_table", keyword_only, required),
103
+ ("prompt_lens", keyword_only, None),
104
+ ("start_pos", keyword_only, None),
105
+ ("empty_slots", keyword_only, None),
106
+ ("sampling_params", keyword_only, None),
107
+ ("prompt_tokens", keyword_only, None),
108
+ ("output_tokens", keyword_only, None),
109
+ ("slot_remap", keyword_only, None),
110
+ ]
111
+ decode_contract = [
112
+ ("self", positional, required),
113
+ ("tokens", keyword_only, required),
114
+ ("start_pos", keyword_only, required),
115
+ ("page_table", keyword_only, required),
116
+ ("sampling_params", keyword_only, None),
117
+ ("prompt_tokens", keyword_only, None),
118
+ ("output_tokens", keyword_only, None),
119
+ ("slot_remap", keyword_only, None),
120
+ ("reset_batch", keyword_only, False),
121
+ ]
122
+ decode_forward_contract = [
123
+ *decode_contract,
124
+ ("read_from_device", keyword_only, True),
125
+ ]
126
+
127
+ def parameter_contract(method):
128
+ return [
129
+ (parameter.name, parameter.kind, parameter.default)
130
+ for parameter in inspect.signature(method).parameters.values()
131
+ ]
132
+
133
+ for executor_type in (EagerExecutor, TracedExecutor):
134
+ assert parameter_contract(executor_type.compile_prefill) == prefill_contract
135
+ assert parameter_contract(executor_type.prefill_forward) == prefill_contract
136
+ assert parameter_contract(executor_type.compile_decode) == decode_contract
137
+ assert parameter_contract(executor_type.decode_forward) == decode_forward_contract
138
+
139
+ assert parameter_contract(EagerExecutor._prepare_prefill) == prefill_contract
140
+ assert parameter_contract(EagerExecutor._prepare_decode) == decode_contract
141
+ for method_name in ("compile_prefill", "prefill_forward", "compile_decode", "decode_forward"):
142
+ assert inspect.signature(getattr(EagerExecutor, method_name)) == inspect.signature(
143
+ getattr(TracedExecutor, method_name)
144
+ )
145
+
146
+
147
+ def test_execution_request_methods_reject_kv_cache(expect_error):
148
+ prefill_fields = {
149
+ "tokens": torch.zeros(1, 1),
150
+ "page_table": torch.zeros(1, 1),
151
+ }
152
+ decode_fields = {
153
+ "tokens": torch.zeros(1, 1),
154
+ "start_pos": torch.zeros(1),
155
+ "page_table": torch.zeros(1, 1),
156
+ }
157
+
158
+ for executor_type in (EagerExecutor, TracedExecutor):
159
+ executor = object.__new__(executor_type)
160
+ for method_name, fields in (
161
+ ("compile_prefill", prefill_fields),
162
+ ("prefill_forward", prefill_fields),
163
+ ("compile_decode", decode_fields),
164
+ ("decode_forward", decode_fields),
165
+ ):
166
+ with expect_error(TypeError, "kv_cache"):
167
+ getattr(executor, method_name)(**fields, kv_cache=object())
168
+
169
+
170
+ def test_traced_constructor_rejects_a_different_program_compiler(monkeypatch, expect_error):
171
+ eager = EagerExecutor(
172
+ prefill=_runtime(PrefillRuntime),
173
+ decode=_runtime(DecodeRuntime),
174
+ program_compiler=_compiler(monkeypatch),
175
+ )
176
+ unrelated_trace_compiler = _trace_compiler(_compiler(monkeypatch))
177
+
178
+ with expect_error(ValueError, "compose eager.program_compiler"):
179
+ TracedExecutor(eager=eager, trace_compiler=unrelated_trace_compiler)
180
+
181
+
182
+ def test_eager_prefill_prepares_once_and_compiles_all_signatures_from_same_object(monkeypatch):
183
+ prepared = _prepared_prefill(
184
+ signatures=(_Signature("prefill", 1), _Signature("prefill", 2)),
185
+ )
186
+ prepared_seen = []
187
+ prepare_calls = []
188
+
189
+ def prepare(
190
+ *,
191
+ tokens,
192
+ page_table,
193
+ prompt_lens=None,
194
+ start_pos=None,
195
+ empty_slots=None,
196
+ sampling_params=None,
197
+ ):
198
+ prepare_calls.append(
199
+ {
200
+ "tokens": tokens,
201
+ "page_table": page_table,
202
+ "prompt_lens": prompt_lens,
203
+ "start_pos": start_pos,
204
+ "empty_slots": empty_slots,
205
+ "sampling_params": sampling_params,
206
+ }
207
+ )
208
+ return (prepared,)
209
+
210
+ prefill = _runtime(
211
+ PrefillRuntime,
212
+ prepare=prepare,
213
+ invoke=lambda prepared, *, count_tokens=True: prepared_seen.append((prepared, count_tokens))
214
+ or PrefillInvocationResult(torch.zeros(1), ()),
215
+ )
216
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=_compiler(monkeypatch))
217
+ tokens = torch.zeros(1, 1)
218
+ page_table = torch.zeros(1, 1)
219
+ prompt_lens = torch.tensor([1])
220
+ start_pos = torch.tensor([0])
221
+ empty_slots = [0]
222
+ sampling_params = object()
223
+
224
+ eager.compile_prefill(
225
+ tokens=tokens,
226
+ page_table=page_table,
227
+ prompt_lens=prompt_lens,
228
+ start_pos=start_pos,
229
+ empty_slots=empty_slots,
230
+ sampling_params=sampling_params,
231
+ )
232
+
233
+ assert prepare_calls == [
234
+ {
235
+ "tokens": tokens,
236
+ "page_table": page_table,
237
+ "prompt_lens": prompt_lens,
238
+ "start_pos": start_pos,
239
+ "empty_slots": empty_slots,
240
+ "sampling_params": sampling_params,
241
+ }
242
+ ]
243
+ assert prepared_seen == [(prepared, False), (prepared, False)]
244
+
245
+
246
+ def test_traced_prefill_compile_does_not_interpret_request_eligibility(monkeypatch):
247
+ prepared = _prepared_prefill(trace_eligible=False)
248
+ identity_events = []
249
+ operation_plan = SimpleNamespace(
250
+ signature=_Signature("prefill-trace", 1),
251
+ prepare_inputs=lambda: (),
252
+ capture=lambda persistent: torch.zeros(1),
253
+ refresh_fields=("tokens",),
254
+ prime=None,
255
+ release_prime_output=None,
256
+ )
257
+
258
+ def prepare(
259
+ *,
260
+ tokens,
261
+ page_table,
262
+ prompt_lens=None,
263
+ start_pos=None,
264
+ empty_slots=None,
265
+ sampling_params=None,
266
+ ):
267
+ identity_events.append(("prepare", prepared))
268
+ return (prepared,)
269
+
270
+ prefill = _runtime(
271
+ PrefillRuntime,
272
+ prepare=prepare,
273
+ invoke=lambda prepared, *, count_tokens=True: identity_events.append(("invoke", prepared, count_tokens))
274
+ or PrefillInvocationResult(torch.zeros(1), ()),
275
+ capture_plan=lambda prepared: identity_events.append(("capture_plan", prepared)) or operation_plan,
276
+ )
277
+ program_compiler = _compiler(monkeypatch)
278
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
279
+ trace_compiler = _trace_compiler(program_compiler)
280
+ registered = []
281
+ trace_compiler.register_capture_plan = registered.append
282
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
283
+
284
+ traced.compile_prefill(tokens=torch.zeros(1, 1), page_table=torch.zeros(1, 1))
285
+
286
+ assert [event[0] for event in identity_events] == ["prepare", "invoke", "capture_plan"]
287
+ assert all(event[1] is prepared for event in identity_events)
288
+ assert identity_events[1][2] is False
289
+ assert len(registered) == 1
290
+
291
+
292
+ def test_traced_prefill_recompile_reuses_existing_trace_association(monkeypatch):
293
+ prepared = _prepared_prefill()
294
+
295
+ def prepare(
296
+ *,
297
+ tokens,
298
+ page_table,
299
+ prompt_lens=None,
300
+ start_pos=None,
301
+ empty_slots=None,
302
+ sampling_params=None,
303
+ ):
304
+ return (prepared,)
305
+
306
+ prefill = _runtime(
307
+ PrefillRuntime,
308
+ prepare=prepare,
309
+ invoke=lambda prepared, *, count_tokens=True: PrefillInvocationResult(torch.zeros(1), ()),
310
+ capture_plan=lambda prepared: (_ for _ in ()).throw(AssertionError("capture plan rebuilt")),
311
+ )
312
+ program_compiler = _compiler(monkeypatch)
313
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
314
+ trace_compiler = _trace_compiler(program_compiler)
315
+ trace_compiler.trace_key_for_program = lambda program_key: "existing-trace"
316
+ trace_compiler.register_capture_plan = lambda plan: (_ for _ in ()).throw(AssertionError("plan registered"))
317
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
318
+
319
+ traced.compile_prefill(tokens=torch.zeros(1, 1), page_table=torch.zeros(1, 1))
320
+
321
+
322
+ def test_traced_decode_recompile_reuses_existing_trace_association(monkeypatch):
323
+ prepared = _prepared_decode()
324
+
325
+ def prepare(*, tokens, start_pos, page_table, sampling_params=None, reset_batch=False):
326
+ return prepared
327
+
328
+ decode = _runtime(
329
+ DecodeRuntime,
330
+ prepare=prepare,
331
+ program_signature=lambda prepared: _Signature("decode", 1),
332
+ invoke=lambda prepared, *, device_feedback=False, count_tokens=True: DecodeInvocationResult(
333
+ torch.zeros(1), (), False
334
+ ),
335
+ capture_plan=lambda prepared: (_ for _ in ()).throw(AssertionError("capture plan rebuilt")),
336
+ )
337
+ program_compiler = _compiler(monkeypatch)
338
+ eager = EagerExecutor(prefill=_runtime(PrefillRuntime), decode=decode, program_compiler=program_compiler)
339
+ trace_compiler = _trace_compiler(program_compiler)
340
+ trace_compiler.trace_key_for_program = lambda program_key: "existing-trace"
341
+ trace_compiler.register_capture_plan = lambda plan: (_ for _ in ()).throw(AssertionError("plan registered"))
342
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
343
+
344
+ traced.compile_decode(tokens=torch.zeros(1), start_pos=torch.zeros(1), page_table=torch.zeros(1, 1))
345
+
346
+
347
+ def test_execution_target_selection_is_external_to_traced_prefill(monkeypatch):
348
+ prepared = _prepared_prefill(trace_eligible=True)
349
+ invocations = []
350
+
351
+ def prepare(
352
+ *,
353
+ tokens,
354
+ page_table,
355
+ prompt_lens=None,
356
+ start_pos=None,
357
+ empty_slots=None,
358
+ sampling_params=None,
359
+ ):
360
+ return (prepared,)
361
+
362
+ prefill = _runtime(
363
+ PrefillRuntime,
364
+ prepare=prepare,
365
+ invoke=lambda prepared: invocations.append(prepared) or PrefillInvocationResult("eager", ()),
366
+ assemble=lambda prepared_results, *, batch_size, sampling_params=None: prepared_results[0][1].value,
367
+ )
368
+ program_compiler = _compiler(monkeypatch)
369
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
370
+ trace_compiler = _trace_compiler(program_compiler, mode="decode_only")
371
+
372
+ def replay(
373
+ program_key,
374
+ refresh_inputs,
375
+ *,
376
+ reset_batch=False,
377
+ device_feedback_enabled=False,
378
+ feedback_compatible=False,
379
+ page_table_changed=False,
380
+ ):
381
+ raise AssertionError("trace replayed")
382
+
383
+ trace_compiler.replay = replay
384
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
385
+
386
+ assert (
387
+ traced.eager_executor.prefill_forward(
388
+ tokens=torch.zeros(1, 1),
389
+ page_table=torch.zeros(1, 1),
390
+ )
391
+ == "eager"
392
+ )
393
+ assert invocations == [prepared]
394
+
395
+
396
+ def test_prefill_replay_does_not_interpret_request_eligibility(monkeypatch):
397
+ prepared = _prepared_prefill(trace_eligible=False)
398
+ prepared.trace_signature = _Signature("prefill-trace", 1)
399
+ persistent = object()
400
+ hidden = object()
401
+ identity_events = []
402
+
403
+ def prepare(
404
+ *,
405
+ tokens,
406
+ page_table,
407
+ prompt_lens=None,
408
+ start_pos=None,
409
+ empty_slots=None,
410
+ sampling_params=None,
411
+ ):
412
+ identity_events.append(("prepare", prepared))
413
+ return (prepared,)
414
+
415
+ prefill = _runtime(
416
+ PrefillRuntime,
417
+ prepare=prepare,
418
+ refresh_trace=lambda prepared, persistent: identity_events.append(("refresh", prepared, persistent)),
419
+ finish_trace=lambda prepared, hidden, persistent: identity_events.append(
420
+ ("finish", prepared, hidden, persistent)
421
+ )
422
+ or "traced",
423
+ assemble=lambda prepared_results, *, batch_size, sampling_params=None: next(iter(prepared_results))[1],
424
+ )
425
+ program_compiler = _compiler(monkeypatch)
426
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
427
+ trace_compiler = _trace_compiler(program_compiler)
428
+ artifact = SimpleNamespace(persistent_inputs=SimpleNamespace(values=persistent))
429
+ record = SimpleNamespace(artifact=artifact)
430
+ trace_compiler.replay = (
431
+ lambda program_key, refresh_inputs, *, reset_batch=False, device_feedback_enabled=False, feedback_compatible=False, page_table_changed=False: refresh_inputs(
432
+ artifact, object()
433
+ )
434
+ or hidden
435
+ )
436
+ trace_compiler.trace_key_for_program = lambda program_key: "trace-key"
437
+ trace_compiler.get = lambda trace_key: record
438
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
439
+
440
+ result = traced.prefill_forward(tokens=torch.zeros(1, 1), page_table=torch.zeros(1, 1))
441
+
442
+ assert result == "traced"
443
+ assert [event[0] for event in identity_events] == ["prepare", "refresh", "finish"]
444
+ assert all(event[1] is prepared for event in identity_events)
445
+ assert identity_events[1][2] is persistent
446
+ assert identity_events[2][2:] == (hidden, persistent)
447
+
448
+
449
+ def test_prefill_replay_is_consumed_before_shared_trace_output_is_overwritten(monkeypatch):
450
+ prepared = (
451
+ _prepared_prefill(name="first"),
452
+ _prepared_prefill(name="second"),
453
+ )
454
+ persistent = {"output": None}
455
+
456
+ def prepare(**kwargs):
457
+ return prepared
458
+
459
+ def refresh_trace(request, trace_inputs):
460
+ trace_inputs["output"] = request.name
461
+
462
+ def assemble(prepared_results, *, batch_size, sampling_params=None):
463
+ return [result.value["output"] for _, result in prepared_results]
464
+
465
+ prefill = _runtime(
466
+ PrefillRuntime,
467
+ prepare=prepare,
468
+ refresh_trace=refresh_trace,
469
+ finish_trace=lambda request, hidden, trace_inputs: PrefillInvocationResult(trace_inputs, ()),
470
+ assemble=assemble,
471
+ )
472
+ program_compiler = _compiler(monkeypatch)
473
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
474
+ trace_compiler = _trace_compiler(program_compiler)
475
+ artifact = SimpleNamespace(persistent_inputs=SimpleNamespace(values=persistent))
476
+ trace_compiler.replay = lambda program_key, refresh_inputs, **kwargs: refresh_inputs(artifact, object()) or "hidden"
477
+ trace_compiler.trace_key_for_program = lambda program_key: "trace-key"
478
+ trace_compiler.get = lambda trace_key: SimpleNamespace(artifact=artifact)
479
+
480
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
481
+
482
+ assert traced.prefill_forward(tokens=torch.zeros(2, 1), page_table=torch.zeros(2, 1)) == [
483
+ "first",
484
+ "second",
485
+ ]
486
+
487
+
488
+ def test_prefill_missing_trace_artifact_is_an_error_without_eager_reinvocation(monkeypatch, expect_error):
489
+ prepared = _prepared_prefill(trace_eligible=True)
490
+ eager_invocations = []
491
+
492
+ def prepare(
493
+ *,
494
+ tokens,
495
+ page_table,
496
+ prompt_lens=None,
497
+ start_pos=None,
498
+ empty_slots=None,
499
+ sampling_params=None,
500
+ ):
501
+ return (prepared,)
502
+
503
+ prefill = _runtime(
504
+ PrefillRuntime,
505
+ prepare=prepare,
506
+ invoke=lambda prepared: eager_invocations.append(prepared) or PrefillInvocationResult("eager", ()),
507
+ refresh_trace=lambda prepared, persistent: None,
508
+ assemble=lambda prepared_results, *, batch_size, sampling_params=None: next(iter(prepared_results))[1],
509
+ )
510
+ program_compiler = _compiler(monkeypatch)
511
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
512
+ trace_compiler = _trace_compiler(program_compiler)
513
+ artifact = SimpleNamespace(persistent_inputs=SimpleNamespace(values=()))
514
+ trace_compiler.replay = (
515
+ lambda program_key, refresh_inputs, *, reset_batch=False, device_feedback_enabled=False, feedback_compatible=False, page_table_changed=False: refresh_inputs(
516
+ artifact, object()
517
+ )
518
+ or "hidden"
519
+ )
520
+ trace_compiler.trace_key_for_program = lambda program_key: "missing"
521
+ trace_compiler.get = lambda trace_key: None
522
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
523
+
524
+ with expect_error(RuntimeError, "Required prefill trace") as exc_info:
525
+ traced.prefill_forward(tokens=torch.zeros(1, 1), page_table=torch.zeros(1, 1))
526
+
527
+ assert eager_invocations == []
528
+ message = str(exc_info.value)
529
+ for field in (
530
+ "operation=prefill",
531
+ "trace_mode=all",
532
+ "model=PrefillRuntime",
533
+ "signature_material=",
534
+ "signature_digest=",
535
+ "configured_coverage=",
536
+ "TraceConfig(mode='none')",
537
+ ):
538
+ assert field in message
539
+ assert traced.coverage_miss_count == 1
540
+
541
+
542
+ def test_prefill_missing_trace_signature_is_rejected_before_replay(monkeypatch, expect_error):
543
+ prepared = _prepared_prefill(trace_eligible=False)
544
+ program_compiler = _compiler(monkeypatch)
545
+ eager = EagerExecutor(
546
+ prefill=_runtime(PrefillRuntime),
547
+ decode=_runtime(DecodeRuntime),
548
+ program_compiler=program_compiler,
549
+ )
550
+ trace_compiler = _trace_compiler(program_compiler)
551
+ trace_compiler.replay = lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("unexpected replay"))
552
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
553
+
554
+ with expect_error(RuntimeError, "not trace-eligible"):
555
+ traced._execute_prefill(prepared)
556
+
557
+
558
+ def test_decode_missing_trace_artifact_reports_exact_strict_coverage(monkeypatch, expect_error):
559
+ prepared = _prepared_decode()
560
+ signature = _Signature("decode", 7)
561
+ decode = _runtime(
562
+ DecodeRuntime,
563
+ config=SimpleNamespace(position_feedback_capable=True),
564
+ program_signature=lambda value: signature,
565
+ )
566
+ program_compiler = _compiler(monkeypatch)
567
+ eager = EagerExecutor(prefill=_runtime(PrefillRuntime), decode=decode, program_compiler=program_compiler)
568
+ trace_compiler = _trace_compiler(program_compiler)
569
+ trace_compiler.trace_key_for_program = lambda key: None
570
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler, trace_mode="decode_only")
571
+
572
+ with expect_error(RuntimeError, "Required decode trace") as exc_info:
573
+ traced._execute_decode(prepared)
574
+
575
+ message = str(exc_info.value)
576
+ for field in (
577
+ "operation=decode",
578
+ "trace_mode=decode_only",
579
+ "model=DecodeRuntime",
580
+ "signature_material=",
581
+ "signature_digest=",
582
+ "program_key=",
583
+ "trace_key=unavailable",
584
+ "configured_coverage=",
585
+ "TraceConfig(mode='none')",
586
+ ):
587
+ assert field in message
588
+ assert traced.coverage_miss_count == 1
589
+
590
+
591
+ def test_prefill_whole_call_preflights_every_trace_before_first_replay(monkeypatch, expect_error):
592
+ first = _prepared_prefill(signatures=(_Signature("prefill", 1),), name="first")
593
+ second = _prepared_prefill(signatures=(_Signature("prefill", 2),), name="second")
594
+ prefill = _runtime(
595
+ PrefillRuntime,
596
+ prepare=lambda **kwargs: (first, second),
597
+ assemble=lambda results, **kwargs: tuple(results),
598
+ )
599
+ program_compiler = _compiler(monkeypatch)
600
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
601
+ trace_compiler = _trace_compiler(program_compiler)
602
+ first_key = program_compiler.key_for(first.program_signatures[0])
603
+ second_key = program_compiler.key_for(second.program_signatures[0])
604
+ artifact = SimpleNamespace(persistent_inputs=SimpleNamespace(values=()))
605
+ trace_compiler.trace_key_for_program = lambda key: "first" if key == first_key else "second"
606
+ trace_compiler.get = lambda key: SimpleNamespace(artifact=artifact) if key == "first" else None
607
+ replays = []
608
+ trace_compiler.replay = lambda *args, **kwargs: replays.append(args) or "hidden"
609
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
610
+
611
+ with expect_error(RuntimeError, second_key.digest):
612
+ traced.prefill_forward(tokens=torch.zeros(2, 1), page_table=torch.zeros(2, 1))
613
+
614
+ assert replays == []
615
+
616
+
617
+ def test_prefill_fixed_chunk_replays_every_step_and_finishes_only_final_hidden(monkeypatch):
618
+ prepared = _prepared_prefill()
619
+ chunks = ("chunk-0", "chunk-1", "chunk-2")
620
+ prepared.request = SimpleNamespace(chunks=chunks)
621
+ events = []
622
+ prefill = _runtime(
623
+ PrefillRuntime,
624
+ refresh_trace=lambda request, hidden_inputs, workspace, chunk: events.append(("refresh", chunk, workspace)),
625
+ finish_trace=lambda request, hidden, persistent: events.append(("finish", hidden, persistent))
626
+ or PrefillInvocationResult(hidden, ()),
627
+ )
628
+ program_compiler = _compiler(monkeypatch)
629
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
630
+ trace_compiler = _trace_compiler(program_compiler)
631
+ persistent = object()
632
+ artifact = SimpleNamespace(persistent_inputs=SimpleNamespace(values=persistent))
633
+ record = SimpleNamespace(artifact=artifact)
634
+ trace_compiler.trace_key_for_program = lambda key: "trace"
635
+ trace_compiler.get = lambda key: record
636
+ trace_compiler.workspace_for_program = lambda key: persistent
637
+ replayed = []
638
+
639
+ def replay(program_key, refresh_inputs, **kwargs):
640
+ refresh_inputs(artifact, object())
641
+ replayed.append(program_key)
642
+ return f"hidden-{len(replayed) - 1}"
643
+
644
+ trace_compiler.replay = replay
645
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
646
+
647
+ result = traced._execute_prefill(prepared)
648
+
649
+ assert result.value == "hidden-2"
650
+ assert [event[:2] for event in events] == [
651
+ ("refresh", "chunk-0"),
652
+ ("refresh", "chunk-1"),
653
+ ("refresh", "chunk-2"),
654
+ ("finish", "hidden-2"),
655
+ ]
656
+ assert all(event[-1] is persistent for event in events)
657
+
658
+
659
+ def test_prefill_replay_emits_structured_serving_evidence(monkeypatch):
660
+ class Request:
661
+ source_rows = tuple(range(15))
662
+ padded_batch_size = 16
663
+ padded_sequence_length = 128
664
+ chunks = ("step",)
665
+
666
+ monkeypatch.setattr(execution_module, "PrefillRequest", Request)
667
+ signature = SimpleNamespace(
668
+ operation_variant="regular-batched",
669
+ key_material=(("operation_variant", "regular-batched"),),
670
+ )
671
+ prepared = SimpleNamespace(
672
+ request=Request(),
673
+ program_signatures=(signature,),
674
+ trace_signature=SimpleNamespace(key_material=(("padded_batch_size", 16),)),
675
+ sampling_path="topk",
676
+ sampling_params=object(),
677
+ )
678
+ persistent = object()
679
+ prefill = _runtime(
680
+ PrefillRuntime,
681
+ refresh_trace=lambda request, hidden_inputs, workspace, chunk: None,
682
+ finish_trace=lambda request, hidden, workspace: PrefillInvocationResult(hidden, ()),
683
+ assemble=lambda results, **kwargs: next(iter(results))[1],
684
+ )
685
+ program_compiler = _compiler(monkeypatch)
686
+ eager = EagerExecutor(prefill=prefill, decode=_runtime(DecodeRuntime), program_compiler=program_compiler)
687
+ trace_compiler = _trace_compiler(program_compiler)
688
+ program_key = program_compiler.key_for(signature)
689
+ trace_key = SimpleNamespace(digest="1" * 64)
690
+ artifact = SimpleNamespace(persistent_inputs=SimpleNamespace(values=persistent))
691
+ record = SimpleNamespace(artifact=artifact)
692
+ trace_compiler.trace_key_for_program = lambda key: trace_key
693
+ trace_compiler.workspace_for_program = lambda key: persistent
694
+ trace_compiler.replay = lambda key, refresh, **kwargs: refresh(artifact, object()) or "hidden"
695
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
696
+
697
+ result = traced.execute_prepared_prefill(
698
+ ((prepared, ((program_key, record),)),),
699
+ batch_size=15,
700
+ sampling_params=prepared.sampling_params,
701
+ lane=3,
702
+ )
703
+
704
+ assert result.value == "hidden"
705
+ assert len(traced.recent_prefill_replay_evidence) == 1
706
+ evidence = traced.recent_prefill_replay_evidence[0]
707
+ assert evidence.operation == "prefill"
708
+ assert evidence.variant == "regular-batched"
709
+ assert evidence.sampling_path == "topk"
710
+ assert evidence.execution == "trace_replay"
711
+ assert (evidence.active_batch_size, evidence.padded_batch_size) == (15, 16)
712
+ assert evidence.padded_sequence_length == 128
713
+ assert (evidence.lane, evidence.rank) == (3, 3)
714
+ assert evidence.program_key == program_key.digest
715
+ assert evidence.trace_key == "1" * 64
716
+ assert evidence.replay_steps == 1
717
+ assert traced.runtime_summary() == {
718
+ "eager_prefill_executions": 0,
719
+ "semantic_program_count": 0,
720
+ "rejected_post_activation_compile_attempts": 0,
721
+ "ttnn_program_cache_count": None,
722
+ "successful_trace_replays": 0,
723
+ "trace_replays_by_operation": {"prefill": 0, "decode": 0},
724
+ "strict_coverage_misses": 0,
725
+ "semantic_trace_count": 0,
726
+ "trace_association_count": 0,
727
+ }
728
+
729
+
730
+ def test_decode_replay_prepares_once_and_uses_same_object_for_refresh_submission_and_consume(monkeypatch):
731
+ prepared = _prepared_decode()
732
+ events = []
733
+
734
+ def prepare(*, tokens, start_pos, page_table, sampling_params=None, reset_batch=False):
735
+ events.append(("prepare", prepared, sampling_params, reset_batch))
736
+ return prepared
737
+
738
+ decode = _runtime(
739
+ DecodeRuntime,
740
+ config=SimpleNamespace(position_feedback_capable=True),
741
+ prepare=prepare,
742
+ program_signature=lambda prepared: events.append(("signature", prepared)) or _Signature("decode", 1),
743
+ refresh_trace=lambda artifact, prepared, decision: events.append(("refresh", prepared)),
744
+ note_submitted=lambda prepared: events.append(("submitted", prepared)),
745
+ consume=lambda result, *, read_from_device=True: events.append(("consume", result, read_from_device))
746
+ or result.value,
747
+ )
748
+ program_compiler = _compiler(monkeypatch)
749
+ eager = EagerExecutor(prefill=_runtime(PrefillRuntime), decode=decode, program_compiler=program_compiler)
750
+ trace_compiler = _trace_compiler(program_compiler)
751
+ trace_key = SimpleNamespace(digest="2" * 64)
752
+ trace_compiler.trace_key_for_program = lambda key: trace_key
753
+ trace_compiler.get = lambda key: SimpleNamespace(artifact=object())
754
+ trace_compiler.replay = (
755
+ lambda program_key, refresh_inputs, *, reset_batch=False, device_feedback_enabled=False, feedback_compatible=False, page_table_changed=False: refresh_inputs(
756
+ object(), object()
757
+ )
758
+ or "token"
759
+ )
760
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
761
+ sampling_params = object()
762
+
763
+ result = traced.decode_forward(
764
+ tokens=torch.zeros(1, 1),
765
+ start_pos=torch.zeros(1),
766
+ page_table=torch.zeros(1, 1),
767
+ sampling_params=sampling_params,
768
+ reset_batch=True,
769
+ read_from_device=False,
770
+ )
771
+
772
+ assert result == "token"
773
+ assert [event[0] for event in events] == ["prepare", "signature", "refresh", "submitted", "consume"]
774
+ assert all(event[1] is prepared for event in events[:-1])
775
+ assert events[0][2:] == (sampling_params, True)
776
+ assert isinstance(events[-1][1], DecodeInvocationResult)
777
+ assert events[-1][1].owned is None
778
+ assert events[-1][2] is False
779
+
780
+
781
+ def test_explicit_eager_decode_delegates_once_and_execution_objects_do_not_cleanup(monkeypatch):
782
+ prepared = _prepared_decode()
783
+ calls = []
784
+
785
+ def prepare(*, tokens, start_pos, page_table, sampling_params=None, reset_batch=False):
786
+ calls.append(("prepare", prepared, sampling_params, reset_batch))
787
+ return prepared
788
+
789
+ decode = _runtime(
790
+ DecodeRuntime,
791
+ prepare=prepare,
792
+ invoke=lambda prepared, *, device_feedback=False: calls.append(("invoke", prepared, device_feedback))
793
+ or DecodeInvocationResult("eager", (), False),
794
+ consume=lambda result, *, read_from_device=True: calls.append(("consume", result.value, read_from_device))
795
+ or result.value,
796
+ )
797
+ program_compiler = _compiler(monkeypatch)
798
+ eager = EagerExecutor(prefill=_runtime(PrefillRuntime), decode=decode, program_compiler=program_compiler)
799
+ trace_compiler = _trace_compiler(program_compiler)
800
+
801
+ def replay(
802
+ program_key,
803
+ refresh_inputs,
804
+ *,
805
+ reset_batch=False,
806
+ device_feedback_enabled=False,
807
+ feedback_compatible=False,
808
+ page_table_changed=False,
809
+ ):
810
+ raise AssertionError("trace replayed")
811
+
812
+ trace_compiler.replay = replay
813
+ traced = TracedExecutor(eager=eager, trace_compiler=trace_compiler)
814
+ sampling_params = object()
815
+
816
+ assert (
817
+ traced.eager_executor.decode_forward(
818
+ tokens=torch.zeros(1, 1),
819
+ start_pos=torch.zeros(1),
820
+ page_table=torch.zeros(1, 1),
821
+ sampling_params=sampling_params,
822
+ reset_batch=True,
823
+ read_from_device=False,
824
+ )
825
+ == "eager"
826
+ )
827
+ assert calls == [
828
+ ("prepare", prepared, sampling_params, True),
829
+ ("invoke", prepared, False),
830
+ ("consume", "eager", False),
831
+ ]
code/models/common/tests/llm_runtime/test_executor_integration.py ADDED
@@ -0,0 +1,1833 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Generic host-only executor/generator integration for migrated siblings."""
5
+
6
+ import inspect
7
+ from dataclasses import replace
8
+ from importlib import import_module
9
+ from types import SimpleNamespace
10
+ from unittest.mock import MagicMock, create_autospec
11
+
12
+ import pytest
13
+ import torch
14
+
15
+ import ttnn
16
+ from models.common.llm_runtime.config import PagedKVCacheConfig, PageTableLayout, TraceConfig, WarmupConfig
17
+ from models.common.llm_runtime.decode import DecodeTraceSignature
18
+ from models.common.llm_runtime.execution import EagerExecutor
19
+ from models.common.llm_runtime.lane_group import LaneGroupExecutor
20
+ from models.common.llm_runtime.prefill.signatures import PrefillProgramSignature
21
+ from models.common.llm_runtime.program_compiler import ProgramKey
22
+ from models.common.llm_runtime.warmup import _build_plan
23
+ from models.common.models import executor as shared_model_executor
24
+ from models.common.models import llama3_executor as llama3_family_executor
25
+ from models.common.models import qwen2_executor as qwen2_family_executor
26
+ from models.common.models.deepseek_r1_distill_qwen_14b import executor as deepseek_executor
27
+ from models.common.models.deepseek_r1_distill_qwen_14b import generator as deepseek_generator
28
+ from models.common.models.llama32_1b import executor as llama32_executor
29
+ from models.common.models.llama32_1b import generator as llama32_generator
30
+ from models.common.models.llama32_3b import executor as llama32_3b_executor
31
+ from models.common.models.llama32_3b import generator as llama32_3b_generator
32
+ from models.common.models.llama33_70b import executor as llama33_70b_executor
33
+ from models.common.models.llama33_70b import generator as llama33_70b_generator
34
+ from models.common.models.mistral_7b import executor as mistral_executor
35
+ from models.common.models.mistral_7b import generator as mistral_generator
36
+ from models.common.models.phi4 import executor as phi4_executor
37
+ from models.common.models.phi4 import generator as phi4_generator
38
+ from models.common.models.qwen2_7b import executor as qwen2_executor
39
+ from models.common.models.qwen2_7b import generator as qwen2_generator
40
+ from models.common.models.qwen3_32b import executor as qwen3_32b_executor
41
+ from models.common.models.qwen3_32b import generator as qwen3_32b_generator
42
+ from models.common.models.qwen25_7b import executor as qwen25_executor
43
+ from models.common.models.qwen25_7b import generator as qwen25_generator
44
+ from models.common.models.qwen25_72b import executor as qwen25_72b_executor
45
+ from models.common.models.qwen25_72b import generator as qwen25_72b_generator
46
+ from models.common.models.qwen25_coder_32b import executor as qwen25_coder_32b_executor
47
+ from models.common.models.qwen25_coder_32b import generator as qwen25_coder_32b_generator
48
+
49
+ EXECUTOR_BINDINGS = {
50
+ "llama32_1b": SimpleNamespace(
51
+ executor_module=llama32_executor,
52
+ executor_class=llama32_executor.Llama32_1BExecutor,
53
+ executor_config_class=llama32_executor.Llama32_1BExecutorConfig,
54
+ generator_module=llama32_generator,
55
+ generator_class=llama32_generator.Llama32_1BGenerator,
56
+ generator_config_class=llama32_generator.Llama32_1BGeneratorConfig,
57
+ build_generator_name="build_llama32_1b_generator",
58
+ build_executor_name="build_llama32_1b_executor",
59
+ make_model=lambda **kwargs: _make_llama32_model(**kwargs),
60
+ make_runtime_config=lambda: _make_llama32_runtime_config(),
61
+ make_executor_config=lambda mode="none": _make_llama32_executor_config(mode),
62
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_llama32_model(), **kwargs),
63
+ make_product=lambda mesh_device, max_batch_size: _make_llama32_product(mesh_device, max_batch_size),
64
+ make_lane=lambda llm, config: _FakeLane(llm, config),
65
+ hf_model="meta-llama/Llama-3.2-1B-Instruct",
66
+ ),
67
+ "llama32_3b": SimpleNamespace(
68
+ executor_module=llama32_3b_executor,
69
+ executor_class=llama32_3b_executor.Llama32_3BExecutor,
70
+ executor_config_class=llama32_3b_executor.Llama32_3BExecutorConfig,
71
+ generator_module=llama32_3b_generator,
72
+ generator_class=llama32_3b_generator.Llama32_3BGenerator,
73
+ generator_config_class=llama32_3b_generator.Llama32_3BGeneratorConfig,
74
+ build_generator_name="build_llama32_3b_generator",
75
+ build_executor_name="build_llama32_3b_executor",
76
+ make_model=lambda **kwargs: _make_llama32_model(**kwargs),
77
+ make_runtime_config=lambda: _make_llama32_runtime_config(),
78
+ make_executor_config=lambda mode="none": _make_llama32_executor_config(mode, module=llama32_3b_executor),
79
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_llama32_model(), **kwargs),
80
+ make_product=lambda mesh_device, max_batch_size: _make_llama32_product(mesh_device, max_batch_size),
81
+ make_lane=lambda llm, config: _FakeLane(llm, config),
82
+ hf_model="meta-llama/Llama-3.2-3B-Instruct",
83
+ ),
84
+ "llama33_70b": SimpleNamespace(
85
+ executor_module=llama33_70b_executor,
86
+ executor_class=llama33_70b_executor.Llama33_70BExecutor,
87
+ executor_config_class=llama33_70b_executor.Llama33_70BExecutorConfig,
88
+ generator_module=llama33_70b_generator,
89
+ generator_class=llama33_70b_generator.Llama33_70BGenerator,
90
+ generator_config_class=llama33_70b_generator.Llama33_70BGeneratorConfig,
91
+ build_generator_name="build_llama33_70b_generator",
92
+ build_executor_name="build_llama33_70b_executor",
93
+ make_model=lambda **kwargs: _make_llama32_model(**kwargs),
94
+ make_runtime_config=lambda: _make_llama32_runtime_config(),
95
+ make_executor_config=lambda mode="none": _make_llama32_executor_config(mode, module=llama33_70b_executor),
96
+ make_recording_target=lambda **kwargs: _RecordingTarget(
97
+ _make_llama32_model(),
98
+ request_state_fields=llama33_70b_executor.Llama33_70BExecutor.request_state_fields,
99
+ **kwargs,
100
+ ),
101
+ make_product=lambda mesh_device, max_batch_size: _make_llama32_product(mesh_device, max_batch_size),
102
+ make_lane=lambda llm, config: _FakeLane(
103
+ llm, config, request_state_fields=llama33_70b_executor.Llama33_70BExecutor.request_state_fields
104
+ ),
105
+ hf_model="meta-llama/Llama-3.3-70B-Instruct",
106
+ ),
107
+ "qwen2_7b": SimpleNamespace(
108
+ executor_module=qwen2_executor,
109
+ executor_class=qwen2_executor.Qwen2Executor,
110
+ executor_config_class=qwen2_executor.Qwen2ExecutorConfig,
111
+ generator_module=qwen2_generator,
112
+ generator_class=qwen2_generator.Qwen2Generator,
113
+ generator_config_class=qwen2_generator.Qwen2GeneratorConfig,
114
+ build_generator_name="build_qwen2_7b_generator",
115
+ build_executor_name="build_qwen2_7b_executor",
116
+ make_model=lambda **kwargs: _make_qwen2_model(**kwargs),
117
+ make_runtime_config=lambda: _make_qwen2_runtime_config(),
118
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode),
119
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen2_model(), **kwargs),
120
+ make_product=lambda mesh_device, max_batch_size: _make_qwen2_product(mesh_device, max_batch_size),
121
+ make_lane=lambda llm, config: _FakeLane(llm, config),
122
+ hf_model="Qwen/Qwen2-7B-Instruct",
123
+ ),
124
+ "qwen25_7b": SimpleNamespace(
125
+ executor_module=qwen25_executor,
126
+ executor_class=qwen25_executor.Qwen25Executor,
127
+ executor_config_class=qwen25_executor.Qwen25ExecutorConfig,
128
+ generator_module=qwen25_generator,
129
+ generator_class=qwen25_generator.Qwen25Generator,
130
+ generator_config_class=qwen25_generator.Qwen25GeneratorConfig,
131
+ build_generator_name="build_qwen25_7b_generator",
132
+ build_executor_name="build_qwen25_7b_executor",
133
+ make_model=lambda **kwargs: _make_qwen2_model(**kwargs),
134
+ make_runtime_config=lambda: _make_qwen2_runtime_config(max_prefill_batch_size=8),
135
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=qwen25_executor),
136
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen2_model(), **kwargs),
137
+ make_product=lambda mesh_device, max_batch_size: _make_qwen2_product(
138
+ mesh_device, max_batch_size, max_prefill_batch_size=8
139
+ ),
140
+ make_lane=lambda llm, config: _FakeLane(llm, config),
141
+ hf_model="Qwen/Qwen2.5-7B-Instruct",
142
+ ),
143
+ "qwen25_72b": SimpleNamespace(
144
+ executor_module=qwen25_72b_executor,
145
+ executor_class=qwen25_72b_executor.Qwen25_72BExecutor,
146
+ executor_config_class=qwen25_72b_executor.Qwen25_72BExecutorConfig,
147
+ generator_module=qwen25_72b_generator,
148
+ generator_class=qwen25_72b_generator.Qwen25_72BGenerator,
149
+ generator_config_class=qwen25_72b_generator.Qwen25_72BGeneratorConfig,
150
+ build_generator_name="build_qwen25_72b_generator",
151
+ build_executor_name="build_qwen25_72b_executor",
152
+ make_model=lambda **kwargs: _make_qwen25_72b_model(**kwargs),
153
+ make_runtime_config=lambda: _make_qwen25_72b_runtime_config(),
154
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=qwen25_72b_executor),
155
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen25_72b_model(), **kwargs),
156
+ make_product=lambda mesh_device, max_batch_size: _make_qwen25_72b_product(mesh_device, max_batch_size),
157
+ make_lane=lambda llm, config: _FakeLane(llm, config),
158
+ hf_model="Qwen/Qwen2.5-72B-Instruct",
159
+ ),
160
+ "qwen25_coder_32b": SimpleNamespace(
161
+ executor_module=qwen25_coder_32b_executor,
162
+ executor_class=qwen25_coder_32b_executor.Qwen25Coder32BExecutor,
163
+ executor_config_class=qwen25_coder_32b_executor.Qwen25Coder32BExecutorConfig,
164
+ generator_module=qwen25_coder_32b_generator,
165
+ generator_class=qwen25_coder_32b_generator.Qwen25Coder32BGenerator,
166
+ generator_config_class=qwen25_coder_32b_generator.Qwen25Coder32BGeneratorConfig,
167
+ build_generator_name="build_qwen25_coder_32b_generator",
168
+ build_executor_name="build_qwen25_coder_32b_executor",
169
+ make_model=lambda **kwargs: _make_qwen25_coder_32b_model(**kwargs),
170
+ make_runtime_config=lambda: _make_qwen25_coder_32b_runtime_config(),
171
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=qwen25_coder_32b_executor),
172
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen25_coder_32b_model(), **kwargs),
173
+ make_product=lambda mesh_device, max_batch_size: _make_qwen25_coder_32b_product(mesh_device, max_batch_size),
174
+ make_lane=lambda llm, config: _FakeLane(llm, config),
175
+ hf_model="Qwen/Qwen2.5-Coder-32B-Instruct",
176
+ ),
177
+ "qwen3_32b": SimpleNamespace(
178
+ executor_module=qwen3_32b_executor,
179
+ executor_class=qwen3_32b_executor.Qwen3_32BExecutor,
180
+ executor_config_class=qwen3_32b_executor.Qwen3_32BExecutorConfig,
181
+ generator_module=qwen3_32b_generator,
182
+ generator_class=qwen3_32b_generator.Qwen3_32BGenerator,
183
+ generator_config_class=qwen3_32b_generator.Qwen3_32BGeneratorConfig,
184
+ build_generator_name="build_qwen3_32b_generator",
185
+ build_executor_name="build_qwen3_32b_executor",
186
+ make_model=lambda **kwargs: _make_qwen3_32b_model(**kwargs),
187
+ make_runtime_config=lambda: _make_qwen3_32b_runtime_config(),
188
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=qwen3_32b_executor),
189
+ make_recording_target=lambda **kwargs: _RecordingTarget(
190
+ _make_qwen3_32b_model(),
191
+ request_state_fields=qwen3_32b_executor.Qwen3_32BExecutor.request_state_fields,
192
+ **kwargs,
193
+ ),
194
+ make_product=lambda mesh_device, max_batch_size: _make_qwen3_32b_product(mesh_device, max_batch_size),
195
+ make_lane=lambda llm, config: _FakeLane(
196
+ llm, config, request_state_fields=qwen3_32b_executor.Qwen3_32BExecutor.request_state_fields
197
+ ),
198
+ hf_model="Qwen/Qwen3-32B",
199
+ ),
200
+ "deepseek_r1_distill_qwen_14b": SimpleNamespace(
201
+ executor_module=deepseek_executor,
202
+ executor_class=deepseek_executor.DeepSeekR1Qwen14BExecutor,
203
+ executor_config_class=deepseek_executor.DeepSeekR1Qwen14BExecutorConfig,
204
+ generator_module=deepseek_generator,
205
+ generator_class=deepseek_generator.DeepSeekR1Qwen14BGenerator,
206
+ generator_config_class=deepseek_generator.DeepSeekR1Qwen14BGeneratorConfig,
207
+ build_generator_name="build_deepseek_r1_distill_qwen_14b_generator",
208
+ build_executor_name="build_deepseek_r1_distill_qwen_14b_executor",
209
+ make_model=lambda **kwargs: _make_qwen2_model(**kwargs),
210
+ make_runtime_config=lambda: _make_qwen2_runtime_config(max_prefill_batch_size=32),
211
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=deepseek_executor),
212
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen2_model(), **kwargs),
213
+ make_product=lambda mesh_device, max_batch_size: _make_qwen2_product(
214
+ mesh_device, max_batch_size, max_prefill_batch_size=32
215
+ ),
216
+ make_lane=lambda llm, config: _FakeLane(llm, config),
217
+ hf_model="deepseek-ai/DeepSeek-R1-Distill-Qwen-14B",
218
+ ),
219
+ "mistral_7b": SimpleNamespace(
220
+ executor_module=mistral_executor,
221
+ executor_class=mistral_executor.Mistral7BExecutor,
222
+ executor_config_class=mistral_executor.Mistral7BExecutorConfig,
223
+ generator_module=mistral_generator,
224
+ generator_class=mistral_generator.Mistral7BGenerator,
225
+ generator_config_class=mistral_generator.Mistral7BGeneratorConfig,
226
+ build_generator_name="build_mistral_7b_generator",
227
+ build_executor_name="build_mistral_7b_executor",
228
+ make_model=lambda **kwargs: _make_qwen2_model(**kwargs),
229
+ make_runtime_config=lambda: _make_qwen2_runtime_config(max_prefill_batch_size=8),
230
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=mistral_executor),
231
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen2_model(), **kwargs),
232
+ make_product=lambda mesh_device, max_batch_size: _make_qwen2_product(
233
+ mesh_device, max_batch_size, max_prefill_batch_size=8
234
+ ),
235
+ make_lane=lambda llm, config: _FakeLane(llm, config),
236
+ hf_model="mistralai/Mistral-7B-Instruct-v0.3",
237
+ ),
238
+ "phi4": SimpleNamespace(
239
+ executor_module=phi4_executor,
240
+ executor_class=phi4_executor.Phi4Executor,
241
+ executor_config_class=phi4_executor.Phi4ExecutorConfig,
242
+ generator_module=phi4_generator,
243
+ generator_class=phi4_generator.Phi4Generator,
244
+ generator_config_class=phi4_generator.Phi4GeneratorConfig,
245
+ build_generator_name="build_phi4_generator",
246
+ build_executor_name="build_phi4_executor",
247
+ make_model=lambda **kwargs: _make_qwen2_model(**kwargs),
248
+ make_runtime_config=lambda: _make_qwen2_runtime_config(max_prefill_batch_size=8),
249
+ make_executor_config=lambda mode="none": _make_qwen2_executor_config(mode, module=phi4_executor),
250
+ make_recording_target=lambda **kwargs: _RecordingTarget(_make_qwen2_model(), **kwargs),
251
+ make_product=lambda mesh_device, max_batch_size: _make_qwen2_product(
252
+ mesh_device, max_batch_size, max_prefill_batch_size=8
253
+ ),
254
+ make_lane=lambda llm, config: _FakeLane(llm, config),
255
+ hf_model="microsoft/phi-4",
256
+ ),
257
+ }
258
+
259
+ GENERATOR_PATHS = {
260
+ "llama32_1b": "models.common.models.llama32_1b.generator:Llama32_1BGenerator",
261
+ "llama32_3b": "models.common.models.llama32_3b.generator:Llama32_3BGenerator",
262
+ "llama33_70b": "models.common.models.llama33_70b.generator:Llama33_70BGenerator",
263
+ "mistral_7b": "models.common.models.mistral_7b.generator:Mistral7BGenerator",
264
+ "phi4": "models.common.models.phi4.generator:Phi4Generator",
265
+ "qwen2_7b": "models.common.models.qwen2_7b.generator:Qwen2Generator",
266
+ "qwen25_7b": "models.common.models.qwen25_7b.generator:Qwen25Generator",
267
+ "qwen25_72b": "models.common.models.qwen25_72b.generator:Qwen25_72BGenerator",
268
+ "qwen25_coder_32b": "models.common.models.qwen25_coder_32b.generator:Qwen25Coder32BGenerator",
269
+ "qwen3_32b": "models.common.models.qwen3_32b.generator:Qwen3_32BGenerator",
270
+ "deepseek_r1_distill_qwen_14b": (
271
+ "models.common.models.deepseek_r1_distill_qwen_14b.generator:DeepSeekR1Qwen14BGenerator"
272
+ ),
273
+ }
274
+
275
+
276
+ @pytest.fixture(params=EXECUTOR_BINDINGS.items(), ids=lambda item: item[0])
277
+ def binding(request):
278
+ return request.param[1]
279
+
280
+
281
+ _LLAMA_FAMILY_EXECUTOR_MODULES = (llama32_executor, llama32_3b_executor, llama33_70b_executor)
282
+ _QWEN2_FAMILY_EXECUTOR_MODULES = (
283
+ qwen2_executor,
284
+ qwen25_executor,
285
+ qwen25_72b_executor,
286
+ qwen25_coder_32b_executor,
287
+ )
288
+ _SHARED_MODEL_EXECUTOR_MODULES = (
289
+ *_LLAMA_FAMILY_EXECUTOR_MODULES,
290
+ *_QWEN2_FAMILY_EXECUTOR_MODULES,
291
+ qwen3_32b_executor,
292
+ )
293
+
294
+
295
+ def _composition_module(binding):
296
+ if binding.executor_module in _SHARED_MODEL_EXECUTOR_MODULES:
297
+ return shared_model_executor
298
+ return binding.executor_module
299
+
300
+
301
+ def _sampling_policy_module(binding):
302
+ if binding.executor_module in _LLAMA_FAMILY_EXECUTOR_MODULES:
303
+ return llama3_family_executor
304
+ if binding.executor_module in _QWEN2_FAMILY_EXECUTOR_MODULES:
305
+ return qwen2_family_executor
306
+ return binding.executor_module
307
+
308
+
309
+ class _Mesh:
310
+ shape = (1, 1)
311
+
312
+ @staticmethod
313
+ def get_num_devices():
314
+ return 1
315
+
316
+
317
+ class _Mesh2:
318
+ shape = (1, 2)
319
+
320
+ @staticmethod
321
+ def get_num_devices():
322
+ return 2
323
+
324
+
325
+ class _Mesh8:
326
+ shape = (1, 8)
327
+
328
+ @staticmethod
329
+ def get_num_devices():
330
+ return 8
331
+
332
+
333
+ def _make_llama32_model(max_batch_size=4):
334
+ paged = SimpleNamespace(block_size=32, max_num_blocks=132)
335
+ attention = SimpleNamespace(
336
+ n_kv_heads=8,
337
+ head_dim=64,
338
+ kv_cache_dtype=ttnn.bfloat8_b,
339
+ paged_attention_config=paged,
340
+ use_vllm_paged_kv_cache=True,
341
+ kv_cache=None,
342
+ )
343
+ live = SimpleNamespace(config=attention, kv_cache=None)
344
+ model = SimpleNamespace(
345
+ config=SimpleNamespace(
346
+ mesh_device=_Mesh(),
347
+ max_batch_size=max_batch_size,
348
+ max_seq_len=4096,
349
+ n_layers=1,
350
+ num_devices=1,
351
+ block_configs=(SimpleNamespace(attention_config=attention),),
352
+ ),
353
+ layers=(SimpleNamespace(attention=live),),
354
+ iter_executor_named_modules=lambda: (),
355
+ vocab_size=128256,
356
+ num_devices=1,
357
+ )
358
+
359
+ def configure_paged_attention(*, block_size, max_num_blocks):
360
+ assert attention.kv_cache is None
361
+ assert live.kv_cache is None
362
+ attention.paged_attention_config = SimpleNamespace(
363
+ block_size=block_size,
364
+ max_num_blocks=max_num_blocks,
365
+ )
366
+
367
+ model.configure_paged_attention = configure_paged_attention
368
+ return model
369
+
370
+
371
+ def _make_llama32_runtime_config():
372
+ return SimpleNamespace(
373
+ model_cache_path="cache",
374
+ max_prefill_chunk_size=2048,
375
+ trace_prefill_supported_seq_lens=(128,),
376
+ can_enable_trace=lambda length, num_cached_tokens=0: length == 128,
377
+ supports_batched_prefill=True,
378
+ disable_batched_prefill=False,
379
+ max_prefill_batch_size=32,
380
+ batched_prefill_batched_extract=True,
381
+ )
382
+
383
+
384
+ def _make_llama32_executor_config(mode="none", *, module=llama32_executor):
385
+ config_class = next(
386
+ getattr(module, name)
387
+ for name in (
388
+ "Llama32_1BExecutorConfig",
389
+ "Llama32_3BExecutorConfig",
390
+ "Llama33_70BExecutorConfig",
391
+ )
392
+ if hasattr(module, name)
393
+ )
394
+ return config_class(
395
+ trace=TraceConfig(mode),
396
+ warmup=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
397
+ paged_kv_cache=PagedKVCacheConfig(block_size=32, max_num_blocks=132, dtype=ttnn.bfloat8_b),
398
+ device_sampling_enabled=False,
399
+ )
400
+
401
+
402
+ def _make_llama32_product(mesh_device, max_batch_size):
403
+ model = _make_llama32_model(max_batch_size=max_batch_size)
404
+ model.config.mesh_device = mesh_device
405
+ return SimpleNamespace(model=model, runtime_config=_make_llama32_runtime_config())
406
+
407
+
408
+ def _make_qwen2_model(max_batch_size=4):
409
+ model = _make_llama32_model(max_batch_size=max_batch_size)
410
+ model.config.mesh_device = _Mesh2()
411
+ model.config.num_devices = 2
412
+ model.num_devices = 2
413
+ attention = model.layers[0].attention.config
414
+ attention.n_kv_heads = 4
415
+ attention.head_dim = 128
416
+ return model
417
+
418
+
419
+ def _make_qwen2_runtime_config(*, max_prefill_batch_size=32):
420
+ runtime = _make_llama32_runtime_config()
421
+ runtime.trace_prefill_supported_seq_lens = (128, 1024)
422
+ runtime.can_enable_trace = lambda length, num_cached_tokens=0: length in (128, 1024)
423
+ runtime.max_prefill_batch_size = max_prefill_batch_size
424
+ return runtime
425
+
426
+
427
+ def _make_qwen2_executor_config(mode="none", *, module=qwen2_executor):
428
+ config_class = (
429
+ getattr(module, "Qwen2ExecutorConfig", None)
430
+ or getattr(module, "Qwen25ExecutorConfig", None)
431
+ or getattr(module, "Qwen25_72BExecutorConfig", None)
432
+ or getattr(module, "Qwen25Coder32BExecutorConfig", None)
433
+ or getattr(module, "Qwen3_32BExecutorConfig", None)
434
+ or getattr(module, "DeepSeekR1Qwen14BExecutorConfig", None)
435
+ or getattr(module, "Mistral7BExecutorConfig", None)
436
+ or module.Phi4ExecutorConfig
437
+ )
438
+ return config_class(
439
+ trace=TraceConfig(mode),
440
+ warmup=WarmupConfig(prefill_seq_lens=(128, 1024), prefill_batch_sizes=(1,)),
441
+ paged_kv_cache=PagedKVCacheConfig(block_size=32, max_num_blocks=132, dtype=ttnn.bfloat8_b),
442
+ device_sampling_enabled=False,
443
+ )
444
+
445
+
446
+ def _make_qwen2_product(mesh_device, max_batch_size, *, max_prefill_batch_size=32):
447
+ model = _make_qwen2_model(max_batch_size=max_batch_size)
448
+ model.config.mesh_device = mesh_device
449
+ return SimpleNamespace(
450
+ model=model,
451
+ runtime_config=_make_qwen2_runtime_config(max_prefill_batch_size=max_prefill_batch_size),
452
+ )
453
+
454
+
455
+ def _make_qwen25_72b_model(max_batch_size=4):
456
+ model = _make_llama32_model(max_batch_size=max_batch_size)
457
+ model.config.mesh_device = _Mesh8()
458
+ model.config.num_devices = 8
459
+ model.num_devices = 8
460
+ attention = model.layers[0].attention.config
461
+ attention.n_kv_heads = 8
462
+ attention.head_dim = 128
463
+ return model
464
+
465
+
466
+ def _make_qwen25_72b_runtime_config():
467
+ return _make_qwen2_runtime_config(max_prefill_batch_size=32)
468
+
469
+
470
+ def _make_qwen25_72b_product(mesh_device, max_batch_size):
471
+ model = _make_qwen25_72b_model(max_batch_size=max_batch_size)
472
+ model.config.mesh_device = mesh_device
473
+ return SimpleNamespace(model=model, runtime_config=_make_qwen25_72b_runtime_config())
474
+
475
+
476
+ def _make_qwen25_coder_32b_model(max_batch_size=4):
477
+ model = _make_qwen25_72b_model(max_batch_size=max_batch_size)
478
+ model.config.dim = 5120
479
+ model.config.n_heads = 40
480
+ model.config.hidden_dim = 27648
481
+ model.config.hf_model_id = "Qwen/Qwen2.5-Coder-32B-Instruct"
482
+ return model
483
+
484
+
485
+ def _make_qwen25_coder_32b_runtime_config():
486
+ runtime = _make_qwen25_72b_runtime_config()
487
+ runtime.max_prefill_chunk_size = 4096
488
+ runtime.can_enable_trace = lambda length, num_cached_tokens=0: num_cached_tokens == 0 and length in (128, 1024)
489
+ return runtime
490
+
491
+
492
+ def _make_qwen25_coder_32b_product(mesh_device, max_batch_size):
493
+ model = _make_qwen25_coder_32b_model(max_batch_size=max_batch_size)
494
+ model.config.mesh_device = mesh_device
495
+ return SimpleNamespace(model=model, runtime_config=_make_qwen25_coder_32b_runtime_config())
496
+
497
+
498
+ def _make_qwen3_32b_model(max_batch_size=4):
499
+ model = _make_qwen25_72b_model(max_batch_size=max_batch_size)
500
+ model.config.dim = 5120
501
+ model.config.n_heads = 64
502
+ model.config.hidden_dim = 27648
503
+ model.config.vocab_size = 151936
504
+ model.config.hf_model_id = "Qwen/Qwen3-32B"
505
+ model.padded_vocab_size = 152064
506
+ return model
507
+
508
+
509
+ def _make_qwen3_32b_runtime_config():
510
+ runtime = _make_qwen25_coder_32b_runtime_config()
511
+ runtime.max_prefill_chunk_size = 4096
512
+ return runtime
513
+
514
+
515
+ def _make_qwen3_32b_product(mesh_device, max_batch_size):
516
+ model = _make_qwen3_32b_model(max_batch_size=max_batch_size)
517
+ model.config.mesh_device = mesh_device
518
+ return SimpleNamespace(model=model, runtime_config=_make_qwen3_32b_runtime_config())
519
+
520
+
521
+ def test_qwen2_binding_preserves_tp2_runtime_and_sampling_defaults():
522
+ model = _make_qwen2_model()
523
+ _, num_layers, kv_heads_per_device, head_dim = qwen2_generator._model_kv_metadata(model)
524
+ runtime = _make_qwen2_runtime_config()
525
+ config = qwen2_generator.Qwen2GeneratorConfig(
526
+ hf_model="Qwen/Qwen2-7B-Instruct",
527
+ hf_revision="test-revision",
528
+ mesh_device=model.config.mesh_device,
529
+ max_batch_size=32,
530
+ max_seq_len=4096,
531
+ )
532
+
533
+ assert model.config.mesh_device.shape == (1, 2)
534
+ assert (num_layers, kv_heads_per_device, head_dim) == (1, 2, 128)
535
+ assert runtime.trace_prefill_supported_seq_lens == (128, 1024)
536
+ assert runtime.max_prefill_chunk_size == 2048
537
+ assert runtime.max_prefill_batch_size == 32
538
+ assert config.hf_revision == "test-revision"
539
+ assert config.device_sampling_enabled is False
540
+
541
+
542
+ def test_qwen25_72b_binding_preserves_tp8_runtime_and_sampling_defaults():
543
+ model = _make_qwen25_72b_model()
544
+ _, num_layers, kv_heads_per_device, head_dim = qwen25_72b_generator._model_kv_metadata(model)
545
+ runtime = _make_qwen25_72b_runtime_config()
546
+ config = qwen25_72b_generator.Qwen25_72BGeneratorConfig(
547
+ hf_model="Qwen/Qwen2.5-72B-Instruct",
548
+ mesh_device=model.config.mesh_device,
549
+ max_batch_size=32,
550
+ max_seq_len=4096,
551
+ )
552
+
553
+ assert model.config.mesh_device.shape == (1, 8)
554
+ assert (num_layers, kv_heads_per_device, head_dim) == (1, 1, 128)
555
+ assert runtime.trace_prefill_supported_seq_lens == (128, 1024)
556
+ assert runtime.max_prefill_chunk_size == 2048
557
+ assert runtime.max_prefill_batch_size == 32
558
+ assert config.hf_revision == qwen25_72b_generator.DEFAULT_HF_REVISION
559
+ assert config.device_sampling_enabled is False
560
+
561
+
562
+ def test_qwen25_coder_32b_binding_preserves_tp8_runtime_and_sampling_defaults():
563
+ model = _make_qwen25_coder_32b_model()
564
+ _, num_layers, kv_heads_per_device, head_dim = qwen25_coder_32b_generator._model_kv_metadata(model)
565
+ runtime = _make_qwen25_coder_32b_runtime_config()
566
+ config = qwen25_coder_32b_generator.Qwen25Coder32BGeneratorConfig(
567
+ hf_model="Qwen/Qwen2.5-Coder-32B-Instruct",
568
+ mesh_device=model.config.mesh_device,
569
+ max_batch_size=32,
570
+ max_seq_len=4096,
571
+ )
572
+
573
+ assert model.config.mesh_device.shape == (1, 8)
574
+ assert (num_layers, kv_heads_per_device, head_dim) == (1, 1, 128)
575
+ assert runtime.trace_prefill_supported_seq_lens == (128, 1024)
576
+ assert runtime.max_prefill_chunk_size == 4096
577
+ assert runtime.can_enable_trace(128, 0) is True
578
+ assert runtime.can_enable_trace(128, 1) is False
579
+ assert runtime.max_prefill_batch_size == 32
580
+ assert config.hf_revision == qwen25_coder_32b_generator.DEFAULT_HF_REVISION
581
+ assert config.device_sampling_enabled is False
582
+
583
+
584
+ @pytest.mark.parametrize(
585
+ "product_binding",
586
+ (EXECUTOR_BINDINGS["qwen25_72b"], EXECUTOR_BINDINGS["qwen25_coder_32b"]),
587
+ ids=("qwen25_72b", "qwen25_coder_32b"),
588
+ )
589
+ @pytest.mark.parametrize("device_sampling_enabled", (False, True), ids=("sampling-off", "sampling-on"))
590
+ def test_large_qwen_builder_threads_exact_decode_sampling_coverage(
591
+ monkeypatch,
592
+ product_binding,
593
+ device_sampling_enabled,
594
+ ):
595
+ mesh_device = _Mesh()
596
+ product = product_binding.make_product(mesh_device, 4)
597
+ executor_configs = []
598
+
599
+ monkeypatch.setattr(product_binding.generator_module, "from_pretrained", lambda **kwargs: product)
600
+ monkeypatch.setattr(
601
+ product_binding.generator_module,
602
+ "_model_kv_metadata",
603
+ lambda model: ((ttnn.bfloat8_b,), 1, 1, 128),
604
+ )
605
+
606
+ def build_executor(llm, config):
607
+ executor_configs.append(config)
608
+ return product_binding.make_lane(llm, config)
609
+
610
+ monkeypatch.setattr(product_binding.generator_module, product_binding.build_executor_name, build_executor)
611
+ generator = getattr(product_binding.generator_module, product_binding.build_generator_name)(
612
+ product_binding.generator_config_class(
613
+ hf_model=product_binding.hf_model,
614
+ mesh_device=mesh_device,
615
+ max_batch_size=4,
616
+ max_seq_len=1024,
617
+ n_layers=1,
618
+ device_sampling_enabled=device_sampling_enabled,
619
+ )
620
+ )
621
+
622
+ try:
623
+ assert len(executor_configs) == 1
624
+ warmup = executor_configs[0].warmup
625
+ assert warmup.include_decode_top_k is device_sampling_enabled
626
+ plan = _build_plan(
627
+ warmup=warmup,
628
+ layout=PageTableLayout(block_size=32, raw_capacity_width=32, prefill_width=64, decode_width=32),
629
+ prefill_sequence_lengths=(128,),
630
+ lane_batch_size=4,
631
+ allow_force_argmax=True,
632
+ can_sample_on_device=device_sampling_enabled,
633
+ )
634
+ assert [case.sampling_path for case in plan.decode] == (
635
+ ["logits", "argmax", "topk"] if device_sampling_enabled else ["logits"]
636
+ )
637
+ finally:
638
+ generator.cleanup()
639
+
640
+
641
+ def test_qwen3_32b_binding_preserves_tp8_runtime_and_padded_vocab_defaults():
642
+ model = _make_qwen3_32b_model()
643
+ _, num_layers, kv_heads_per_device, head_dim = qwen3_32b_generator._model_kv_metadata(model)
644
+ runtime = _make_qwen3_32b_runtime_config()
645
+ config = qwen3_32b_generator.Qwen3_32BGeneratorConfig(
646
+ hf_model="Qwen/Qwen3-32B",
647
+ mesh_device=model.config.mesh_device,
648
+ max_batch_size=32,
649
+ max_seq_len=4096,
650
+ )
651
+
652
+ assert model.config.mesh_device.shape == (1, 8)
653
+ assert (num_layers, kv_heads_per_device, head_dim) == (1, 1, 128)
654
+ assert model.config.vocab_size == 151936
655
+ assert model.padded_vocab_size == 152064
656
+ assert runtime.trace_prefill_supported_seq_lens == (128, 1024)
657
+ assert runtime.max_prefill_chunk_size == 4096
658
+ assert runtime.can_enable_trace(128, 0) is True
659
+ assert runtime.can_enable_trace(128, 1) is False
660
+ assert config.hf_revision == qwen3_32b_generator.DEFAULT_HF_REVISION
661
+ assert config.device_sampling_enabled is False
662
+
663
+
664
+ @pytest.mark.parametrize(
665
+ ("cluster_shape", "disable_batched_prefill", "advertised_lengths", "expected"),
666
+ [
667
+ ([1, 8], False, (128, 1024), (128, 1024)),
668
+ ([1, 8], True, (128, 1024), (128, 1024)),
669
+ ([1, 4], False, (128, 1024), ()),
670
+ ([1, 8], False, (128,), (128,)),
671
+ ],
672
+ ids=("t3k-batched", "t3k-sequential", "bh-batched", "t3k-low-ceiling"),
673
+ )
674
+ def test_qwen3_prefill_capture_primes_are_t3k_product_owned(
675
+ cluster_shape,
676
+ disable_batched_prefill,
677
+ advertised_lengths,
678
+ expected,
679
+ ):
680
+ runtime = SimpleNamespace(
681
+ cluster_shape=cluster_shape,
682
+ disable_batched_prefill=disable_batched_prefill,
683
+ trace_prefill_supported_seq_lens=advertised_lengths,
684
+ can_enable_trace=lambda length, cached: length in advertised_lengths and cached == 0,
685
+ )
686
+
687
+ num_devices = int(cluster_shape[0]) * int(cluster_shape[1])
688
+ assert (
689
+ qwen3_32b_executor._resolve_trace_capture_prime_sequence_lengths(
690
+ runtime,
691
+ num_devices=num_devices,
692
+ )
693
+ == expected
694
+ )
695
+
696
+
697
+ def test_phi4_binding_preserves_cap8_trace_buckets_and_pinned_revision():
698
+ runtime = _make_qwen2_runtime_config(max_prefill_batch_size=8)
699
+ config = phi4_generator.Phi4GeneratorConfig(
700
+ hf_model="microsoft/phi-4",
701
+ mesh_device=_Mesh2(),
702
+ max_batch_size=32,
703
+ max_seq_len=4096,
704
+ )
705
+
706
+ assert runtime.trace_prefill_supported_seq_lens == (128, 1024)
707
+ assert runtime.max_prefill_chunk_size == 2048
708
+ assert runtime.max_prefill_batch_size == 8
709
+ assert config.hf_revision == phi4_generator.DEFAULT_HF_REVISION
710
+ assert config.device_sampling_enabled is False
711
+
712
+
713
+ @pytest.mark.parametrize("mode", ["none", "decode_only", "all"])
714
+ def test_model_owned_executor_has_exact_composition_and_owner_counts(binding, mode, monkeypatch):
715
+ owner_names = (
716
+ "PagedKVCacheManager",
717
+ "OutputReader",
718
+ "PrefillRuntime",
719
+ "DecodeRuntime",
720
+ "ProgramCompiler",
721
+ "EagerExecutor",
722
+ "TraceCompiler",
723
+ "TracedExecutor",
724
+ "WarmupCoordinator",
725
+ )
726
+ composition_module = _composition_module(binding)
727
+ owner_factories = {}
728
+ for name in owner_names:
729
+ factory = MagicMock(wraps=getattr(composition_module, name))
730
+ monkeypatch.setattr(composition_module, name, factory)
731
+ owner_factories[name] = factory
732
+
733
+ executor = binding.executor_class(
734
+ binding.make_model(),
735
+ binding.make_runtime_config(),
736
+ binding.make_executor_config(mode),
737
+ )
738
+ expected_counts = {name: 1 for name in owner_names}
739
+ if mode == "none":
740
+ expected_counts["TraceCompiler"] = 0
741
+ expected_counts["TracedExecutor"] = 0
742
+ assert {name: factory.call_count for name, factory in owner_factories.items()} == expected_counts
743
+ assert executor.eager_executor.program_compiler is executor.program_compiler
744
+ assert executor.eager_executor.prefill is executor.prefill_runtime
745
+ assert executor.eager_executor.decode is executor.decode_runtime
746
+ assert executor.warmup.eager is executor.eager_executor
747
+ assert executor.warmup.trace_compiler is executor.trace_compiler
748
+ assert executor.eager_execution is executor.eager_executor
749
+ assert executor.prefill_runtime.config.trace_capture_prime_sequence_lengths == (
750
+ (128, 1024) if binding.executor_module is qwen3_32b_executor else ()
751
+ )
752
+ if mode == "none":
753
+ assert executor.warmup.execution is executor.eager_executor
754
+ assert executor.trace_compiler is None
755
+ assert executor.traced_executor is None
756
+ assert executor.traced_prefill_execution is None
757
+ assert executor.traced_decode_execution is None
758
+ else:
759
+ assert executor.warmup.execution is executor.traced_executor
760
+ assert executor.traced_executor.eager_executor is executor.eager_executor
761
+ assert executor.traced_executor.trace_compiler is executor.trace_compiler
762
+ assert executor.trace_compiler.program_compiler is executor.program_compiler
763
+ assert (executor.traced_prefill_execution is not None) is (mode == "all")
764
+ assert executor.traced_decode_execution is executor.traced_executor
765
+
766
+
767
+ def test_llama3_create_sampling_state_disables_duplicate_seed_salting(monkeypatch):
768
+ captured = {}
769
+
770
+ class FakeSampling1D:
771
+ def __init__(self):
772
+ self.config = SimpleNamespace(is_resolved=lambda: True)
773
+
774
+ class FakeSamplingState1D:
775
+ def __init__(self, sampling, *, salt_duplicate_seeds=True, **_kwargs):
776
+ captured["salt_duplicate_seeds"] = salt_duplicate_seeds
777
+ self.sampling = sampling
778
+
779
+ def create_state(self):
780
+ return object()
781
+
782
+ monkeypatch.setattr(llama3_family_executor, "Sampling1D", FakeSampling1D)
783
+ monkeypatch.setattr(llama3_family_executor, "SamplingState1D", FakeSamplingState1D)
784
+
785
+ controller, state = llama3_family_executor._create_sampling_state(SimpleNamespace(sampling=FakeSampling1D()), True)
786
+
787
+ assert captured["salt_duplicate_seeds"] is False
788
+ assert controller is not None
789
+ assert state is not None
790
+
791
+
792
+ def test_qwen3_create_sampling_state_disables_duplicate_seed_salting(monkeypatch):
793
+ captured = {}
794
+
795
+ class FakeSampling1D:
796
+ def __init__(self):
797
+ self.config = SimpleNamespace(is_resolved=lambda: True)
798
+
799
+ class FakeSamplingState1D:
800
+ def __init__(self, sampling, *, salt_duplicate_seeds=True, **_kwargs):
801
+ captured["salt_duplicate_seeds"] = salt_duplicate_seeds
802
+ self.sampling = sampling
803
+
804
+ def create_state(self):
805
+ return object()
806
+
807
+ monkeypatch.setattr(qwen3_32b_executor, "Sampling1D", FakeSampling1D)
808
+ monkeypatch.setattr(qwen3_32b_executor, "SamplingState1D", FakeSamplingState1D)
809
+
810
+ controller, state = qwen3_32b_executor._create_sampling_state(SimpleNamespace(sampling=FakeSampling1D()), True)
811
+
812
+ assert captured["salt_duplicate_seeds"] is False
813
+ assert controller is not None
814
+ assert state is not None
815
+
816
+
817
+ def _device_sampling_executor(binding, monkeypatch, *, runtime_disable: bool):
818
+ class FakeSampling1D:
819
+ config = SimpleNamespace(
820
+ is_resolved=lambda: True,
821
+ allow_force_argmax=False,
822
+ max_batch_size=32,
823
+ max_top_k=32,
824
+ )
825
+
826
+ def decode_forward(self):
827
+ raise AssertionError("construction-policy test must not execute sampling")
828
+
829
+ sampling_policy_module = _sampling_policy_module(binding)
830
+ monkeypatch.setattr(sampling_policy_module, "Sampling1D", FakeSampling1D)
831
+ model = binding.make_model()
832
+ model.sampling = FakeSampling1D()
833
+ if binding.executor_module in (llama33_70b_executor, qwen3_32b_executor):
834
+
835
+ class FakeSamplingState1D:
836
+ def __init__(self, sampling, **_kwargs):
837
+ self.sampling = sampling
838
+ self.seed_manager = SimpleNamespace()
839
+
840
+ def create_state(self):
841
+ return SimpleNamespace(seed_state=SimpleNamespace(capacity=32))
842
+
843
+ def admit(self, *args, **kwargs):
844
+ return None
845
+
846
+ def decode_forward(self, *args, **kwargs):
847
+ return None
848
+
849
+ def release(self, *args, **kwargs):
850
+ return None
851
+
852
+ monkeypatch.setattr(sampling_policy_module, "SamplingState1D", FakeSamplingState1D)
853
+ runtime_config = binding.make_runtime_config()
854
+ runtime_config.disable_batched_prefill = runtime_disable
855
+ config = replace(
856
+ binding.make_executor_config("none"),
857
+ device_sampling_enabled=True,
858
+ )
859
+ return binding.executor_class(model, runtime_config, config)
860
+
861
+
862
+ @pytest.mark.parametrize(
863
+ ("runtime_disable", "environment_disable", "expected_kinds"),
864
+ [
865
+ (False, False, ("batched",)),
866
+ (True, False, ("single", "single")),
867
+ (False, True, ("single", "single")),
868
+ ],
869
+ ids=("device-sampled-batched", "runtime-disabled", "environment-disabled"),
870
+ )
871
+ def test_device_sampling_prefill_batch_policy_is_model_owned(
872
+ binding,
873
+ monkeypatch,
874
+ runtime_disable,
875
+ environment_disable,
876
+ expected_kinds,
877
+ ):
878
+ if environment_disable:
879
+ monkeypatch.setenv("DISABLE_BATCHED_PREFILL", "1")
880
+ else:
881
+ monkeypatch.delenv("DISABLE_BATCHED_PREFILL", raising=False)
882
+ executor = _device_sampling_executor(
883
+ binding,
884
+ monkeypatch,
885
+ runtime_disable=runtime_disable,
886
+ )
887
+
888
+ prepared = executor.prefill_runtime.prepare(
889
+ tokens=torch.ones((2, 128), dtype=torch.long),
890
+ page_table=torch.arange(8, dtype=torch.int32).reshape(2, 4),
891
+ prompt_lens=torch.full((2,), 128, dtype=torch.long),
892
+ empty_slots=[0, 1],
893
+ )
894
+
895
+ if binding.executor_module in (llama33_70b_executor, qwen3_32b_executor):
896
+ expected_kinds = ("single", "single")
897
+ assert tuple(item.request.kind for item in prepared) == expected_kinds
898
+ if expected_kinds == ("batched",):
899
+ assert prepared[0].request.source_rows == (0, 1)
900
+ assert not prepared[0].request.uses_chunked_prefill
901
+
902
+
903
+ def test_llama32_1b_warms_every_q128_topk_tile_start_once_per_execution_mode():
904
+ executor = object.__new__(llama32_executor.Llama32_1BExecutor)
905
+ executor._q128_topk_tile_ends_warmed = set()
906
+ executor.eager_executor = object()
907
+ executor.traced_executor = object()
908
+ executor.page_table_layout = SimpleNamespace(block_size=32)
909
+ executor.prefill_runtime = SimpleNamespace(config=SimpleNamespace(static_q128_topk_supported=True))
910
+ executor.warmup = SimpleNamespace(
911
+ config=SimpleNamespace(
912
+ prefill_sequence_lengths=(128,),
913
+ prime_q128_tile_ends=False,
914
+ )
915
+ )
916
+ executor.compile_prefill = MagicMock()
917
+ kv_cache = object()
918
+
919
+ for enable_trace in (False, False, True, True):
920
+ executor._warmup_q128_topk_tile_ends(
921
+ kv_cache=kv_cache,
922
+ can_sample_on_device=True,
923
+ enable_trace=enable_trace,
924
+ )
925
+
926
+ assert executor.compile_prefill.call_count == 6
927
+ calls = executor.compile_prefill.call_args_list
928
+ assert [call.kwargs["tokens"].shape[1] for call in calls] == [32, 64, 96, 32, 64, 96]
929
+ assert [call.kwargs["page_table"].shape[1] for call in calls] == [1, 2, 3, 1, 2, 3]
930
+ assert all(call.kwargs["kv_cache"] is kv_cache for call in calls)
931
+ assert all(call.kwargs["execution"] is executor.eager_executor for call in calls[:3])
932
+ assert all(call.kwargs["execution"] is executor.traced_executor for call in calls[3:])
933
+ assert executor._q128_topk_tile_ends_warmed == {False, True}
934
+
935
+
936
+ @pytest.mark.parametrize(
937
+ ("enable_trace", "expected_order"),
938
+ [
939
+ (False, ("default", 96, 0, 32, 64)),
940
+ (True, (0, 32, 64, "default", 96)),
941
+ ],
942
+ ids=("eager", "traced"),
943
+ )
944
+ def test_qwen3_lane4_warms_every_runtime_q128_topk_signature_before_activation(enable_trace, expected_order):
945
+ executor = SimpleNamespace(
946
+ _q128_topk_tile_ends_warmed=set(),
947
+ eager_executor=object(),
948
+ traced_executor=object(),
949
+ page_table_layout=SimpleNamespace(block_size=32),
950
+ prefill_runtime=SimpleNamespace(config=SimpleNamespace(static_q128_topk_supported=True)),
951
+ warmup=SimpleNamespace(
952
+ config=SimpleNamespace(
953
+ prefill_sequence_lengths=(128, 1024),
954
+ prime_q128_tile_ends=False,
955
+ )
956
+ ),
957
+ )
958
+ compiled = []
959
+ order = []
960
+ activation = []
961
+
962
+ def record_signature(prompt_length):
963
+ assert not activation
964
+ order.append(((prompt_length - 1) // 32) * 32)
965
+ compiled.append(
966
+ PrefillProgramSignature(
967
+ operation_variant="regular-single",
968
+ padded_batch_size=1,
969
+ invocation_sequence_length=128,
970
+ page_table_width=64,
971
+ chunk_page_table_width=None,
972
+ sampling_path="topk",
973
+ penalties_enabled=False,
974
+ logprobs_enabled=False,
975
+ last_token_tile_start=((prompt_length - 1) // 32) * 32,
976
+ )
977
+ )
978
+
979
+ def compile_prefill(**kwargs):
980
+ expected_execution = executor.traced_executor if enable_trace else executor.eager_executor
981
+ assert kwargs["execution"] is expected_execution
982
+ assert kwargs["kv_cache"] == "cache"
983
+ assert kwargs["sampling_params"].top_k.tolist() == [32]
984
+ record_signature(int(kwargs["prompt_lens"][0]))
985
+
986
+ executor.compile_prefill = compile_prefill
987
+
988
+ default_warmed = False
989
+
990
+ def default_warmup():
991
+ nonlocal default_warmed
992
+ if default_warmed:
993
+ return
994
+ default_warmed = True
995
+ order.append("default")
996
+ # The coordinator's ordinary Q128 top-k case covers the final tile.
997
+ record_signature(128)
998
+
999
+ for _ in range(2):
1000
+ qwen3_32b_executor._warmup_q128_around_prefill(
1001
+ executor,
1002
+ default_warmup,
1003
+ kv_cache="cache",
1004
+ can_sample_on_device=True,
1005
+ enable_trace=enable_trace,
1006
+ )
1007
+ activation.append(True)
1008
+
1009
+ assert tuple(order) == expected_order
1010
+ assert {signature.last_token_tile_start for signature in compiled} == {0, 32, 64, 96}
1011
+ compiled_keys = {ProgramKey.from_signature(signature) for signature in compiled}
1012
+ runtime_signatures = {replace(compiled[0], last_token_tile_start=tile_start) for tile_start in (0, 32, 64, 96)}
1013
+ assert {ProgramKey.from_signature(signature) for signature in runtime_signatures} == compiled_keys
1014
+ assert executor._q128_topk_tile_ends_warmed == {enable_trace}
1015
+
1016
+
1017
+ def test_qwen3_lane4_executor_installs_model_owned_q128_warmup(monkeypatch):
1018
+ executor = _device_sampling_executor(EXECUTOR_BINDINGS["qwen3_32b"], monkeypatch, runtime_disable=False)
1019
+
1020
+ assert executor._prefill_warmup is qwen3_32b_executor._warmup_q128_around_prefill
1021
+ assert executor._q128_topk_tile_ends_warmed == set()
1022
+
1023
+
1024
+ @pytest.mark.parametrize("device_sampling_enabled", (False, True), ids=("disabled", "enabled"))
1025
+ def test_qwen3_generator_sampling_policy_controls_decode_topk_warmup(monkeypatch, device_sampling_enabled):
1026
+ mesh_device = _Mesh8()
1027
+ product = _make_qwen3_32b_product(mesh_device, max_batch_size=4)
1028
+ executor_configs = []
1029
+
1030
+ monkeypatch.setattr(qwen3_32b_generator, "from_pretrained", lambda **kwargs: product)
1031
+ monkeypatch.setattr(
1032
+ qwen3_32b_generator,
1033
+ "_model_kv_metadata",
1034
+ lambda model: ((ttnn.bfloat8_b,), 1, 8, 64),
1035
+ )
1036
+
1037
+ def build_executor(llm, config):
1038
+ executor_configs.append(config)
1039
+ return _FakeLane(
1040
+ llm,
1041
+ config,
1042
+ request_state_fields=qwen3_32b_executor.Qwen3_32BExecutor.request_state_fields,
1043
+ )
1044
+
1045
+ monkeypatch.setattr(qwen3_32b_generator, "build_qwen3_32b_executor", build_executor)
1046
+ generator = qwen3_32b_generator.build_qwen3_32b_generator(
1047
+ qwen3_32b_generator.Qwen3_32BGeneratorConfig(
1048
+ hf_model="Qwen/Qwen3-32B",
1049
+ mesh_device=mesh_device,
1050
+ max_batch_size=4,
1051
+ max_seq_len=1024,
1052
+ trace_mode="decode_only",
1053
+ device_sampling_enabled=device_sampling_enabled,
1054
+ )
1055
+ )
1056
+
1057
+ assert len(executor_configs) == 1
1058
+ assert executor_configs[0].warmup.include_decode_top_k is device_sampling_enabled
1059
+ if device_sampling_enabled:
1060
+ observed_missing_signature = DecodeTraceSignature(
1061
+ batch_size=4,
1062
+ page_table_width=32,
1063
+ sampling_path="topk",
1064
+ device_feedback=True,
1065
+ )
1066
+ assert (
1067
+ ProgramKey.from_signature(observed_missing_signature).digest
1068
+ == "6f8351f51a0c90eaea5fca6700b3887e380a015dad7ee6f6a9e8be971dfebbd5"
1069
+ )
1070
+ generator.cleanup()
1071
+
1072
+
1073
+ @pytest.mark.parametrize(
1074
+ "method,positional,keyword_only",
1075
+ [
1076
+ (
1077
+ "compile_prefill",
1078
+ ["self"],
1079
+ [
1080
+ "tokens",
1081
+ "page_table",
1082
+ "prompt_lens",
1083
+ "start_pos",
1084
+ "empty_slots",
1085
+ "kv_cache",
1086
+ "sampling_params",
1087
+ "prompt_tokens",
1088
+ "output_tokens",
1089
+ "slot_remap",
1090
+ "execution",
1091
+ ],
1092
+ ),
1093
+ (
1094
+ "compile_decode",
1095
+ ["self"],
1096
+ [
1097
+ "tokens",
1098
+ "start_pos",
1099
+ "page_table",
1100
+ "kv_cache",
1101
+ "sampling_params",
1102
+ "prompt_tokens",
1103
+ "output_tokens",
1104
+ "slot_remap",
1105
+ "reset_batch",
1106
+ "execution",
1107
+ ],
1108
+ ),
1109
+ (
1110
+ "prefill_forward",
1111
+ ["self", "tokens", "page_table"],
1112
+ [
1113
+ "prompt_lens",
1114
+ "start_pos",
1115
+ "empty_slots",
1116
+ "kv_cache",
1117
+ "sampling_params",
1118
+ "prompt_tokens",
1119
+ "output_tokens",
1120
+ "slot_remap",
1121
+ "execution",
1122
+ ],
1123
+ ),
1124
+ (
1125
+ "decode_forward",
1126
+ ["self", "tokens", "start_pos", "page_table"],
1127
+ [
1128
+ "kv_cache",
1129
+ "sampling_params",
1130
+ "prompt_tokens",
1131
+ "output_tokens",
1132
+ "slot_remap",
1133
+ "reset_batch",
1134
+ "read_from_device",
1135
+ "execution",
1136
+ ],
1137
+ ),
1138
+ ("read_decode_output", ["self", "tt_out"], ["async_read"]),
1139
+ ("process_decode_output_host", ["self", "tt_out"], ["is_tokens"]),
1140
+ ("can_trace_prefill", ["self"], ["tokens", "prompt_lens", "start_pos", "empty_slots"]),
1141
+ ("warmup_model_prefill", ["self"], ["kv_cache", "can_sample_on_device", "enable_trace"]),
1142
+ (
1143
+ "warmup_model_decode",
1144
+ ["self"],
1145
+ ["kv_cache", "max_batch_size", "num_blocks", "can_sample_on_device", "enable_trace"],
1146
+ ),
1147
+ ],
1148
+ )
1149
+ def test_executor_call_contract(binding, method, positional, keyword_only):
1150
+ if binding.executor_module not in (llama33_70b_executor, qwen3_32b_executor):
1151
+ keyword_only = [name for name in keyword_only if name not in {"prompt_tokens", "output_tokens", "slot_remap"}]
1152
+ signature = inspect.signature(getattr(binding.executor_class, method))
1153
+ parameters = signature.parameters
1154
+ required = {
1155
+ "compile_prefill": {"tokens", "page_table"},
1156
+ "compile_decode": {"tokens", "start_pos", "page_table"},
1157
+ "prefill_forward": {"tokens", "page_table"},
1158
+ "decode_forward": {"tokens", "start_pos", "page_table"},
1159
+ "read_decode_output": {"tt_out"},
1160
+ "process_decode_output_host": {"tt_out"},
1161
+ "can_trace_prefill": {"tokens"},
1162
+ "warmup_model_prefill": {"kv_cache", "can_sample_on_device", "enable_trace"},
1163
+ "warmup_model_decode": {
1164
+ "kv_cache",
1165
+ "max_batch_size",
1166
+ "num_blocks",
1167
+ "can_sample_on_device",
1168
+ "enable_trace",
1169
+ },
1170
+ }[method]
1171
+ non_none_defaults = {"reset_batch": False, "read_from_device": True, "async_read": False, "is_tokens": False}
1172
+
1173
+ assert list(parameters) == positional + keyword_only
1174
+ assert all(parameters[name].kind is inspect.Parameter.POSITIONAL_OR_KEYWORD for name in positional)
1175
+ assert all(parameters[name].kind is inspect.Parameter.KEYWORD_ONLY for name in keyword_only)
1176
+ for name, parameter in tuple(parameters.items())[1:]:
1177
+ expected_default = inspect.Parameter.empty if name in required else non_none_defaults.get(name)
1178
+ assert parameter.default == expected_default
1179
+ assert parameter.annotation is not inspect.Parameter.empty
1180
+ assert signature.return_annotation is not inspect.Signature.empty
1181
+
1182
+
1183
+ @pytest.mark.parametrize(
1184
+ "executor_class",
1185
+ [
1186
+ qwen3_32b_executor.Qwen3_32BExecutor,
1187
+ ],
1188
+ )
1189
+ def test_qwen3_delegates_resolved_sampling_warmup_and_activation_to_coordinator(executor_class):
1190
+ warmup = SimpleNamespace(warmup_prefill=MagicMock(), warmup_decode=MagicMock())
1191
+ executor = SimpleNamespace(_ensure_active=MagicMock(), warmup=warmup)
1192
+
1193
+ executor_class.warmup_model_prefill(
1194
+ executor,
1195
+ kv_cache="cache",
1196
+ can_sample_on_device=True,
1197
+ enable_trace=True,
1198
+ )
1199
+ executor_class.warmup_model_decode(
1200
+ executor,
1201
+ kv_cache="cache",
1202
+ max_batch_size=32,
1203
+ num_blocks=64,
1204
+ can_sample_on_device=True,
1205
+ enable_trace=True,
1206
+ )
1207
+
1208
+ warmup.warmup_prefill.assert_called_once_with(
1209
+ kv_cache="cache",
1210
+ can_sample_on_device=True,
1211
+ enable_trace=True,
1212
+ )
1213
+ warmup.warmup_decode.assert_called_once_with(
1214
+ kv_cache="cache",
1215
+ max_batch_size=32,
1216
+ num_blocks=64,
1217
+ can_sample_on_device=True,
1218
+ enable_trace=True,
1219
+ )
1220
+ assert "capture_all" not in inspect.getsource(executor_class.warmup_model_decode)
1221
+
1222
+
1223
+ class _RecordingTarget:
1224
+ model_args = object()
1225
+ mesh_device = object()
1226
+ cache_path = "cache"
1227
+ already_warmed_up_prefill = False
1228
+ eager_execution = object()
1229
+ traced_prefill_execution = object()
1230
+ traced_decode_execution = object()
1231
+
1232
+ def __init__(self, model, traceable=True, request_state_fields=()):
1233
+ self.model = model
1234
+ self.traceable = traceable
1235
+ self._request_state_fields = tuple(request_state_fields)
1236
+ self.calls = []
1237
+
1238
+ def can_trace_prefill(self, **kwargs):
1239
+ self.calls.append(("can_trace_prefill", kwargs))
1240
+ return self.traceable
1241
+
1242
+ def prefill_forward(self, **kwargs):
1243
+ self.calls.append(("prefill_forward", kwargs))
1244
+ return kwargs["execution"]
1245
+
1246
+ def decode_forward(self, **kwargs):
1247
+ self.calls.append(("decode_forward", kwargs))
1248
+ return kwargs["execution"]
1249
+
1250
+ def cleanup(self):
1251
+ self.calls.append(("cleanup", {}))
1252
+
1253
+
1254
+ def test_generator_preserves_required_trace_intent_for_ineligible_prefill(binding):
1255
+ target = binding.make_recording_target(traceable=False)
1256
+ target.config = binding.make_executor_config("all")
1257
+ generator = binding.generator_class(target, binding.generator_module._build_vllm_adapter(target))
1258
+ tokens = __import__("torch").tensor([[1]])
1259
+ page_table = __import__("torch").tensor([[0]], dtype=__import__("torch").int32)
1260
+ assert generator.prefill_forward(tokens, page_table, enable_trace=True) is target.traced_prefill_execution
1261
+ assert [name for name, _ in target.calls] == ["prefill_forward"]
1262
+
1263
+
1264
+ def test_generator_routes_external_decode_only_policy_with_all_trace_targets(binding):
1265
+ target = binding.make_recording_target()
1266
+ target.config = binding.make_executor_config("all")
1267
+ generator = binding.generator_class(target, binding.generator_module._build_vllm_adapter(target))
1268
+ torch = __import__("torch")
1269
+ tokens = torch.tensor([[1]])
1270
+ page_table = torch.tensor([[0]], dtype=torch.int32)
1271
+ start_pos = torch.tensor([0])
1272
+
1273
+ assert generator.prefill_forward(tokens, page_table, enable_trace=False) is target.eager_execution
1274
+ assert (
1275
+ generator.decode_forward(tokens[:, 0], start_pos, page_table, enable_trace=True)
1276
+ is target.traced_decode_execution
1277
+ )
1278
+ assert [name for name, _ in target.calls] == ["prefill_forward", "decode_forward"]
1279
+
1280
+
1281
+ def test_executor_validates_borrowed_cache_then_omits_it_from_execution(binding):
1282
+ execution = create_autospec(EagerExecutor, instance=True)
1283
+ events = []
1284
+ execution.compile_prefill.side_effect = lambda **kwargs: events.append("dispatch_compile_prefill")
1285
+ execution.compile_decode.side_effect = lambda **kwargs: events.append("dispatch_compile_decode")
1286
+ execution.prefill_forward.side_effect = lambda **kwargs: events.append("dispatch_prefill") or "prefill"
1287
+ execution.decode_forward.side_effect = lambda **kwargs: events.append("dispatch_decode") or "decode"
1288
+ executor = object.__new__(binding.executor_class)
1289
+ executor._prefill_execution = execution
1290
+ executor._decode_execution = execution
1291
+ executor._ensure_active = lambda: None
1292
+ executor._validate_bound_cache = lambda cache: events.append(("validate_cache", cache))
1293
+ executor._ensure_sampling_for = lambda params: events.append(("validate_sampling", params))
1294
+
1295
+ tokens = torch.zeros((1, 4), dtype=torch.long)
1296
+ start_pos = torch.zeros((1,), dtype=torch.long)
1297
+ page_table = torch.zeros((1, 1), dtype=torch.int32)
1298
+ prompt_lens = torch.full((1,), 4, dtype=torch.long)
1299
+ empty_slots = [0]
1300
+ kv_cache = object()
1301
+ sampling_params = object()
1302
+
1303
+ executor.compile_prefill(
1304
+ tokens=tokens,
1305
+ page_table=page_table,
1306
+ prompt_lens=prompt_lens,
1307
+ start_pos=start_pos,
1308
+ empty_slots=empty_slots,
1309
+ kv_cache=kv_cache,
1310
+ sampling_params=sampling_params,
1311
+ )
1312
+ executor.compile_decode(
1313
+ tokens=tokens,
1314
+ start_pos=start_pos,
1315
+ page_table=page_table,
1316
+ kv_cache=kv_cache,
1317
+ sampling_params=sampling_params,
1318
+ reset_batch=True,
1319
+ )
1320
+ assert (
1321
+ executor.prefill_forward(
1322
+ tokens,
1323
+ page_table,
1324
+ prompt_lens=prompt_lens,
1325
+ start_pos=start_pos,
1326
+ empty_slots=empty_slots,
1327
+ kv_cache=kv_cache,
1328
+ sampling_params=sampling_params,
1329
+ )
1330
+ == "prefill"
1331
+ )
1332
+ assert (
1333
+ executor.decode_forward(
1334
+ tokens,
1335
+ start_pos,
1336
+ page_table,
1337
+ kv_cache=kv_cache,
1338
+ sampling_params=sampling_params,
1339
+ reset_batch=True,
1340
+ read_from_device=False,
1341
+ )
1342
+ == "decode"
1343
+ )
1344
+
1345
+ expected_validation = [("validate_cache", kv_cache), ("validate_sampling", sampling_params)]
1346
+ assert events == [
1347
+ *expected_validation,
1348
+ "dispatch_compile_prefill",
1349
+ *expected_validation,
1350
+ "dispatch_compile_decode",
1351
+ *expected_validation,
1352
+ "dispatch_prefill",
1353
+ *expected_validation,
1354
+ "dispatch_decode",
1355
+ ]
1356
+ request_state_names = (
1357
+ ("prompt_tokens", "output_tokens", "slot_remap")
1358
+ if "prompt_tokens" in inspect.signature(binding.executor_class.compile_prefill).parameters
1359
+ else ()
1360
+ )
1361
+ for target, expected_names in (
1362
+ (
1363
+ execution.compile_prefill,
1364
+ (
1365
+ "tokens",
1366
+ "page_table",
1367
+ "prompt_lens",
1368
+ "start_pos",
1369
+ "empty_slots",
1370
+ "sampling_params",
1371
+ *request_state_names,
1372
+ ),
1373
+ ),
1374
+ (
1375
+ execution.compile_decode,
1376
+ ("tokens", "start_pos", "page_table", "sampling_params", *request_state_names, "reset_batch"),
1377
+ ),
1378
+ (
1379
+ execution.prefill_forward,
1380
+ (
1381
+ "tokens",
1382
+ "page_table",
1383
+ "prompt_lens",
1384
+ "start_pos",
1385
+ "empty_slots",
1386
+ "sampling_params",
1387
+ *request_state_names,
1388
+ ),
1389
+ ),
1390
+ (
1391
+ execution.decode_forward,
1392
+ (
1393
+ "tokens",
1394
+ "start_pos",
1395
+ "page_table",
1396
+ "sampling_params",
1397
+ *request_state_names,
1398
+ "reset_batch",
1399
+ "read_from_device",
1400
+ ),
1401
+ ),
1402
+ ):
1403
+ assert target.call_count == 1
1404
+ assert tuple(target.call_args.kwargs) == expected_names
1405
+ assert "kv_cache" not in target.call_args.kwargs
1406
+
1407
+
1408
+ def test_late_capacity_reconfigures_existing_owners_before_allocation(binding, monkeypatch):
1409
+ executor = binding.executor_class(
1410
+ binding.make_model(), binding.make_runtime_config(), binding.make_executor_config()
1411
+ )
1412
+ owner_ids = tuple(
1413
+ id(owner)
1414
+ for owner in (executor.prefill_runtime, executor.decode_runtime, executor.warmup, executor.program_compiler)
1415
+ )
1416
+ assert executor.page_table_layout.raw_capacity_width == 128
1417
+
1418
+ executor.configure_paged_kv_cache(
1419
+ PagedKVCacheConfig(
1420
+ block_size=16,
1421
+ max_num_blocks=200,
1422
+ dtype=ttnn.bfloat8_b,
1423
+ num_blocks=200,
1424
+ )
1425
+ )
1426
+
1427
+ assert (
1428
+ tuple(
1429
+ id(owner)
1430
+ for owner in (executor.prefill_runtime, executor.decode_runtime, executor.warmup, executor.program_compiler)
1431
+ )
1432
+ == owner_ids
1433
+ )
1434
+ assert executor.config.paged_kv_cache is executor.kv_cache_manager.config
1435
+ assert executor.kv_cache_manager.config.block_size == 16
1436
+ assert executor.kv_cache_manager.config.max_num_blocks == executor.kv_cache_manager.config.num_blocks == 200
1437
+ assert executor.model.layers[0].attention.config.paged_attention_config.block_size == 16
1438
+ assert executor.model.layers[0].attention.config.paged_attention_config.max_num_blocks == 200
1439
+ assert executor.page_table_layout.block_size == 16
1440
+ assert executor.page_table_layout.raw_capacity_width == 200
1441
+ assert executor.prefill_runtime.config.page_table_layout is executor.page_table_layout
1442
+ assert executor.decode_runtime.config.page_table_layout is executor.page_table_layout
1443
+ assert executor.warmup.config.page_table_layout is executor.page_table_layout
1444
+ assert executor.prefill_runtime.config.trace_capture_prime_sequence_lengths == (
1445
+ (128, 1024) if binding.executor_module is qwen3_32b_executor else ()
1446
+ )
1447
+
1448
+ def fake_allocate():
1449
+ assert executor._runtime_configuration_sealed
1450
+ assert executor.warmup._configuration_sealed
1451
+ return ["allocated"]
1452
+
1453
+ monkeypatch.setattr(executor.kv_cache_manager, "allocate", fake_allocate)
1454
+ assert executor.allocate_kv_cache() == ["allocated"]
1455
+
1456
+
1457
+ def test_late_capacity_failure_is_atomic(binding, expect_error):
1458
+ executor = binding.executor_class(
1459
+ binding.make_model(), binding.make_runtime_config(), binding.make_executor_config()
1460
+ )
1461
+ executor._seal_runtime_configuration()
1462
+ unresolved = executor.kv_cache_manager.config
1463
+ original_layout = executor.page_table_layout
1464
+ original_model_paged = executor.model.layers[0].attention.config.paged_attention_config
1465
+
1466
+ with expect_error(RuntimeError, "runtime configuration is sealed"):
1467
+ executor.configure_paged_kv_cache(
1468
+ PagedKVCacheConfig(
1469
+ block_size=32,
1470
+ max_num_blocks=132,
1471
+ dtype=ttnn.bfloat8_b,
1472
+ num_blocks=64,
1473
+ )
1474
+ )
1475
+
1476
+ assert executor.kv_cache_manager.config is unresolved
1477
+ assert not unresolved.is_resolved()
1478
+ assert executor.page_table_layout is original_layout
1479
+ assert executor.model.layers[0].attention.config.paged_attention_config is original_model_paged
1480
+
1481
+
1482
+ def test_generator_resolves_configures_then_allocates_vllm_kv_shape(binding):
1483
+ events = []
1484
+ resolved = object()
1485
+ cache = object()
1486
+ shape = (129, 8, 64, 128)
1487
+ dtype = object()
1488
+ target = SimpleNamespace(
1489
+ configure_paged_kv_cache=lambda config: events.append(("configure", config)),
1490
+ allocate_kv_cache=lambda: events.append(("allocate",)) or cache,
1491
+ )
1492
+ adapter = SimpleNamespace(
1493
+ resolve_legacy_kv_cache_config=lambda *args: events.append(("resolve", args)) or resolved,
1494
+ )
1495
+ generator = binding.generator_class(target, adapter)
1496
+
1497
+ assert generator.allocate_kv_cache(shape, dtype, 32) is cache
1498
+ assert events == [
1499
+ ("resolve", (shape, dtype, 32)),
1500
+ ("configure", resolved),
1501
+ ("allocate",),
1502
+ ]
1503
+
1504
+
1505
+ def test_generator_allocates_model_owned_kv_without_reconfiguration(binding):
1506
+ events = []
1507
+ cache = object()
1508
+ target = SimpleNamespace(
1509
+ configure_paged_kv_cache=lambda config: events.append(("configure", config)),
1510
+ allocate_kv_cache=lambda: events.append(("allocate",)) or cache,
1511
+ )
1512
+ adapter = SimpleNamespace(
1513
+ resolve_legacy_kv_cache_config=lambda *args: events.append(("resolve", args)),
1514
+ )
1515
+ generator = binding.generator_class(target, adapter)
1516
+
1517
+ assert generator.allocate_kv_cache() is cache
1518
+ assert events == [("allocate",)]
1519
+
1520
+
1521
+ @pytest.mark.parametrize(
1522
+ "arguments",
1523
+ (
1524
+ ((64, 8, 32, 128), None, None),
1525
+ (None, object(), None),
1526
+ (None, None, 32),
1527
+ ((64, 8, 32, 128), object(), None),
1528
+ ),
1529
+ )
1530
+ def test_generator_rejects_partial_vllm_kv_shape_atomically(binding, arguments, expect_error):
1531
+ events = []
1532
+ target = SimpleNamespace(
1533
+ configure_paged_kv_cache=lambda config: events.append(("configure", config)),
1534
+ allocate_kv_cache=lambda: events.append(("allocate",)),
1535
+ )
1536
+ adapter = SimpleNamespace(
1537
+ resolve_legacy_kv_cache_config=lambda *args: events.append(("resolve", args)),
1538
+ )
1539
+ generator = binding.generator_class(target, adapter)
1540
+
1541
+ with expect_error(TypeError, "must be supplied together"):
1542
+ generator.allocate_kv_cache(*arguments)
1543
+
1544
+ assert events == []
1545
+
1546
+
1547
+ def test_generator_does_not_configure_or_allocate_after_vllm_kv_resolution_failure(binding, expect_error):
1548
+ events = []
1549
+
1550
+ def fail_resolution(*args):
1551
+ events.append(("resolve", args))
1552
+ raise ValueError("invalid vLLM KV geometry")
1553
+
1554
+ target = SimpleNamespace(
1555
+ configure_paged_kv_cache=lambda config: events.append(("configure", config)),
1556
+ allocate_kv_cache=lambda: events.append(("allocate",)),
1557
+ )
1558
+ generator = binding.generator_class(
1559
+ target,
1560
+ SimpleNamespace(resolve_legacy_kv_cache_config=fail_resolution),
1561
+ )
1562
+
1563
+ with expect_error(ValueError, "invalid vLLM KV geometry"):
1564
+ generator.allocate_kv_cache((129, 8, 64, 128), object(), 32)
1565
+
1566
+ assert tuple(name for name, *_ in events) == ("resolve",)
1567
+
1568
+
1569
+ def test_generator_reports_unmultiplied_per_submesh_token_capacity(binding):
1570
+ assert (
1571
+ binding.generator_class.get_max_tokens_all_users(
1572
+ model_name="ignored",
1573
+ num_devices=8,
1574
+ tt_data_parallel=4,
1575
+ max_model_len=32768,
1576
+ max_num_seqs=64,
1577
+ )
1578
+ == 32768
1579
+ )
1580
+
1581
+
1582
+ def test_generator_rejects_unavailable_traced_execution(binding, expect_error):
1583
+ target = binding.make_recording_target()
1584
+ target.traced_decode_execution = None
1585
+ target.config = binding.make_executor_config("none")
1586
+ generator = binding.generator_class(target, binding.generator_module._build_vllm_adapter(target))
1587
+
1588
+ with expect_error(RuntimeError, "unavailable traced decode execution"):
1589
+ generator._select_execution("decode", True)
1590
+
1591
+
1592
+ def test_initialize_vllm_model_threads_policy(binding, monkeypatch):
1593
+ captured = []
1594
+ sentinel = object()
1595
+ mesh_device = object()
1596
+ monkeypatch.setattr(
1597
+ binding.generator_module,
1598
+ binding.build_generator_name,
1599
+ lambda config: captured.append(config) or sentinel,
1600
+ )
1601
+
1602
+ result = binding.generator_class.initialize_vllm_model(
1603
+ SimpleNamespace(_name_or_path=binding.hf_model),
1604
+ mesh_device,
1605
+ 8,
1606
+ 4096,
1607
+ n_layers=3,
1608
+ tt_data_parallel=2,
1609
+ optimizations="accuracy",
1610
+ trace_mode="decode_only",
1611
+ device_sampling_enabled=True,
1612
+ )
1613
+
1614
+ assert result is sentinel
1615
+ config = captured[0]
1616
+ assert isinstance(config, binding.generator_config_class)
1617
+ assert config.hf_model == binding.hf_model
1618
+ assert config.mesh_device is mesh_device
1619
+ assert config.max_batch_size == 8
1620
+ assert config.max_seq_len == 4096
1621
+ assert config.n_layers == 3
1622
+ assert config.tt_data_parallel == 2
1623
+ assert config.optimizations == "accuracy"
1624
+ assert config.trace_mode == "decode_only"
1625
+ assert config.device_sampling_enabled is True
1626
+
1627
+
1628
+ @pytest.mark.parametrize("model_id,generator_path", GENERATOR_PATHS.items(), ids=GENERATOR_PATHS)
1629
+ def test_vllm_generator_path_and_construction_defaults(model_id, generator_path):
1630
+ module_name, class_name = generator_path.split(":", maxsplit=1)
1631
+ generator_class = getattr(import_module(module_name), class_name)
1632
+
1633
+ assert generator_class is EXECUTOR_BINDINGS[model_id].generator_class
1634
+ assert callable(getattr(generator_class, "initialize_vllm_model", None))
1635
+
1636
+ parameters = inspect.signature(generator_class.initialize_vllm_model).parameters
1637
+ assert parameters["trace_mode"].default == "all"
1638
+ assert parameters["device_sampling_enabled"].default is True
1639
+
1640
+
1641
+ class _FakeLane:
1642
+ requires_prefill_trace_warmup = True
1643
+
1644
+ def __init__(self, llm, config, request_state_fields=()):
1645
+ self.model = llm.model
1646
+ self.model_args = llm.runtime_config
1647
+ self.mesh_device = llm.model.config.mesh_device
1648
+ self.cache_path = llm.runtime_config.model_cache_path
1649
+ self.config = config
1650
+ self._request_state_fields = tuple(request_state_fields)
1651
+ self.paged_kv_cache_config = config.paged_kv_cache
1652
+ self.already_warmed_up_prefill = False
1653
+ self.eager_execution = object()
1654
+ self.traced_prefill_execution = object()
1655
+ self.traced_decode_execution = object()
1656
+ self.cleanup_calls = 0
1657
+
1658
+ def cleanup(self):
1659
+ self.cleanup_calls += 1
1660
+
1661
+
1662
+ def test_generator_constructs_data_parallel_lane_group(binding, monkeypatch):
1663
+ executor_calls = []
1664
+ built_lanes = []
1665
+ pretrained_calls = []
1666
+ parent_mesh = object()
1667
+ submeshes = [_Mesh(), _Mesh()]
1668
+ create_submeshes = MagicMock(return_value=submeshes)
1669
+ monkeypatch.setattr(binding.generator_module, "_create_submeshes", create_submeshes)
1670
+
1671
+ def fake_from_pretrained(mesh_device, **kwargs):
1672
+ pretrained_calls.append((mesh_device, kwargs))
1673
+ return binding.make_product(mesh_device, kwargs["max_batch_size"])
1674
+
1675
+ def fake_build_executor(llm, config):
1676
+ executor_calls.append((llm, config))
1677
+ lane = binding.make_lane(llm, config)
1678
+ built_lanes.append(lane)
1679
+ return lane
1680
+
1681
+ monkeypatch.setattr(binding.generator_module, "from_pretrained", fake_from_pretrained)
1682
+ monkeypatch.setattr(binding.generator_module, binding.build_executor_name, fake_build_executor)
1683
+ monkeypatch.setattr(
1684
+ binding.generator_module,
1685
+ "_model_kv_metadata",
1686
+ lambda model: ((ttnn.bfloat8_b,), 1, 8, 64),
1687
+ )
1688
+
1689
+ generator = getattr(binding.generator_module, binding.build_generator_name)(
1690
+ binding.generator_config_class(
1691
+ hf_model=binding.hf_model,
1692
+ mesh_device=parent_mesh,
1693
+ max_batch_size=4,
1694
+ max_seq_len=4096,
1695
+ n_layers=1,
1696
+ tt_data_parallel=2,
1697
+ trace_mode="all",
1698
+ device_sampling_enabled=True,
1699
+ )
1700
+ )
1701
+
1702
+ try:
1703
+ create_submeshes.assert_called_once_with(parent_mesh, 2)
1704
+ assert [mesh for mesh, _ in pretrained_calls] == submeshes
1705
+ assert all(call[1]["max_batch_size"] == 2 for call in pretrained_calls)
1706
+ assert all(call[1]["max_seq_len"] == 4096 for call in pretrained_calls)
1707
+ assert all(call[1]["n_layers"] == 1 for call in pretrained_calls)
1708
+ assert isinstance(generator.target, LaneGroupExecutor)
1709
+ assert generator.target.mesh_device is parent_mesh
1710
+ assert generator.target.tt_data_parallel == 2
1711
+ assert len(executor_calls) == 2
1712
+ assert executor_calls[0][0] is not executor_calls[1][0]
1713
+ assert generator.target.lanes == built_lanes
1714
+ assert [lane.model for lane in generator.target.lanes] == [llm.model for llm, _ in executor_calls]
1715
+ assert [lane.mesh_device for lane in generator.target.lanes] == submeshes
1716
+ assert len({id(lane) for lane in generator.target.lanes}) == 2
1717
+ assert all(isinstance(config, binding.executor_config_class) for _, config in executor_calls)
1718
+ assert all(llm.model.config.max_batch_size == 2 for llm, _ in executor_calls)
1719
+ assert generator._adapter.config.trace.mode == "all"
1720
+ assert generator._adapter.config.expected_num_layers == 1
1721
+ assert generator._adapter.config.expected_kv_heads_per_device == 8
1722
+ assert generator._adapter.config.expected_head_dim == 64
1723
+ finally:
1724
+ generator.cleanup()
1725
+
1726
+
1727
+ def test_executor_cleanup_is_ordered_retryable_and_idempotent(binding, expect_error):
1728
+ calls = []
1729
+ failures = {"reader", "trace"}
1730
+
1731
+ class _Owner:
1732
+ def __init__(self, name):
1733
+ self.name = name
1734
+
1735
+ def cleanup(self, *args):
1736
+ calls.append(self.name)
1737
+ if self.name in failures:
1738
+ raise RuntimeError(self.name)
1739
+
1740
+ drain = cleanup
1741
+ drain_external_outputs = cleanup
1742
+ cleanup_transients = cleanup
1743
+ release = cleanup
1744
+
1745
+ executor = object.__new__(binding.executor_class)
1746
+ executor._terminal = False
1747
+ executor._cleaned_up = False
1748
+ executor.decode_runtime = _Owner("decode-external")
1749
+ executor.output_reader = _Owner("reader")
1750
+ executor.prefill_runtime = _Owner("prefill")
1751
+ executor.trace_compiler = _Owner("trace")
1752
+ executor.program_compiler = _Owner("program")
1753
+ executor.config = SimpleNamespace(device_sampling_enabled=True)
1754
+ executor.model = SimpleNamespace(sampling=_Owner("sampling"))
1755
+ if binding.executor_module in (llama33_70b_executor, qwen3_32b_executor):
1756
+ executor.sampling_state_controller = _Owner("sampling-state")
1757
+ executor.sampling_state = object()
1758
+ else:
1759
+ executor.sampling_state_controller = None
1760
+ executor.sampling_state = None
1761
+ executor.kv_cache_manager = _Owner("kv")
1762
+
1763
+ with expect_error(RuntimeError, "reader") as raised:
1764
+ executor.cleanup()
1765
+
1766
+ expected_order = [
1767
+ "decode-external",
1768
+ "reader",
1769
+ "prefill",
1770
+ "decode-external",
1771
+ "trace",
1772
+ "program",
1773
+ ]
1774
+ if binding.executor_module in (llama33_70b_executor, qwen3_32b_executor):
1775
+ expected_order.append("sampling-state")
1776
+ expected_order.extend(["sampling", "kv"])
1777
+ assert calls == expected_order
1778
+ assert tuple(error.args[0] for error in raised.value.cleanup_failures) == ("trace",)
1779
+ assert executor.terminal
1780
+ assert not executor._cleaned_up
1781
+
1782
+ failures.clear()
1783
+ executor.cleanup()
1784
+ assert calls == expected_order * 2
1785
+ assert executor._cleaned_up
1786
+
1787
+ executor.cleanup()
1788
+ assert calls == expected_order * 2
1789
+
1790
+
1791
+ def test_llama33_generator_emits_runtime_summary_before_owned_cleanup():
1792
+ events = []
1793
+ traced = SimpleNamespace(log_runtime_summary=lambda **kwargs: events.append(("summary", kwargs)))
1794
+ target = SimpleNamespace(
1795
+ traced_executor=traced,
1796
+ cleanup=lambda: events.append("cleanup"),
1797
+ )
1798
+ generator = llama33_70b_generator.Llama33_70BGenerator(target, SimpleNamespace())
1799
+
1800
+ generator.cleanup()
1801
+
1802
+ assert events == [("summary", {"phase": "shutdown"}), "cleanup"]
1803
+
1804
+
1805
+ def test_llama33_generator_emits_serving_ready_and_idempotent_shutdown_summaries():
1806
+ phases = []
1807
+ trace_compiler = SimpleNamespace(trace_active=False)
1808
+ traced = SimpleNamespace(
1809
+ trace_compiler=trace_compiler,
1810
+ log_runtime_summary=lambda **kwargs: phases.append(kwargs["phase"]),
1811
+ )
1812
+ target = SimpleNamespace(
1813
+ traced_executor=traced,
1814
+ warmup_model_prefill=lambda **kwargs: None,
1815
+ warmup_model_decode=lambda **kwargs: None,
1816
+ cleanup=lambda: None,
1817
+ )
1818
+ generator = llama33_70b_generator.Llama33_70BGenerator(target, SimpleNamespace())
1819
+
1820
+ generator.warmup_model_prefill(kv_cache="cache", can_sample_on_device=True, enable_trace=True)
1821
+ trace_compiler.trace_active = True
1822
+ generator.warmup_model_decode(
1823
+ kv_cache="cache",
1824
+ max_batch_size=16,
1825
+ num_blocks=128,
1826
+ can_sample_on_device=True,
1827
+ enable_trace=True,
1828
+ )
1829
+ generator.warmup_model_prefill(kv_cache="cache", can_sample_on_device=True, enable_trace=True)
1830
+ generator._shutdown_summary_callback()
1831
+ generator.cleanup()
1832
+
1833
+ assert phases == ["serving_ready", "shutdown"]
code/models/common/tests/llm_runtime/test_lane_group.py ADDED
@@ -0,0 +1,1373 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import inspect
5
+ from contextlib import contextmanager
6
+ from dataclasses import dataclass
7
+ from types import SimpleNamespace
8
+
9
+ import pytest
10
+ import torch
11
+
12
+ from models.common.llm_runtime.lane_group import LaneGroupExecutor
13
+
14
+
15
+ @dataclass(frozen=True)
16
+ class _SamplingParams:
17
+ temperature: list[float]
18
+ top_k: torch.Tensor
19
+ top_p: float
20
+
21
+
22
+ @dataclass(frozen=True)
23
+ class _SingletonSamplingParams:
24
+ temperature: list[float]
25
+ top_k: torch.Tensor
26
+ top_p: tuple[float, ...]
27
+
28
+
29
+ @dataclass(frozen=True)
30
+ class _StatefulSamplingParams:
31
+ temperature: list[float]
32
+ top_k: torch.Tensor
33
+ top_p: list[float]
34
+ presence_penalty: list[float]
35
+ frequency_penalty: list[float]
36
+ repetition_penalty: list[float]
37
+ seed: list[int | None]
38
+ enable_log_probs: list[bool]
39
+ num_logprobs: list[int]
40
+
41
+
42
+ class _Lane:
43
+ def __init__(self, lane_idx, *, capacity=2):
44
+ self.lane_idx = lane_idx
45
+ self.model = SimpleNamespace(config=SimpleNamespace(max_batch_size=capacity))
46
+ self.model_args = f"args-{lane_idx}"
47
+ self.mesh_device = f"mesh-{lane_idx}"
48
+ self.cache_path = f"cache-{lane_idx}"
49
+ self.already_warmed_up_prefill = False
50
+ self.calls = []
51
+ self.cleanup_calls = 0
52
+ self.fail_method = None
53
+ self.cleanup_error = None
54
+
55
+ def _call(self, method, kwargs):
56
+ self.calls.append((method, kwargs))
57
+ if self.fail_method == method:
58
+ raise RuntimeError(f"{method} boom {self.lane_idx}")
59
+
60
+ def configure_paged_kv_cache(self, config):
61
+ self._call("configure", {"config": config})
62
+ self.paged_kv_cache_config = config
63
+
64
+ def allocate_kv_cache(self, kv_cache_shape=None, dtype=None, num_layers=None):
65
+ self._call(
66
+ "allocate",
67
+ {
68
+ "kv_cache_shape": kv_cache_shape,
69
+ "dtype": dtype,
70
+ "num_layers": num_layers,
71
+ },
72
+ )
73
+ return f"cache-handle-{self.lane_idx}"
74
+
75
+ def compile_prefill(
76
+ self,
77
+ *,
78
+ tokens,
79
+ page_table,
80
+ prompt_lens=None,
81
+ start_pos=None,
82
+ empty_slots=None,
83
+ kv_cache=None,
84
+ sampling_params=None,
85
+ prompt_tokens=None,
86
+ output_tokens=None,
87
+ slot_remap=None,
88
+ execution=None,
89
+ ):
90
+ self._call(
91
+ "compile_prefill",
92
+ {
93
+ "tokens": tokens,
94
+ "page_table": page_table,
95
+ "prompt_lens": prompt_lens,
96
+ "start_pos": start_pos,
97
+ "empty_slots": empty_slots,
98
+ "kv_cache": kv_cache,
99
+ "sampling_params": sampling_params,
100
+ "prompt_tokens": prompt_tokens,
101
+ "output_tokens": output_tokens,
102
+ "slot_remap": slot_remap,
103
+ "execution": execution,
104
+ },
105
+ )
106
+
107
+ def compile_decode(
108
+ self,
109
+ *,
110
+ tokens,
111
+ start_pos,
112
+ page_table,
113
+ kv_cache=None,
114
+ sampling_params=None,
115
+ prompt_tokens=None,
116
+ output_tokens=None,
117
+ slot_remap=None,
118
+ reset_batch=False,
119
+ execution=None,
120
+ ):
121
+ self._call(
122
+ "compile_decode",
123
+ {
124
+ "tokens": tokens,
125
+ "start_pos": start_pos,
126
+ "page_table": page_table,
127
+ "kv_cache": kv_cache,
128
+ "sampling_params": sampling_params,
129
+ "prompt_tokens": prompt_tokens,
130
+ "output_tokens": output_tokens,
131
+ "slot_remap": slot_remap,
132
+ "reset_batch": reset_batch,
133
+ "execution": execution,
134
+ },
135
+ )
136
+
137
+ def warmup_model_prefill(self, *, kv_cache, can_sample_on_device, enable_trace):
138
+ self._call(
139
+ "warmup_prefill",
140
+ {
141
+ "kv_cache": kv_cache,
142
+ "can_sample_on_device": can_sample_on_device,
143
+ "enable_trace": enable_trace,
144
+ },
145
+ )
146
+ self.already_warmed_up_prefill = True
147
+
148
+ def warmup_model_decode(
149
+ self,
150
+ *,
151
+ kv_cache,
152
+ max_batch_size,
153
+ num_blocks,
154
+ can_sample_on_device,
155
+ enable_trace,
156
+ ):
157
+ self._call(
158
+ "warmup_decode",
159
+ {
160
+ "kv_cache": kv_cache,
161
+ "max_batch_size": max_batch_size,
162
+ "num_blocks": num_blocks,
163
+ "can_sample_on_device": can_sample_on_device,
164
+ "enable_trace": enable_trace,
165
+ },
166
+ )
167
+
168
+ def prefill_forward(
169
+ self,
170
+ tokens,
171
+ page_table,
172
+ *,
173
+ prompt_lens=None,
174
+ start_pos=None,
175
+ empty_slots=None,
176
+ kv_cache=None,
177
+ sampling_params=None,
178
+ prompt_tokens=None,
179
+ output_tokens=None,
180
+ slot_remap=None,
181
+ execution=None,
182
+ ):
183
+ kwargs = {
184
+ "tokens": tokens,
185
+ "page_table": page_table,
186
+ "prompt_lens": prompt_lens,
187
+ "start_pos": start_pos,
188
+ "empty_slots": empty_slots,
189
+ "kv_cache": kv_cache,
190
+ "sampling_params": sampling_params,
191
+ "prompt_tokens": prompt_tokens,
192
+ "output_tokens": output_tokens,
193
+ "slot_remap": slot_remap,
194
+ "execution": execution,
195
+ }
196
+ self._call("prefill", kwargs)
197
+ values = tokens[:, 0] + self.lane_idx * 100
198
+ if sampling_params is not None:
199
+ return values.to(torch.int64), None
200
+ return values.float().view(-1, 1, 1)
201
+
202
+ def decode_forward(
203
+ self,
204
+ tokens,
205
+ start_pos,
206
+ page_table,
207
+ *,
208
+ kv_cache=None,
209
+ sampling_params=None,
210
+ prompt_tokens=None,
211
+ output_tokens=None,
212
+ slot_remap=None,
213
+ reset_batch=False,
214
+ read_from_device=True,
215
+ execution=None,
216
+ ):
217
+ kwargs = {
218
+ "tokens": tokens,
219
+ "start_pos": start_pos,
220
+ "page_table": page_table,
221
+ "kv_cache": kv_cache,
222
+ "sampling_params": sampling_params,
223
+ "prompt_tokens": prompt_tokens,
224
+ "output_tokens": output_tokens,
225
+ "slot_remap": slot_remap,
226
+ "reset_batch": reset_batch,
227
+ "read_from_device": read_from_device,
228
+ "execution": execution,
229
+ }
230
+ self._call("decode", kwargs)
231
+ values = tokens + self.lane_idx * 100
232
+ if not read_from_device:
233
+ return f"raw-{self.lane_idx}", None
234
+ if sampling_params is not None:
235
+ return values.to(torch.int64), None
236
+ return values.float().view(-1, 1, 1), None
237
+
238
+ def can_trace_prefill(self, *, tokens, prompt_lens=None, start_pos=None, empty_slots=None):
239
+ self._call(
240
+ "can_trace_prefill",
241
+ {
242
+ "tokens": tokens,
243
+ "prompt_lens": prompt_lens,
244
+ "start_pos": start_pos,
245
+ "empty_slots": empty_slots,
246
+ },
247
+ )
248
+ return False
249
+
250
+ def read_decode_output(self, tt_out, *, async_read=False):
251
+ self._call("read", {"tt_out": tt_out, "async_read": async_read})
252
+ host = (torch.tensor([self.lane_idx * 2, self.lane_idx * 2 + 1]), None)
253
+ if async_read:
254
+ return host, [f"event-{self.lane_idx}"]
255
+ return host
256
+
257
+ def process_decode_output_host(self, tt_out, *, is_tokens=False):
258
+ self._call("process", {"tt_out": tt_out, "is_tokens": is_tokens})
259
+ if is_tokens:
260
+ return torch.tensor([self.lane_idx * 2, self.lane_idx * 2 + 1], dtype=torch.int64), None
261
+ return torch.full((2, 1, 3), float(self.lane_idx)), None
262
+
263
+ def cleanup(self):
264
+ self.cleanup_calls += 1
265
+ if self.cleanup_error is not None:
266
+ raise self.cleanup_error
267
+
268
+
269
+ def _sampling():
270
+ return _SamplingParams(
271
+ temperature=[0.1, 0.2, 0.3, 0.4],
272
+ top_k=torch.tensor([1, 2, 3, 4]),
273
+ top_p=0.9,
274
+ )
275
+
276
+
277
+ _POSITIONAL = inspect.Parameter.POSITIONAL_OR_KEYWORD
278
+ _KEYWORD_ONLY = inspect.Parameter.KEYWORD_ONLY
279
+ _REQUIRED = inspect.Parameter.empty
280
+
281
+ _PUBLIC_SIGNATURES = {
282
+ "allocate_kv_cache": (
283
+ ("kv_cache_shape", _POSITIONAL, None),
284
+ ("dtype", _POSITIONAL, None),
285
+ ("num_layers", _POSITIONAL, None),
286
+ ),
287
+ "compile_prefill": (
288
+ ("tokens", _POSITIONAL, _REQUIRED),
289
+ ("page_table", _POSITIONAL, _REQUIRED),
290
+ ("prompt_lens", _KEYWORD_ONLY, None),
291
+ ("start_pos", _KEYWORD_ONLY, None),
292
+ ("empty_slots", _KEYWORD_ONLY, None),
293
+ ("kv_cache", _KEYWORD_ONLY, None),
294
+ ("sampling_params", _KEYWORD_ONLY, None),
295
+ ("prompt_tokens", _KEYWORD_ONLY, None),
296
+ ("output_tokens", _KEYWORD_ONLY, None),
297
+ ("slot_remap", _KEYWORD_ONLY, None),
298
+ ("execution", _KEYWORD_ONLY, None),
299
+ ),
300
+ "compile_decode": (
301
+ ("tokens", _POSITIONAL, _REQUIRED),
302
+ ("start_pos", _POSITIONAL, _REQUIRED),
303
+ ("page_table", _POSITIONAL, _REQUIRED),
304
+ ("kv_cache", _KEYWORD_ONLY, None),
305
+ ("sampling_params", _KEYWORD_ONLY, None),
306
+ ("prompt_tokens", _KEYWORD_ONLY, None),
307
+ ("output_tokens", _KEYWORD_ONLY, None),
308
+ ("slot_remap", _KEYWORD_ONLY, None),
309
+ ("reset_batch", _KEYWORD_ONLY, False),
310
+ ("execution", _KEYWORD_ONLY, None),
311
+ ),
312
+ "warmup_model_prefill": (
313
+ ("kv_cache", _KEYWORD_ONLY, _REQUIRED),
314
+ ("can_sample_on_device", _KEYWORD_ONLY, _REQUIRED),
315
+ ("enable_trace", _KEYWORD_ONLY, _REQUIRED),
316
+ ),
317
+ "warmup_model_decode": (
318
+ ("kv_cache", _KEYWORD_ONLY, _REQUIRED),
319
+ ("max_batch_size", _KEYWORD_ONLY, _REQUIRED),
320
+ ("num_blocks", _KEYWORD_ONLY, _REQUIRED),
321
+ ("can_sample_on_device", _KEYWORD_ONLY, _REQUIRED),
322
+ ("enable_trace", _KEYWORD_ONLY, _REQUIRED),
323
+ ),
324
+ "prefill_forward": (
325
+ ("tokens", _POSITIONAL, _REQUIRED),
326
+ ("page_table", _POSITIONAL, _REQUIRED),
327
+ ("prompt_lens", _KEYWORD_ONLY, None),
328
+ ("start_pos", _KEYWORD_ONLY, None),
329
+ ("empty_slots", _KEYWORD_ONLY, None),
330
+ ("kv_cache", _KEYWORD_ONLY, None),
331
+ ("sampling_params", _KEYWORD_ONLY, None),
332
+ ("prompt_tokens", _KEYWORD_ONLY, None),
333
+ ("output_tokens", _KEYWORD_ONLY, None),
334
+ ("slot_remap", _KEYWORD_ONLY, None),
335
+ ("execution", _KEYWORD_ONLY, None),
336
+ ),
337
+ "can_trace_prefill": (
338
+ ("tokens", _KEYWORD_ONLY, _REQUIRED),
339
+ ("prompt_lens", _KEYWORD_ONLY, None),
340
+ ("start_pos", _KEYWORD_ONLY, None),
341
+ ("empty_slots", _KEYWORD_ONLY, None),
342
+ ),
343
+ "decode_forward": (
344
+ ("tokens", _POSITIONAL, _REQUIRED),
345
+ ("start_pos", _POSITIONAL, _REQUIRED),
346
+ ("page_table", _POSITIONAL, _REQUIRED),
347
+ ("kv_cache", _KEYWORD_ONLY, None),
348
+ ("sampling_params", _KEYWORD_ONLY, None),
349
+ ("prompt_tokens", _KEYWORD_ONLY, None),
350
+ ("output_tokens", _KEYWORD_ONLY, None),
351
+ ("slot_remap", _KEYWORD_ONLY, None),
352
+ ("reset_batch", _KEYWORD_ONLY, False),
353
+ ("read_from_device", _KEYWORD_ONLY, True),
354
+ ("execution", _KEYWORD_ONLY, None),
355
+ ),
356
+ "read_decode_output": (
357
+ ("tt_out", _POSITIONAL, _REQUIRED),
358
+ ("async_read", _KEYWORD_ONLY, False),
359
+ ),
360
+ "process_decode_output_host": (
361
+ ("tt_out", _POSITIONAL, _REQUIRED),
362
+ ("is_tokens", _KEYWORD_ONLY, False),
363
+ ),
364
+ }
365
+
366
+
367
+ @pytest.mark.parametrize(("method_name", "expected"), _PUBLIC_SIGNATURES.items())
368
+ def test_public_api_signatures_are_exact(method_name, expected):
369
+ signature = inspect.signature(getattr(LaneGroupExecutor, method_name))
370
+ parameters = tuple(signature.parameters.values())[1:]
371
+
372
+ assert tuple((parameter.name, parameter.kind, parameter.default) for parameter in parameters) == expected
373
+ assert all(parameter.annotation is not inspect.Parameter.empty for parameter in parameters)
374
+ assert signature.return_annotation is not inspect.Signature.empty
375
+
376
+
377
+ def test_prefill_routes_global_slots_and_restores_source_row_order():
378
+ lanes = [_Lane(0), _Lane(1)]
379
+ group = LaneGroupExecutor(lanes)
380
+ tokens = torch.tensor([[10], [11], [12]])
381
+ page_table = torch.tensor([[0], [1], [2]], dtype=torch.int32)
382
+
383
+ output = group.prefill_forward(
384
+ tokens,
385
+ page_table,
386
+ empty_slots=[3, 0, 2],
387
+ prompt_lens=torch.tensor([1, 1, 1]),
388
+ sampling_params=_sampling(),
389
+ prompt_tokens=torch.tensor([[100], [101], [102], [103]]),
390
+ output_tokens=[[200], [201], [202], [203]],
391
+ slot_remap=torch.tensor([0, 1, 2, 3]),
392
+ kv_cache=["kv-0", "kv-1"],
393
+ )
394
+
395
+ assert isinstance(output, tuple)
396
+ assert output[0].tolist() == [110, 11, 112]
397
+ assert output[0].dtype == torch.int64
398
+ assert output[1] is None
399
+
400
+ lane0_kwargs = next(kwargs for method, kwargs in lanes[0].calls if method == "prefill")
401
+ lane1_kwargs = next(kwargs for method, kwargs in lanes[1].calls if method == "prefill")
402
+ assert lane0_kwargs["empty_slots"] == [0]
403
+ assert lane1_kwargs["empty_slots"] == [1, 0]
404
+ assert lane0_kwargs["tokens"].flatten().tolist() == [11]
405
+ assert lane1_kwargs["tokens"].flatten().tolist() == [10, 12]
406
+ assert lane0_kwargs["sampling_params"].temperature == [0.2]
407
+ assert lane1_kwargs["sampling_params"].temperature == [0.1, 0.3]
408
+ assert lane0_kwargs["sampling_params"].top_k.tolist() == [2]
409
+ assert lane1_kwargs["sampling_params"].top_k.tolist() == [1, 3]
410
+ assert lane0_kwargs["sampling_params"].top_p == 0.9
411
+ assert lane0_kwargs["prompt_tokens"].flatten().tolist() == [100, -1]
412
+ assert lane1_kwargs["prompt_tokens"].flatten().tolist() == [102, 103]
413
+ assert lane0_kwargs["output_tokens"] == [[200], [-1]]
414
+ assert lane1_kwargs["output_tokens"] == [[202], [203]]
415
+ assert lane0_kwargs["slot_remap"].tolist() == [0, 1]
416
+ assert lane1_kwargs["slot_remap"].tolist() == [0, 1]
417
+ assert lane0_kwargs["kv_cache"] == "kv-0"
418
+ assert lane1_kwargs["kv_cache"] == "kv-1"
419
+
420
+
421
+ def test_prefill_log_probs_restore_source_rows_and_fill_nonrequesting_lane():
422
+ lanes = [_Lane(0), _Lane(1)]
423
+
424
+ def prefill_with_log_probs(
425
+ tokens,
426
+ page_table,
427
+ *,
428
+ prompt_lens=None,
429
+ start_pos=None,
430
+ empty_slots=None,
431
+ kv_cache=None,
432
+ sampling_params=None,
433
+ execution=None,
434
+ ):
435
+ assert tokens.shape[0] == 2
436
+ return torch.zeros(2, dtype=torch.int64), torch.tensor([0.25, 0.75])
437
+
438
+ lanes[1].prefill_forward = prefill_with_log_probs
439
+ group = LaneGroupExecutor(lanes)
440
+
441
+ output, log_probs = group.prefill_forward(
442
+ tokens=torch.tensor([[10], [11], [12]]),
443
+ page_table=torch.tensor([[0], [1], [2]], dtype=torch.int32),
444
+ empty_slots=[3, 0, 2],
445
+ sampling_params=_sampling(),
446
+ )
447
+
448
+ assert output.tolist() == [0, 11, 0]
449
+ assert isinstance(log_probs, torch.Tensor)
450
+ assert log_probs.tolist() == [0.25, 1.0, 0.75]
451
+
452
+
453
+ def test_decode_slices_fixed_lane_capacity_and_normalizes_sampled_tokens():
454
+ lanes = [_Lane(0), _Lane(1)]
455
+ group = LaneGroupExecutor(lanes)
456
+
457
+ output = group.decode_forward(
458
+ torch.tensor([10, 11, 12, 13]),
459
+ torch.tensor([1, 2, 3, 4]),
460
+ torch.arange(4, dtype=torch.int32).view(4, 1),
461
+ sampling_params=_sampling(),
462
+ prompt_tokens=torch.tensor([[100], [101], [102], [103]]),
463
+ output_tokens=[[200], [201], [202], [203]],
464
+ slot_remap=torch.tensor([1, 0, 3, 2]),
465
+ kv_cache=["kv-0", "kv-1"],
466
+ reset_batch=True,
467
+ )
468
+
469
+ assert output[0].tolist() == [10, 11, 112, 113]
470
+ assert output[0].dtype == torch.int64
471
+ assert output[1] is None
472
+ lane0_kwargs = next(kwargs for method, kwargs in lanes[0].calls if method == "decode")
473
+ lane1_kwargs = next(kwargs for method, kwargs in lanes[1].calls if method == "decode")
474
+ assert lane0_kwargs["tokens"].tolist() == [10, 11]
475
+ assert lane1_kwargs["tokens"].tolist() == [12, 13]
476
+ assert lane0_kwargs["sampling_params"].temperature == [0.1, 0.2]
477
+ assert lane1_kwargs["sampling_params"].temperature == [0.3, 0.4]
478
+ assert lane0_kwargs["prompt_tokens"].flatten().tolist() == [100, 101]
479
+ assert lane1_kwargs["prompt_tokens"].flatten().tolist() == [102, 103]
480
+ assert lane0_kwargs["output_tokens"] == [[200], [201]]
481
+ assert lane1_kwargs["output_tokens"] == [[202], [203]]
482
+ assert lane0_kwargs["slot_remap"].tolist() == [1, 0]
483
+ assert lane1_kwargs["slot_remap"].tolist() == [1, 0]
484
+ assert lane0_kwargs["reset_batch"] is lane1_kwargs["reset_batch"] is True
485
+
486
+
487
+ @pytest.mark.parametrize(
488
+ (
489
+ "layout",
490
+ "prefill_tokens",
491
+ "empty_slots",
492
+ "prefill_sampling",
493
+ "decode_tokens",
494
+ "decode_positions",
495
+ "decode_sampling",
496
+ "slot_remap",
497
+ "expected_prefill_tokens",
498
+ "expected_decode_tokens",
499
+ "expected_lane_slots",
500
+ "expected_lane_seeds",
501
+ ),
502
+ [
503
+ pytest.param(
504
+ "front_packed_gathered",
505
+ torch.tensor([[10], [11], [12], [13]]),
506
+ [0, 1, 2, 3],
507
+ _StatefulSamplingParams(
508
+ temperature=[0.0, 0.8, 1.0, 1.2],
509
+ top_k=torch.tensor([1, 5, 7, 9]),
510
+ top_p=[0.0, 0.7, 0.8, 0.9],
511
+ presence_penalty=[0.0, 0.1, 0.2, 0.3],
512
+ frequency_penalty=[0.0, 0.4, 0.5, 0.6],
513
+ repetition_penalty=[1.0, 1.1, 1.2, 1.3],
514
+ seed=[100, 101, 102, 103],
515
+ enable_log_probs=[False, True, False, True],
516
+ num_logprobs=[0, 0, 0, 0],
517
+ ),
518
+ torch.tensor([10, 11, 12, 13]),
519
+ torch.tensor([1, 2, 3, 4]),
520
+ _StatefulSamplingParams(
521
+ temperature=[0.0, 0.8, 1.0, 1.2],
522
+ top_k=torch.tensor([1, 5, 7, 9]),
523
+ top_p=[0.0, 0.7, 0.8, 0.9],
524
+ presence_penalty=[0.0, 0.1, 0.2, 0.3],
525
+ frequency_penalty=[0.0, 0.4, 0.5, 0.6],
526
+ repetition_penalty=[1.0, 1.1, 1.2, 1.3],
527
+ seed=[100, 101, 102, 103],
528
+ enable_log_probs=[False, True, False, True],
529
+ num_logprobs=[0, 0, 0, 0],
530
+ ),
531
+ torch.tensor([1, 0, 3, 2]),
532
+ [10, 11, 112, 113],
533
+ [10, 11, 112, 113],
534
+ ([0, 1], [0, 1]),
535
+ ([100, 101], [102, 103]),
536
+ id="front_packed_gathered",
537
+ ),
538
+ pytest.param(
539
+ "stable_gap_lane",
540
+ torch.tensor([[11], [13]]),
541
+ [1, 3],
542
+ _StatefulSamplingParams(
543
+ temperature=[0.8, 1.2],
544
+ top_k=torch.tensor([5, 9]),
545
+ top_p=[0.7, 0.9],
546
+ presence_penalty=[0.1, 0.3],
547
+ frequency_penalty=[0.4, 0.6],
548
+ repetition_penalty=[1.1, 1.3],
549
+ seed=[101, 103],
550
+ enable_log_probs=[True, True],
551
+ num_logprobs=[0, 0],
552
+ ),
553
+ torch.tensor([0, 11, 0, 13]),
554
+ torch.tensor([-1, 2, -1, 4]),
555
+ _StatefulSamplingParams(
556
+ temperature=[0.0, 0.8, 0.0, 1.2],
557
+ top_k=torch.tensor([1, 5, 1, 9]),
558
+ top_p=[0.0, 0.7, 0.0, 0.9],
559
+ presence_penalty=[0.0, 0.1, 0.0, 0.3],
560
+ frequency_penalty=[0.0, 0.4, 0.0, 0.6],
561
+ repetition_penalty=[1.0, 1.1, 1.0, 1.3],
562
+ seed=[None, 101, None, 103],
563
+ enable_log_probs=[False, True, False, True],
564
+ num_logprobs=[0, 0, 0, 0],
565
+ ),
566
+ torch.tensor([0, 1, 2, 3]),
567
+ [11, 113],
568
+ [0, 11, 100, 113],
569
+ ([1], [1]),
570
+ ([None, 101], [None, 103]),
571
+ id="stable_gap_lane",
572
+ ),
573
+ ],
574
+ )
575
+ def test_sampling_state_boundary_contract_for_gathered_and_lane_layouts(
576
+ layout,
577
+ prefill_tokens,
578
+ empty_slots,
579
+ prefill_sampling,
580
+ decode_tokens,
581
+ decode_positions,
582
+ decode_sampling,
583
+ slot_remap,
584
+ expected_prefill_tokens,
585
+ expected_decode_tokens,
586
+ expected_lane_slots,
587
+ expected_lane_seeds,
588
+ ):
589
+ """Validate layouts already constructed at the tt-metal target boundary.
590
+
591
+ ``front_packed_gathered`` models the dense, rank-segmented payload supplied
592
+ after gathered-DP packing. ``stable_gap_lane`` models the fixed stable-slot
593
+ grid supplied by lane-DP. This deliberately does not claim to test the vLLM
594
+ code that constructs or gathers either payload.
595
+ """
596
+
597
+ lanes = [_Lane(0), _Lane(1)]
598
+ group = LaneGroupExecutor(lanes)
599
+ prompt_tokens = torch.tensor([[200, -1], [210, 211], [220, -1], [230, 231]])
600
+ output_tokens = [[300, -1], [310, 311], [320, -1], [330, 331]]
601
+ execution = [f"trace-{layout}-0", f"trace-{layout}-1"]
602
+
603
+ prefill_result = group.prefill_forward(
604
+ prefill_tokens,
605
+ torch.arange(len(prefill_tokens), dtype=torch.int32).view(-1, 1),
606
+ empty_slots=empty_slots,
607
+ sampling_params=prefill_sampling,
608
+ prompt_tokens=prompt_tokens,
609
+ output_tokens=output_tokens,
610
+ slot_remap=slot_remap,
611
+ execution=execution,
612
+ )
613
+ assert prefill_result[0].tolist() == expected_prefill_tokens
614
+
615
+ decode_result = group.decode_forward(
616
+ decode_tokens,
617
+ decode_positions,
618
+ torch.arange(4, dtype=torch.int32).view(4, 1),
619
+ sampling_params=decode_sampling,
620
+ prompt_tokens=prompt_tokens,
621
+ output_tokens=output_tokens,
622
+ slot_remap=slot_remap,
623
+ reset_batch=True,
624
+ execution=execution,
625
+ )
626
+ assert decode_result[0].tolist() == expected_decode_tokens
627
+
628
+ if layout == "front_packed_gathered":
629
+ prefill_rows = ([0, 1], [2, 3])
630
+ prefill_prompt = (
631
+ prompt_tokens[0:2].tolist(),
632
+ prompt_tokens[2:4].tolist(),
633
+ )
634
+ prefill_output = (output_tokens[0:2], output_tokens[2:4])
635
+ lane_remap = ([1, 0], [1, 0])
636
+ else:
637
+ prefill_rows = ([0], [1])
638
+ prefill_prompt = (
639
+ [[-1, -1], prompt_tokens[1].tolist()],
640
+ [[-1, -1], prompt_tokens[3].tolist()],
641
+ )
642
+ prefill_output = (
643
+ [[-1, -1], output_tokens[1]],
644
+ [[-1, -1], output_tokens[3]],
645
+ )
646
+ lane_remap = ([0, 1], [0, 1])
647
+
648
+ for lane_idx, lane in enumerate(lanes):
649
+ prefill = next(kwargs for method, kwargs in lane.calls if method == "prefill")
650
+ decode = next(kwargs for method, kwargs in lane.calls if method == "decode")
651
+ source_rows = prefill_rows[lane_idx]
652
+ prefill_params = prefill["sampling_params"]
653
+ decode_params = decode["sampling_params"]
654
+ assert prefill["empty_slots"] == expected_lane_slots[lane_idx]
655
+ assert prefill["execution"] == execution[lane_idx]
656
+ assert prefill_params.temperature == [prefill_sampling.temperature[row] for row in source_rows]
657
+ assert prefill_params.top_k.tolist() == [int(prefill_sampling.top_k[row]) for row in source_rows]
658
+ assert prefill_params.top_p == [prefill_sampling.top_p[row] for row in source_rows]
659
+ assert prefill_params.presence_penalty == [prefill_sampling.presence_penalty[row] for row in source_rows]
660
+ assert prefill_params.frequency_penalty == [prefill_sampling.frequency_penalty[row] for row in source_rows]
661
+ assert prefill_params.repetition_penalty == [prefill_sampling.repetition_penalty[row] for row in source_rows]
662
+ assert prefill_params.seed == [prefill_sampling.seed[row] for row in source_rows]
663
+ assert prefill_params.enable_log_probs == [prefill_sampling.enable_log_probs[row] for row in source_rows]
664
+ assert prefill_params.num_logprobs == [prefill_sampling.num_logprobs[row] for row in source_rows]
665
+ assert prefill["prompt_tokens"].tolist() == prefill_prompt[lane_idx]
666
+ assert prefill["output_tokens"] == prefill_output[lane_idx]
667
+ assert prefill["slot_remap"].tolist() == lane_remap[lane_idx]
668
+ assert decode["execution"] == execution[lane_idx]
669
+ lane_slice = slice(lane_idx * 2, lane_idx * 2 + 2)
670
+ assert decode_params.temperature == decode_sampling.temperature[lane_slice]
671
+ assert decode_params.top_k.tolist() == decode_sampling.top_k[lane_slice].tolist()
672
+ assert decode_params.top_p == decode_sampling.top_p[lane_slice]
673
+ assert decode_params.presence_penalty == decode_sampling.presence_penalty[lane_slice]
674
+ assert decode_params.frequency_penalty == decode_sampling.frequency_penalty[lane_slice]
675
+ assert decode_params.repetition_penalty == decode_sampling.repetition_penalty[lane_slice]
676
+ assert decode_params.seed == expected_lane_seeds[lane_idx]
677
+ assert decode_params.enable_log_probs == decode_sampling.enable_log_probs[lane_slice]
678
+ assert decode_params.num_logprobs == decode_sampling.num_logprobs[lane_slice]
679
+ assert decode["prompt_tokens"].tolist() == prompt_tokens[lane_idx * 2 : lane_idx * 2 + 2].tolist()
680
+ assert decode["output_tokens"] == output_tokens[lane_idx * 2 : lane_idx * 2 + 2]
681
+ assert decode["slot_remap"].tolist() == lane_remap[lane_idx]
682
+ assert decode["reset_batch"] is True
683
+
684
+ raw = group.decode_forward(
685
+ decode_tokens,
686
+ decode_positions,
687
+ torch.arange(4, dtype=torch.int32).view(4, 1),
688
+ sampling_params=decode_sampling,
689
+ read_from_device=False,
690
+ )
691
+ host_outputs, events = group.read_decode_output(raw, async_read=True)
692
+ completed = group.process_decode_output_host(host_outputs, is_tokens=True)
693
+ assert events == ["event-0", "event-1"]
694
+ assert completed[0].tolist() == [0, 1, 2, 3]
695
+ assert completed[1] is None
696
+
697
+
698
+ def test_decode_broadcasts_singleton_sampling_values_to_later_dp_lanes():
699
+ lanes = [_Lane(0, capacity=1), _Lane(1, capacity=1)]
700
+ group = LaneGroupExecutor(lanes)
701
+ sampling_params = _SingletonSamplingParams(
702
+ temperature=[0.7],
703
+ top_k=torch.tensor([8]),
704
+ top_p=(0.4,),
705
+ )
706
+
707
+ output = group.decode_forward(
708
+ tokens=torch.tensor([10, 11]),
709
+ start_pos=torch.tensor([1, 2]),
710
+ page_table=torch.arange(2, dtype=torch.int32).view(2, 1),
711
+ sampling_params=sampling_params,
712
+ )
713
+
714
+ assert output[0].tolist() == [10, 111]
715
+ lane0_kwargs = next(kwargs for method, kwargs in lanes[0].calls if method == "decode")
716
+ lane1_kwargs = next(kwargs for method, kwargs in lanes[1].calls if method == "decode")
717
+ for lane_kwargs in (lane0_kwargs, lane1_kwargs):
718
+ sliced = lane_kwargs["sampling_params"]
719
+ assert sliced.temperature == [0.7]
720
+ assert sliced.top_k.tolist() == [8]
721
+ assert sliced.top_p == (0.4,)
722
+
723
+
724
+ def test_decode_without_readback_returns_lane_local_outputs():
725
+ group = LaneGroupExecutor([_Lane(0), _Lane(1)])
726
+
727
+ output = group.decode_forward(
728
+ tokens=torch.tensor([0, 1, 2, 3]),
729
+ start_pos=torch.tensor([0, 0, 0, 0]),
730
+ page_table=torch.zeros((4, 1), dtype=torch.int32),
731
+ read_from_device=False,
732
+ )
733
+
734
+ assert output == [("raw-0", None), ("raw-1", None)]
735
+
736
+
737
+ def test_async_read_and_host_processing_run_per_lane_and_preserve_lane_order():
738
+ group = LaneGroupExecutor([_Lane(0), _Lane(1)])
739
+
740
+ host_outputs, events = group.read_decode_output(
741
+ tt_out=[("raw-0", None), ("raw-1", None)],
742
+ async_read=True,
743
+ )
744
+ processed = group.process_decode_output_host(tt_out=host_outputs, is_tokens=True)
745
+
746
+ assert events == ["event-0", "event-1"]
747
+ assert processed[0].tolist() == [0, 1, 2, 3]
748
+ assert processed[0].dtype == torch.int64
749
+ assert processed[1] is None
750
+
751
+
752
+ def test_warmup_replicates_lane_local_case_and_cache_to_every_lane():
753
+ lanes = [_Lane(0), _Lane(1)]
754
+ group = LaneGroupExecutor(lanes)
755
+
756
+ group.warmup_model_prefill(
757
+ kv_cache=["kv-0", "kv-1"],
758
+ can_sample_on_device=True,
759
+ enable_trace=False,
760
+ )
761
+ group.warmup_model_decode(
762
+ kv_cache=["kv-0", "kv-1"],
763
+ enable_trace=False,
764
+ max_batch_size=4,
765
+ num_blocks=8,
766
+ can_sample_on_device=True,
767
+ )
768
+
769
+ for lane_idx, lane in enumerate(lanes):
770
+ prefill = next(kwargs for method, kwargs in lane.calls if method == "warmup_prefill")
771
+ decode = next(kwargs for method, kwargs in lane.calls if method == "warmup_decode")
772
+ assert prefill["kv_cache"] == f"kv-{lane_idx}"
773
+ assert decode["kv_cache"] == f"kv-{lane_idx}"
774
+ assert decode["max_batch_size"] == 2
775
+ assert group.already_warmed_up_prefill
776
+
777
+
778
+ def test_traced_warmup_without_coordinators_fails_before_invoking_a_lane(expect_error):
779
+ lanes = [_Lane(0), _Lane(1)]
780
+ group = LaneGroupExecutor(lanes)
781
+
782
+ with expect_error(RuntimeError, "Every DP lane must expose the trace activation barrier"):
783
+ group.warmup_model_decode(
784
+ kv_cache=["kv-0", "kv-1"],
785
+ enable_trace=True,
786
+ max_batch_size=4,
787
+ num_blocks=8,
788
+ can_sample_on_device=False,
789
+ )
790
+
791
+ assert [lane.calls for lane in lanes] == [[], []]
792
+
793
+
794
+ class _DeferredCapture:
795
+ def __init__(self, lane_idx, events, *, ready=True):
796
+ self.lane_idx = lane_idx
797
+ self.events = events
798
+ self.ready = ready
799
+ self.capture_pending = False
800
+ self.trace_activated = False
801
+ self.capture_calls = 0
802
+
803
+ @contextmanager
804
+ def defer_capture(self):
805
+ self.events.append(("enter", self.lane_idx))
806
+ try:
807
+ yield self
808
+ finally:
809
+ self.capture_pending = False
810
+ self.events.append(("exit", self.lane_idx))
811
+
812
+ def register(self):
813
+ self.events.append(("register", self.lane_idx))
814
+ self.capture_pending = self.ready
815
+
816
+ def activate_pending_capture(self):
817
+ assert self.capture_pending
818
+ self.events.append(("capture", self.lane_idx))
819
+ self.capture_calls += 1
820
+ self.trace_activated = True
821
+ self.capture_pending = False
822
+
823
+
824
+ class _BarrierLane(_Lane):
825
+ def __init__(self, lane_idx, events, *, ready=True, fail=False):
826
+ super().__init__(lane_idx)
827
+ self.warmup = _DeferredCapture(lane_idx, events, ready=ready)
828
+ self.events = events
829
+ self.fail = fail
830
+
831
+ def warmup_model_decode(self, **kwargs):
832
+ self.events.append(("warmup", self.lane_idx))
833
+ if kwargs["enable_trace"]:
834
+ self.warmup.register()
835
+ if self.fail:
836
+ raise RuntimeError(f"warmup boom {self.lane_idx}")
837
+ return super().warmup_model_decode(**kwargs)
838
+
839
+
840
+ def test_traced_warmup_registers_every_lane_before_any_capture():
841
+ events = []
842
+ lanes = [_BarrierLane(0, events), _BarrierLane(1, events)]
843
+ group = LaneGroupExecutor(lanes)
844
+
845
+ group.warmup_model_decode(
846
+ kv_cache=["kv-0", "kv-1"],
847
+ enable_trace=True,
848
+ max_batch_size=4,
849
+ num_blocks=8,
850
+ can_sample_on_device=False,
851
+ )
852
+
853
+ assert events == [
854
+ ("enter", 0),
855
+ ("enter", 1),
856
+ ("warmup", 0),
857
+ ("register", 0),
858
+ ("warmup", 1),
859
+ ("register", 1),
860
+ ("capture", 0),
861
+ ("capture", 1),
862
+ ("exit", 1),
863
+ ("exit", 0),
864
+ ]
865
+
866
+
867
+ @pytest.mark.parametrize(
868
+ ("mode", "message"),
869
+ [
870
+ ("exception", "warmup boom 1"),
871
+ ("asymmetry", "mixed trace activation readiness"),
872
+ ],
873
+ )
874
+ def test_traced_warmup_failure_or_asymmetry_captures_no_lane(mode, message, expect_error):
875
+ events = []
876
+ group_lanes = [
877
+ _BarrierLane(0, events, ready=True),
878
+ _BarrierLane(1, events, ready=mode != "asymmetry", fail=mode == "exception"),
879
+ ]
880
+ group = LaneGroupExecutor(group_lanes)
881
+
882
+ with expect_error(RuntimeError, message):
883
+ group.warmup_model_decode(
884
+ kv_cache=["kv-0", "kv-1"],
885
+ enable_trace=True,
886
+ max_batch_size=4,
887
+ num_blocks=8,
888
+ can_sample_on_device=False,
889
+ )
890
+
891
+ assert [lane.warmup.capture_calls for lane in group_lanes] == [0, 0]
892
+ assert not any(lane.warmup.capture_pending for lane in group_lanes)
893
+
894
+
895
+ def test_eager_warmup_does_not_enter_trace_activation_deferral():
896
+ events = []
897
+ lanes = [_BarrierLane(0, events), _BarrierLane(1, events)]
898
+ group = LaneGroupExecutor(lanes)
899
+
900
+ group.warmup_model_decode(
901
+ kv_cache=["kv-0", "kv-1"],
902
+ enable_trace=False,
903
+ max_batch_size=4,
904
+ num_blocks=8,
905
+ can_sample_on_device=False,
906
+ )
907
+
908
+ assert events == [
909
+ ("warmup", 0),
910
+ ("warmup", 1),
911
+ ]
912
+ assert [lane.warmup.capture_calls for lane in lanes] == [0, 0]
913
+
914
+
915
+ def test_compile_methods_slice_requests_to_lane_executors():
916
+ lanes = [_Lane(0), _Lane(1)]
917
+ group = LaneGroupExecutor(lanes)
918
+
919
+ group.compile_prefill(
920
+ torch.tensor([[10], [11], [12]]),
921
+ torch.tensor([[0], [1], [2]], dtype=torch.int32),
922
+ empty_slots=[3, 0, 2],
923
+ kv_cache=["kv-0", "kv-1"],
924
+ )
925
+ group.compile_decode(
926
+ torch.tensor([10, 11, 12, 13]),
927
+ torch.tensor([1, 2, 3, 4]),
928
+ torch.arange(4, dtype=torch.int32).view(4, 1),
929
+ kv_cache=["kv-0", "kv-1"],
930
+ )
931
+
932
+ lane0_prefill = next(kwargs for method, kwargs in lanes[0].calls if method == "compile_prefill")
933
+ lane1_prefill = next(kwargs for method, kwargs in lanes[1].calls if method == "compile_prefill")
934
+ lane0_decode = next(kwargs for method, kwargs in lanes[0].calls if method == "compile_decode")
935
+ lane1_decode = next(kwargs for method, kwargs in lanes[1].calls if method == "compile_decode")
936
+ assert lane0_prefill["tokens"].flatten().tolist() == [11]
937
+ assert lane1_prefill["tokens"].flatten().tolist() == [10, 12]
938
+ assert lane0_decode["tokens"].tolist() == [10, 11]
939
+ assert lane1_decode["tokens"].tolist() == [12, 13]
940
+
941
+
942
+ def test_concrete_execution_targets_preflight_every_traced_lane_before_any_execution(expect_error):
943
+ class DispatchLane(_Lane):
944
+ def __init__(self, lane_idx, *, traceable):
945
+ super().__init__(lane_idx)
946
+ self.traceable = traceable
947
+ self.eager_execution = object()
948
+ self.traced_prefill_execution = object()
949
+ self.traced_decode_execution = object()
950
+
951
+ def can_trace_prefill(self, *, tokens, prompt_lens=None, start_pos=None):
952
+ self._call(
953
+ "can_trace_prefill",
954
+ {
955
+ "tokens": tokens,
956
+ "prompt_lens": prompt_lens,
957
+ "start_pos": start_pos,
958
+ },
959
+ )
960
+ return self.traceable
961
+
962
+ def prefill_forward(
963
+ self,
964
+ tokens,
965
+ page_table,
966
+ *,
967
+ prompt_lens=None,
968
+ start_pos=None,
969
+ empty_slots=None,
970
+ kv_cache=None,
971
+ sampling_params=None,
972
+ execution=None,
973
+ ):
974
+ kwargs = {
975
+ "tokens": tokens,
976
+ "page_table": page_table,
977
+ "prompt_lens": prompt_lens,
978
+ "start_pos": start_pos,
979
+ "empty_slots": empty_slots,
980
+ "kv_cache": kv_cache,
981
+ "sampling_params": sampling_params,
982
+ "execution": execution,
983
+ }
984
+ self._call("prefill", kwargs)
985
+ return tokens[:, :1].float().view(-1, 1, 1)
986
+
987
+ lanes = [DispatchLane(0, traceable=True), DispatchLane(1, traceable=False)]
988
+ group = LaneGroupExecutor(lanes)
989
+ tokens = torch.tensor([[10], [11]])
990
+ page_table = torch.tensor([[0], [1]], dtype=torch.int32)
991
+ prompt_lens = torch.tensor([7, 9])
992
+ start_pos = torch.tensor([3, 4])
993
+ empty_slots = [0, 2]
994
+
995
+ assert not group.can_trace_prefill(
996
+ tokens=tokens,
997
+ prompt_lens=prompt_lens,
998
+ start_pos=start_pos,
999
+ empty_slots=empty_slots,
1000
+ )
1001
+ kwargs = {
1002
+ "tokens": tokens,
1003
+ "page_table": page_table,
1004
+ "prompt_lens": prompt_lens,
1005
+ "start_pos": start_pos,
1006
+ "empty_slots": empty_slots,
1007
+ }
1008
+ assert group.prefill_forward(execution=group.eager_execution, **kwargs).flatten().tolist() == [10.0, 11.0]
1009
+ for lane_idx, lane in enumerate(lanes):
1010
+ eager_call = next(call for call in lane.calls if call[0] == "prefill")
1011
+ assert eager_call[1]["execution"] is group.eager_execution[lane_idx]
1012
+ lane.calls.clear()
1013
+
1014
+ with expect_error(RuntimeError, "no DP lane executed"):
1015
+ group.prefill_forward(execution=group.traced_prefill_execution, **kwargs)
1016
+
1017
+ assert [[method for method, _ in lane.calls] for lane in lanes] == [
1018
+ ["can_trace_prefill"],
1019
+ ["can_trace_prefill"],
1020
+ ]
1021
+ for lane_idx, lane in enumerate(lanes):
1022
+ trace_kwargs = lane.calls[0][1]
1023
+ assert set(trace_kwargs) == {"tokens", "prompt_lens", "start_pos"}
1024
+ assert trace_kwargs["tokens"].flatten().tolist() == [10 + lane_idx]
1025
+ assert trace_kwargs["prompt_lens"].tolist() == [7 + 2 * lane_idx]
1026
+ assert trace_kwargs["start_pos"].tolist() == [3 + lane_idx]
1027
+
1028
+
1029
+ def test_dp1_trace_classification_slices_only_the_exact_lane_subset():
1030
+ class TraceLane(_Lane):
1031
+ def can_trace_prefill(self, *, tokens, prompt_lens=None, start_pos=None):
1032
+ self._call(
1033
+ "can_trace_prefill",
1034
+ {
1035
+ "tokens": tokens,
1036
+ "prompt_lens": prompt_lens,
1037
+ "start_pos": start_pos,
1038
+ },
1039
+ )
1040
+ return True
1041
+
1042
+ lane = TraceLane(0)
1043
+ group = LaneGroupExecutor([lane])
1044
+
1045
+ assert group.can_trace_prefill(
1046
+ tokens=torch.tensor([[10], [11]]),
1047
+ prompt_lens=torch.tensor([7, 9]),
1048
+ start_pos=torch.tensor([3, 4]),
1049
+ empty_slots=[1, 0],
1050
+ )
1051
+
1052
+ trace_kwargs = next(kwargs for method, kwargs in lane.calls if method == "can_trace_prefill")
1053
+ assert set(trace_kwargs) == {"tokens", "prompt_lens", "start_pos"}
1054
+ assert trace_kwargs["tokens"].flatten().tolist() == [10, 11]
1055
+ assert trace_kwargs["prompt_lens"].tolist() == [7, 9]
1056
+ assert trace_kwargs["start_pos"].tolist() == [3, 4]
1057
+
1058
+
1059
+ def test_successful_traced_prefill_preflights_all_lanes_before_first_execution():
1060
+ events = []
1061
+
1062
+ class TraceLane(_Lane):
1063
+ def __init__(self, lane_idx):
1064
+ super().__init__(lane_idx)
1065
+ self.traced_prefill_execution = object()
1066
+
1067
+ def can_trace_prefill(self, *, tokens, prompt_lens=None, start_pos=None):
1068
+ events.append(("preflight", self.lane_idx))
1069
+ return True
1070
+
1071
+ def prefill_forward(self, *args, **kwargs):
1072
+ events.append(("execute", self.lane_idx))
1073
+ return super().prefill_forward(*args, **kwargs)
1074
+
1075
+ group = LaneGroupExecutor([TraceLane(0), TraceLane(1)])
1076
+
1077
+ output = group.prefill_forward(
1078
+ tokens=torch.tensor([[10], [11]]),
1079
+ page_table=torch.tensor([[0], [1]], dtype=torch.int32),
1080
+ empty_slots=[0, 2],
1081
+ execution=group.traced_prefill_execution,
1082
+ )
1083
+
1084
+ assert output.flatten().tolist() == [10.0, 111.0]
1085
+ assert events == [("preflight", 0), ("preflight", 1), ("execute", 0), ("execute", 1)]
1086
+
1087
+
1088
+ @pytest.mark.parametrize(
1089
+ "unrelated_kwargs",
1090
+ (
1091
+ {"page_table": torch.zeros((1, 1), dtype=torch.int32)},
1092
+ {"kv_cache": ["kv-0"]},
1093
+ {"sampling_params": None},
1094
+ {"execution": (object(),)},
1095
+ ),
1096
+ )
1097
+ def test_trace_classification_rejects_unrelated_request_fields(unrelated_kwargs, expect_error):
1098
+ group = LaneGroupExecutor([_Lane(0)])
1099
+
1100
+ with expect_error(TypeError, "unexpected keyword argument"):
1101
+ group.can_trace_prefill(tokens=torch.tensor([[10]]), **unrelated_kwargs)
1102
+
1103
+ group.cleanup()
1104
+
1105
+
1106
+ def test_configure_and_allocate_fan_out_with_distinct_lane_configs():
1107
+ lanes = [_Lane(0), _Lane(1)]
1108
+ group = LaneGroupExecutor(lanes)
1109
+ config = SimpleNamespace(num_blocks=8)
1110
+
1111
+ group.configure_paged_kv_cache(config)
1112
+ handles = group.allocate_kv_cache((8, 2, 32, 64), torch.bfloat16, 2)
1113
+
1114
+ assert handles == ["cache-handle-0", "cache-handle-1"]
1115
+ lane_configs = [next(kwargs["config"] for method, kwargs in lane.calls if method == "configure") for lane in lanes]
1116
+ assert lane_configs[0] is not lane_configs[1]
1117
+ assert lane_configs[0].num_blocks == lane_configs[1].num_blocks == 8
1118
+ for lane in lanes:
1119
+ allocation = next(kwargs for method, kwargs in lane.calls if method == "allocate")
1120
+ assert allocation == {
1121
+ "kv_cache_shape": (8, 2, 32, 64),
1122
+ "dtype": torch.bfloat16,
1123
+ "num_layers": 2,
1124
+ }
1125
+
1126
+
1127
+ def test_side_effect_execution_target_methods_return_none_even_if_lanes_return_values():
1128
+ class ReturningLane(_Lane):
1129
+ def configure_paged_kv_cache(self, config):
1130
+ super().configure_paged_kv_cache(config)
1131
+ return self.lane_idx
1132
+
1133
+ def compile_prefill(
1134
+ self,
1135
+ *,
1136
+ tokens,
1137
+ page_table,
1138
+ prompt_lens=None,
1139
+ start_pos=None,
1140
+ empty_slots=None,
1141
+ kv_cache=None,
1142
+ sampling_params=None,
1143
+ execution=None,
1144
+ ):
1145
+ super().compile_prefill(
1146
+ tokens=tokens,
1147
+ page_table=page_table,
1148
+ prompt_lens=prompt_lens,
1149
+ start_pos=start_pos,
1150
+ empty_slots=empty_slots,
1151
+ kv_cache=kv_cache,
1152
+ sampling_params=sampling_params,
1153
+ execution=execution,
1154
+ )
1155
+ return self.lane_idx
1156
+
1157
+ def compile_decode(
1158
+ self,
1159
+ *,
1160
+ tokens,
1161
+ start_pos,
1162
+ page_table,
1163
+ kv_cache=None,
1164
+ sampling_params=None,
1165
+ reset_batch=False,
1166
+ execution=None,
1167
+ ):
1168
+ super().compile_decode(
1169
+ tokens=tokens,
1170
+ start_pos=start_pos,
1171
+ page_table=page_table,
1172
+ kv_cache=kv_cache,
1173
+ sampling_params=sampling_params,
1174
+ reset_batch=reset_batch,
1175
+ execution=execution,
1176
+ )
1177
+ return self.lane_idx
1178
+
1179
+ def warmup_model_prefill(self, *, kv_cache, can_sample_on_device, enable_trace):
1180
+ super().warmup_model_prefill(
1181
+ kv_cache=kv_cache,
1182
+ can_sample_on_device=can_sample_on_device,
1183
+ enable_trace=enable_trace,
1184
+ )
1185
+ return self.lane_idx
1186
+
1187
+ def warmup_model_decode(
1188
+ self,
1189
+ *,
1190
+ kv_cache,
1191
+ max_batch_size,
1192
+ num_blocks,
1193
+ can_sample_on_device,
1194
+ enable_trace,
1195
+ ):
1196
+ super().warmup_model_decode(
1197
+ kv_cache=kv_cache,
1198
+ max_batch_size=max_batch_size,
1199
+ num_blocks=num_blocks,
1200
+ can_sample_on_device=can_sample_on_device,
1201
+ enable_trace=enable_trace,
1202
+ )
1203
+ return self.lane_idx
1204
+
1205
+ group = LaneGroupExecutor([ReturningLane(0), ReturningLane(1)])
1206
+ config = SimpleNamespace(num_blocks=8)
1207
+ results = (
1208
+ group.configure_paged_kv_cache(config),
1209
+ group.compile_prefill(
1210
+ tokens=torch.tensor([[10], [11]]),
1211
+ page_table=torch.tensor([[0], [1]], dtype=torch.int32),
1212
+ empty_slots=[0, 2],
1213
+ ),
1214
+ group.compile_decode(
1215
+ tokens=torch.tensor([10, 11, 12, 13]),
1216
+ start_pos=torch.tensor([1, 2, 3, 4]),
1217
+ page_table=torch.arange(4, dtype=torch.int32).view(4, 1),
1218
+ ),
1219
+ group.warmup_model_prefill(
1220
+ kv_cache=None,
1221
+ can_sample_on_device=False,
1222
+ enable_trace=False,
1223
+ ),
1224
+ group.warmup_model_decode(
1225
+ kv_cache=None,
1226
+ max_batch_size=4,
1227
+ num_blocks=8,
1228
+ can_sample_on_device=False,
1229
+ enable_trace=False,
1230
+ ),
1231
+ )
1232
+
1233
+ assert results == (None, None, None, None, None)
1234
+
1235
+
1236
+ def test_constructor_validation_failure_cleans_every_supplied_lane(expect_error):
1237
+ lanes = [_Lane(0, capacity=2), _Lane(1, capacity=4)]
1238
+
1239
+ with expect_error(ValueError, "same fixed capacity"):
1240
+ LaneGroupExecutor(lanes)
1241
+
1242
+ assert [lane.cleanup_calls for lane in lanes] == [1, 1]
1243
+
1244
+
1245
+ def test_operation_failure_is_primary_group_becomes_terminal_and_all_lanes_cleanup(expect_error):
1246
+ lanes = [_Lane(0), _Lane(1)]
1247
+ lanes[1].fail_method = "allocate"
1248
+ lanes[0].cleanup_error = RuntimeError("cleanup boom")
1249
+ group = LaneGroupExecutor(lanes)
1250
+
1251
+ with expect_error(RuntimeError, "allocate boom 1") as exc_info:
1252
+ group.allocate_kv_cache()
1253
+
1254
+ assert lanes[0].cleanup_calls == 1
1255
+ assert lanes[1].cleanup_calls == 1
1256
+ assert [str(error) for error in exc_info.value.cleanup_failures] == ["cleanup boom"]
1257
+ with expect_error(RuntimeError, "terminal"):
1258
+ group.allocate_kv_cache()
1259
+
1260
+
1261
+ def test_non_null_lane_log_probs_are_aggregated_in_global_row_order():
1262
+ lanes = [_Lane(0), _Lane(1)]
1263
+
1264
+ def decode_with_log_probs(
1265
+ tokens,
1266
+ start_pos,
1267
+ page_table,
1268
+ *,
1269
+ kv_cache=None,
1270
+ sampling_params=None,
1271
+ reset_batch=False,
1272
+ read_from_device=True,
1273
+ execution=None,
1274
+ ):
1275
+ return (tokens + 100).to(torch.int64), torch.tensor([0.25, 0.75])
1276
+
1277
+ lanes[1].decode_forward = decode_with_log_probs
1278
+ group = LaneGroupExecutor(lanes)
1279
+
1280
+ output, log_probs = group.decode_forward(
1281
+ tokens=torch.tensor([0, 1, 2, 3]),
1282
+ start_pos=torch.tensor([0, 0, 0, 0]),
1283
+ page_table=torch.zeros((4, 1), dtype=torch.int32),
1284
+ sampling_params=_sampling(),
1285
+ )
1286
+
1287
+ assert output.tolist() == [0, 1, 102, 103]
1288
+ assert isinstance(log_probs, torch.Tensor)
1289
+ assert log_probs.tolist() == [1.0, 1.0, 0.25, 0.75]
1290
+
1291
+
1292
+ def test_cleanup_is_idempotent_and_terminal(expect_error):
1293
+ lanes = [_Lane(0), _Lane(1)]
1294
+ group = LaneGroupExecutor(lanes)
1295
+
1296
+ group.cleanup()
1297
+ group.cleanup()
1298
+
1299
+ assert [lane.cleanup_calls for lane in lanes] == [1, 1]
1300
+ with expect_error(RuntimeError, "terminal"):
1301
+ group.prefill_forward(
1302
+ tokens=torch.zeros((1, 1), dtype=torch.long),
1303
+ page_table=torch.zeros((1, 1), dtype=torch.int32),
1304
+ )
1305
+
1306
+
1307
+ def test_cleanup_retries_only_lanes_that_failed_cleanup_and_stays_terminal(expect_error):
1308
+ class FailOnceCleanupLane(_Lane):
1309
+ def cleanup(self):
1310
+ self.cleanup_calls += 1
1311
+ if self.cleanup_calls == 1:
1312
+ raise RuntimeError("lane cleanup failed once")
1313
+
1314
+ failing_lane = FailOnceCleanupLane(0)
1315
+ successful_lane = _Lane(1)
1316
+ group = LaneGroupExecutor([failing_lane, successful_lane])
1317
+
1318
+ with expect_error(RuntimeError, "lane cleanup failed once"):
1319
+ group.cleanup()
1320
+
1321
+ assert group.terminal
1322
+ assert [failing_lane.cleanup_calls, successful_lane.cleanup_calls] == [1, 1]
1323
+ with expect_error(RuntimeError, "terminal"):
1324
+ group.allocate_kv_cache()
1325
+
1326
+ group.cleanup()
1327
+ group.cleanup()
1328
+
1329
+ assert group.terminal
1330
+ assert [failing_lane.cleanup_calls, successful_lane.cleanup_calls] == [2, 1]
1331
+
1332
+
1333
+ def test_cleanup_retries_failed_pool_shutdown_without_recleaning_lanes(expect_error):
1334
+ class FailOnceShutdownPool:
1335
+ def __init__(self, pool):
1336
+ self.pool = pool
1337
+ self.shutdown_calls = 0
1338
+
1339
+ def submit(self, *args, **kwargs):
1340
+ return self.pool.submit(*args, **kwargs)
1341
+
1342
+ def shutdown(self, *, wait):
1343
+ self.shutdown_calls += 1
1344
+ if self.shutdown_calls == 1:
1345
+ raise RuntimeError("pool shutdown failed once")
1346
+ self.pool.shutdown(wait=wait)
1347
+
1348
+ lanes = [_Lane(0), _Lane(1)]
1349
+ group = LaneGroupExecutor(lanes)
1350
+ assert group._output_pool is not None
1351
+ pool = FailOnceShutdownPool(group._output_pool)
1352
+ group._output_pool = pool
1353
+
1354
+ with expect_error(RuntimeError, "pool shutdown failed once"):
1355
+ group.cleanup()
1356
+
1357
+ assert group.terminal
1358
+ assert [lane.cleanup_calls for lane in lanes] == [1, 1]
1359
+
1360
+ group.cleanup()
1361
+ group.cleanup()
1362
+
1363
+ assert pool.shutdown_calls == 2
1364
+ assert [lane.cleanup_calls for lane in lanes] == [1, 1]
1365
+
1366
+
1367
+ def test_explicit_mesh_device_is_preserved_by_identity():
1368
+ mesh_device = object()
1369
+ group = LaneGroupExecutor([_Lane(0), _Lane(1)], mesh_device=mesh_device)
1370
+
1371
+ assert group.mesh_device is mesh_device
1372
+
1373
+ group.cleanup()
code/models/common/tests/llm_runtime/test_llama3_8b_integration.py ADDED
@@ -0,0 +1,1065 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from __future__ import annotations
5
+
6
+ import inspect
7
+ from types import SimpleNamespace
8
+ from typing import Any, Sequence
9
+ from unittest.mock import create_autospec
10
+
11
+ import pytest
12
+ import torch
13
+
14
+ import ttnn
15
+ from models.common.llm_runtime.config import PagedKVCacheConfig, PageTableLayout, TraceConfig, WarmupConfig
16
+ from models.common.llm_runtime.execution import EagerExecutor, TracedExecutor
17
+ from models.common.llm_runtime.lane_group import LaneGroupExecutor
18
+ from models.common.llm_runtime.warmup import _build_plan
19
+ from models.common.models import llama3_executor as llama3_family_executor
20
+ from models.common.models.llama3_8b import executor as llama_executor
21
+ from models.common.models.llama3_8b import generator as llama_generator
22
+ from models.common.models.llama3_8b import model as llama_model
23
+ from models.common.tests.demos.llama3_8b import demo as llama_demo
24
+
25
+
26
+ class _Mesh:
27
+ shape = (1, 1)
28
+
29
+ @staticmethod
30
+ def get_num_devices():
31
+ return 1
32
+
33
+
34
+ def _model(*, max_batch_size=4, max_seq_len=4096):
35
+ paged = SimpleNamespace(block_size=32, max_num_blocks=132)
36
+ attention = SimpleNamespace(
37
+ n_kv_heads=8,
38
+ head_dim=128,
39
+ kv_cache_dtype=ttnn.bfloat8_b,
40
+ paged_attention_config=paged,
41
+ use_vllm_paged_kv_cache=True,
42
+ kv_cache=None,
43
+ )
44
+ attention_module = SimpleNamespace(config=attention, kv_cache=None)
45
+ model = SimpleNamespace(
46
+ config=SimpleNamespace(
47
+ mesh_device=_Mesh(),
48
+ max_batch_size=max_batch_size,
49
+ max_seq_len=max_seq_len,
50
+ n_layers=1,
51
+ num_devices=1,
52
+ block_configs=(SimpleNamespace(attention_config=attention),),
53
+ ),
54
+ layers=(SimpleNamespace(attention=attention_module),),
55
+ iter_executor_named_modules=lambda: (),
56
+ vocab_size=128,
57
+ num_devices=1,
58
+ )
59
+
60
+ def configure_paged_attention(*, block_size, max_num_blocks):
61
+ assert attention.kv_cache is None
62
+ assert attention_module.kv_cache is None
63
+ attention.paged_attention_config = SimpleNamespace(
64
+ block_size=block_size,
65
+ max_num_blocks=max_num_blocks,
66
+ )
67
+
68
+ model.configure_paged_attention = configure_paged_attention
69
+ return model
70
+
71
+
72
+ def _runtime_config():
73
+ return SimpleNamespace(
74
+ model_cache_path="cache",
75
+ max_prefill_chunk_size=2048,
76
+ trace_prefill_supported_seq_lens=(128, 1024),
77
+ supports_batched_prefill=True,
78
+ disable_batched_prefill=False,
79
+ max_prefill_batch_size=32,
80
+ batched_prefill_batched_extract=True,
81
+ can_enable_trace=lambda sequence_length, num_cached_tokens=0: (
82
+ num_cached_tokens == 0 and sequence_length in (128, 1024)
83
+ ),
84
+ )
85
+
86
+
87
+ def _config(mode="none", *, num_blocks=None):
88
+ return llama_executor.Llama3ExecutorConfig(
89
+ trace=TraceConfig(mode),
90
+ warmup=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
91
+ paged_kv_cache=PagedKVCacheConfig(
92
+ block_size=32,
93
+ max_num_blocks=132,
94
+ dtype=ttnn.bfloat8_b,
95
+ num_blocks=num_blocks,
96
+ ),
97
+ device_sampling_enabled=False,
98
+ )
99
+
100
+
101
+ @pytest.mark.parametrize(
102
+ ("runtime_disabled", "device_sampling_enabled", "diagnostic_override", "expected_disabled"),
103
+ [
104
+ (False, False, False, False),
105
+ (True, False, False, True),
106
+ (False, True, False, True),
107
+ (True, True, False, True),
108
+ (False, True, True, False),
109
+ (True, True, True, False),
110
+ ],
111
+ )
112
+ def test_executor_resolves_batched_prefill_policy(
113
+ monkeypatch, runtime_disabled, device_sampling_enabled, diagnostic_override, expected_disabled
114
+ ):
115
+ runtime_config = _runtime_config()
116
+ runtime_config.disable_batched_prefill = runtime_disabled
117
+ base_config = _config()
118
+ config = llama_executor.Llama3ExecutorConfig(
119
+ trace=base_config.trace,
120
+ warmup=base_config.warmup,
121
+ paged_kv_cache=base_config.paged_kv_cache,
122
+ device_sampling_enabled=device_sampling_enabled,
123
+ allow_batched_prefill_with_device_sampling_for_diagnostics=diagnostic_override,
124
+ )
125
+ model = _model()
126
+ if device_sampling_enabled:
127
+
128
+ class _Sampling:
129
+ def decode_forward(self):
130
+ raise AssertionError("construction-policy test must not execute sampling")
131
+
132
+ class _SamplingState:
133
+ def __init__(self, sampling, **_kwargs):
134
+ self.sampling = sampling
135
+ self.seed_manager = SimpleNamespace()
136
+
137
+ def create_state(self):
138
+ return SimpleNamespace(seed_state=SimpleNamespace(capacity=32))
139
+
140
+ def admit(self, *args, **kwargs):
141
+ return None
142
+
143
+ def decode_forward(self, *args, **kwargs):
144
+ return None
145
+
146
+ def release(self, *args, **kwargs):
147
+ return None
148
+
149
+ monkeypatch.setattr(llama3_family_executor, "Sampling1D", _Sampling)
150
+ monkeypatch.setattr(llama3_family_executor, "SamplingState1D", _SamplingState)
151
+ sampler = _Sampling()
152
+ sampler.config = SimpleNamespace(
153
+ is_resolved=lambda: True,
154
+ allow_force_argmax=True,
155
+ max_batch_size=32,
156
+ max_top_k=32,
157
+ )
158
+ model.sampling = sampler
159
+
160
+ executor = llama_executor.Llama3Executor(model, runtime_config, config)
161
+
162
+ assert executor.prefill_runtime.config.supports_batched_prefill is True
163
+ assert executor.prefill_runtime.config.disable_batched_prefill is expected_disabled
164
+ assert executor.prefill_runtime.config.max_prefill_batch_size == 32
165
+ assert executor.prefill_runtime.config.batched_prefill_batched_extract is True
166
+
167
+
168
+ def test_executor_rejects_batched_prefill_diagnostic_override_without_device_sampling(expect_error):
169
+ base_config = _config()
170
+
171
+ with expect_error(ValueError, "requires device_sampling_enabled"):
172
+ llama_executor.Llama3ExecutorConfig(
173
+ trace=base_config.trace,
174
+ warmup=base_config.warmup,
175
+ paged_kv_cache=base_config.paged_kv_cache,
176
+ device_sampling_enabled=False,
177
+ allow_batched_prefill_with_device_sampling_for_diagnostics=True,
178
+ )
179
+
180
+
181
+ @pytest.mark.parametrize("mode", ["none", "decode_only", "all"])
182
+ def test_model_owned_executor_constructs_exact_composition(mode):
183
+ executor = llama_executor.Llama3Executor(_model(), _runtime_config(), _config(mode))
184
+
185
+ assert executor.eager_executor.program_compiler is executor.program_compiler
186
+ assert executor.eager_executor.prefill is executor.prefill_runtime
187
+ assert executor.eager_executor.decode is executor.decode_runtime
188
+ assert executor.warmup.eager is executor.eager_executor
189
+ assert executor.warmup.trace_compiler is executor.trace_compiler
190
+ if mode == "none":
191
+ assert executor.warmup.execution is executor.eager_executor
192
+ assert executor.eager_execution is executor.eager_executor
193
+ assert executor.traced_prefill_execution is None
194
+ assert executor.traced_decode_execution is None
195
+ assert executor.trace_compiler is None
196
+ assert executor.traced_executor is None
197
+ else:
198
+ assert executor.warmup.execution is executor.traced_executor
199
+ expected_prefill = executor.traced_executor if mode == "all" else None
200
+ assert executor.eager_execution is executor.eager_executor
201
+ assert executor.traced_prefill_execution is expected_prefill
202
+ assert executor.traced_decode_execution is executor.traced_executor
203
+ assert executor.traced_executor.eager_executor is executor.eager_executor
204
+ assert executor.traced_executor.trace_compiler is executor.trace_compiler
205
+ assert executor.trace_compiler.program_compiler is executor.program_compiler
206
+
207
+
208
+ @pytest.mark.parametrize(
209
+ ("method_name", "positional_names", "keyword_only_names"),
210
+ [
211
+ (
212
+ "compile_prefill",
213
+ ("self",),
214
+ (
215
+ "tokens",
216
+ "page_table",
217
+ "prompt_lens",
218
+ "start_pos",
219
+ "empty_slots",
220
+ "kv_cache",
221
+ "sampling_params",
222
+ "prompt_tokens",
223
+ "output_tokens",
224
+ "slot_remap",
225
+ "execution",
226
+ ),
227
+ ),
228
+ (
229
+ "compile_decode",
230
+ ("self",),
231
+ (
232
+ "tokens",
233
+ "start_pos",
234
+ "page_table",
235
+ "kv_cache",
236
+ "sampling_params",
237
+ "prompt_tokens",
238
+ "output_tokens",
239
+ "slot_remap",
240
+ "reset_batch",
241
+ "execution",
242
+ ),
243
+ ),
244
+ (
245
+ "prefill_forward",
246
+ ("self", "tokens", "page_table"),
247
+ (
248
+ "prompt_lens",
249
+ "start_pos",
250
+ "empty_slots",
251
+ "kv_cache",
252
+ "sampling_params",
253
+ "prompt_tokens",
254
+ "output_tokens",
255
+ "slot_remap",
256
+ "execution",
257
+ ),
258
+ ),
259
+ (
260
+ "decode_forward",
261
+ ("self", "tokens", "start_pos", "page_table"),
262
+ (
263
+ "kv_cache",
264
+ "sampling_params",
265
+ "prompt_tokens",
266
+ "output_tokens",
267
+ "slot_remap",
268
+ "reset_batch",
269
+ "read_from_device",
270
+ "execution",
271
+ ),
272
+ ),
273
+ ("read_decode_output", ("self", "tt_out"), ("async_read",)),
274
+ ("process_decode_output_host", ("self", "tt_out"), ("is_tokens",)),
275
+ (
276
+ "can_trace_prefill",
277
+ ("self",),
278
+ ("tokens", "prompt_lens", "start_pos", "empty_slots"),
279
+ ),
280
+ (
281
+ "warmup_model_prefill",
282
+ ("self",),
283
+ ("kv_cache", "can_sample_on_device", "enable_trace"),
284
+ ),
285
+ (
286
+ "warmup_model_decode",
287
+ ("self",),
288
+ ("kv_cache", "max_batch_size", "num_blocks", "can_sample_on_device", "enable_trace"),
289
+ ),
290
+ ],
291
+ )
292
+ def test_model_owned_executor_has_exact_call_contract(method_name, positional_names, keyword_only_names):
293
+ signature = inspect.signature(getattr(llama_executor.Llama3Executor, method_name))
294
+ parameters = signature.parameters
295
+ required_names = {
296
+ "compile_prefill": {"tokens", "page_table"},
297
+ "compile_decode": {"tokens", "start_pos", "page_table"},
298
+ "prefill_forward": {"tokens", "page_table"},
299
+ "decode_forward": {"tokens", "start_pos", "page_table"},
300
+ "read_decode_output": {"tt_out"},
301
+ "process_decode_output_host": {"tt_out"},
302
+ "can_trace_prefill": {"tokens"},
303
+ "warmup_model_prefill": {"kv_cache", "can_sample_on_device", "enable_trace"},
304
+ "warmup_model_decode": {
305
+ "kv_cache",
306
+ "max_batch_size",
307
+ "num_blocks",
308
+ "can_sample_on_device",
309
+ "enable_trace",
310
+ },
311
+ }[method_name]
312
+ special_defaults = {
313
+ "reset_batch": False,
314
+ "read_from_device": True,
315
+ "async_read": False,
316
+ "is_tokens": False,
317
+ }
318
+
319
+ assert tuple(parameters) == positional_names + keyword_only_names
320
+ assert all(parameters[name].kind is inspect.Parameter.POSITIONAL_OR_KEYWORD for name in positional_names)
321
+ assert all(parameters[name].kind is inspect.Parameter.KEYWORD_ONLY for name in keyword_only_names)
322
+ for name, parameter in tuple(parameters.items())[1:]:
323
+ expected_default = inspect.Parameter.empty if name in required_names else special_defaults.get(name)
324
+ assert parameter.default == expected_default
325
+ assert parameter.annotation is not inspect.Parameter.empty
326
+ assert signature.return_annotation is not inspect.Signature.empty
327
+ if "execution" in parameters:
328
+ assert parameters["execution"].annotation == "EagerExecutor | TracedExecutor | None"
329
+
330
+
331
+ def test_model_owned_executor_validates_cache_then_omits_it_from_execution():
332
+ execution = create_autospec(EagerExecutor, instance=True)
333
+ execution.prefill_forward.return_value = "prefill"
334
+ execution.decode_forward.return_value = "decode"
335
+ executor = object.__new__(llama_executor.Llama3Executor)
336
+ executor._prefill_execution = execution
337
+ executor._decode_execution = execution
338
+ executor._ensure_active = lambda: None
339
+ validated_caches = []
340
+ sampling_values = []
341
+ executor._validate_bound_cache = validated_caches.append
342
+ executor._ensure_sampling_for = sampling_values.append
343
+
344
+ tokens = torch.zeros((1, 4), dtype=torch.long)
345
+ start_pos = torch.zeros((1,), dtype=torch.long)
346
+ page_table = torch.zeros((1, 1), dtype=torch.int32)
347
+ prompt_lens = torch.full((1,), 4, dtype=torch.long)
348
+ empty_slots = [0]
349
+ kv_cache = object()
350
+ sampling_params = object()
351
+
352
+ executor.compile_prefill(
353
+ tokens=tokens,
354
+ page_table=page_table,
355
+ prompt_lens=prompt_lens,
356
+ start_pos=start_pos,
357
+ empty_slots=empty_slots,
358
+ kv_cache=kv_cache,
359
+ sampling_params=sampling_params,
360
+ )
361
+ executor.compile_decode(
362
+ tokens=tokens,
363
+ start_pos=start_pos,
364
+ page_table=page_table,
365
+ kv_cache=kv_cache,
366
+ sampling_params=sampling_params,
367
+ reset_batch=True,
368
+ )
369
+ assert (
370
+ executor.prefill_forward(
371
+ tokens,
372
+ page_table,
373
+ prompt_lens=prompt_lens,
374
+ start_pos=start_pos,
375
+ empty_slots=empty_slots,
376
+ kv_cache=kv_cache,
377
+ sampling_params=sampling_params,
378
+ )
379
+ == "prefill"
380
+ )
381
+ assert (
382
+ executor.decode_forward(
383
+ tokens,
384
+ start_pos,
385
+ page_table,
386
+ kv_cache=kv_cache,
387
+ sampling_params=sampling_params,
388
+ reset_batch=True,
389
+ read_from_device=False,
390
+ )
391
+ == "decode"
392
+ )
393
+
394
+ assert validated_caches == [kv_cache] * 4
395
+ assert sampling_values == [sampling_params] * 4
396
+ for target, expected_names in (
397
+ (
398
+ execution.compile_prefill,
399
+ (
400
+ "tokens",
401
+ "page_table",
402
+ "prompt_lens",
403
+ "start_pos",
404
+ "empty_slots",
405
+ "sampling_params",
406
+ "prompt_tokens",
407
+ "output_tokens",
408
+ "slot_remap",
409
+ ),
410
+ ),
411
+ (
412
+ execution.compile_decode,
413
+ (
414
+ "tokens",
415
+ "start_pos",
416
+ "page_table",
417
+ "sampling_params",
418
+ "prompt_tokens",
419
+ "output_tokens",
420
+ "slot_remap",
421
+ "reset_batch",
422
+ ),
423
+ ),
424
+ (
425
+ execution.prefill_forward,
426
+ (
427
+ "tokens",
428
+ "page_table",
429
+ "prompt_lens",
430
+ "start_pos",
431
+ "empty_slots",
432
+ "sampling_params",
433
+ "prompt_tokens",
434
+ "output_tokens",
435
+ "slot_remap",
436
+ ),
437
+ ),
438
+ (
439
+ execution.decode_forward,
440
+ (
441
+ "tokens",
442
+ "start_pos",
443
+ "page_table",
444
+ "sampling_params",
445
+ "prompt_tokens",
446
+ "output_tokens",
447
+ "slot_remap",
448
+ "reset_batch",
449
+ "read_from_device",
450
+ ),
451
+ ),
452
+ ):
453
+ assert target.call_count == 1
454
+ assert tuple(target.call_args.kwargs) == expected_names
455
+
456
+
457
+ def test_model_owned_executor_trace_output_and_warmup_forwarding_is_named():
458
+ calls = []
459
+
460
+ class _PrefillRuntime:
461
+ def can_trace(
462
+ self,
463
+ *,
464
+ tokens: torch.Tensor, # ↓ Core request
465
+ prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata
466
+ start_pos: torch.Tensor | None = None,
467
+ ) -> bool:
468
+ calls.append(("can_trace", tokens, prompt_lens, start_pos))
469
+ return True
470
+
471
+ class _DecodeRuntime:
472
+ def read_decode_output(self, tt_out: Any, *, async_read: bool = False) -> Any:
473
+ calls.append(("read_decode_output", tt_out, async_read))
474
+ return "read"
475
+
476
+ def process_decode_output_host(self, tt_out: Any, *, is_tokens: bool = False) -> tuple[Any, Any]:
477
+ calls.append(("process_decode_output_host", tt_out, is_tokens))
478
+ return "processed", "event"
479
+
480
+ class _Warmup:
481
+ def warmup_prefill(
482
+ self,
483
+ *,
484
+ kv_cache: Any, # ↓ Borrowed resources
485
+ can_sample_on_device: bool, # ↓ Execution policy
486
+ enable_trace: bool,
487
+ ) -> None:
488
+ calls.append(("warmup_prefill", kv_cache, can_sample_on_device, enable_trace))
489
+
490
+ def warmup_decode(
491
+ self,
492
+ *,
493
+ kv_cache: Any, # ↓ Borrowed resources
494
+ max_batch_size: int, # ↓ Coverage dimensions
495
+ num_blocks: int,
496
+ can_sample_on_device: bool, # ↓ Execution policy
497
+ enable_trace: bool,
498
+ ) -> None:
499
+ calls.append(("warmup_decode", kv_cache, max_batch_size, num_blocks, can_sample_on_device, enable_trace))
500
+
501
+ executor = object.__new__(llama_executor.Llama3Executor)
502
+ executor.traced_executor = object()
503
+ executor.config = SimpleNamespace(trace=SimpleNamespace(prefill_enabled=True))
504
+ executor.prefill_runtime = _PrefillRuntime()
505
+ executor.decode_runtime = _DecodeRuntime()
506
+ executor.warmup = _Warmup()
507
+ executor._ensure_active = lambda: None
508
+
509
+ tokens = torch.zeros((1, 4), dtype=torch.long)
510
+ prompt_lens = torch.full((1,), 4, dtype=torch.long)
511
+ start_pos = torch.zeros((1,), dtype=torch.long)
512
+ kv_cache = object()
513
+
514
+ assert executor.can_trace_prefill(
515
+ tokens=tokens,
516
+ prompt_lens=prompt_lens,
517
+ start_pos=start_pos,
518
+ empty_slots=[0],
519
+ )
520
+ assert executor.read_decode_output("device-output", async_read=True) == "read"
521
+ assert executor.process_decode_output_host("host-output", is_tokens=True) == ("processed", "event")
522
+ executor.warmup_model_prefill(
523
+ kv_cache=kv_cache,
524
+ can_sample_on_device=True,
525
+ enable_trace=False,
526
+ )
527
+ executor.warmup_model_decode(
528
+ kv_cache=kv_cache,
529
+ max_batch_size=4,
530
+ num_blocks=8,
531
+ can_sample_on_device=True,
532
+ enable_trace=False,
533
+ )
534
+
535
+ assert calls == [
536
+ ("can_trace", tokens, prompt_lens, start_pos),
537
+ ("read_decode_output", "device-output", True),
538
+ ("process_decode_output_host", "host-output", True),
539
+ ("warmup_prefill", kv_cache, True, False),
540
+ ("warmup_decode", kv_cache, 4, 8, True, False),
541
+ ]
542
+
543
+
544
+ def test_vllm_capacity_resolution_reconfigures_existing_runtime_owners_before_allocation(monkeypatch):
545
+ executor = llama_executor.Llama3Executor(_model(), _runtime_config(), _config())
546
+ owner_ids = tuple(
547
+ id(owner)
548
+ for owner in (
549
+ executor.prefill_runtime,
550
+ executor.decode_runtime,
551
+ executor.warmup,
552
+ executor.program_compiler,
553
+ )
554
+ )
555
+ assert executor.page_table_layout.raw_capacity_width == 128
556
+
557
+ executor.configure_paged_kv_cache(
558
+ PagedKVCacheConfig(
559
+ block_size=16,
560
+ max_num_blocks=200,
561
+ dtype=ttnn.bfloat8_b,
562
+ num_blocks=200,
563
+ )
564
+ )
565
+
566
+ assert (
567
+ tuple(
568
+ id(owner)
569
+ for owner in (
570
+ executor.prefill_runtime,
571
+ executor.decode_runtime,
572
+ executor.warmup,
573
+ executor.program_compiler,
574
+ )
575
+ )
576
+ == owner_ids
577
+ )
578
+ assert executor.config.paged_kv_cache is executor.kv_cache_manager.config
579
+ assert executor.kv_cache_manager.config.block_size == 16
580
+ assert executor.kv_cache_manager.config.max_num_blocks == executor.kv_cache_manager.config.num_blocks == 200
581
+ assert executor.model.layers[0].attention.config.paged_attention_config.block_size == 16
582
+ assert executor.model.layers[0].attention.config.paged_attention_config.max_num_blocks == 200
583
+ assert executor.page_table_layout.block_size == 16
584
+ assert executor.page_table_layout.raw_capacity_width == 200
585
+ assert executor.prefill_runtime.config.page_table_layout is executor.page_table_layout
586
+ assert executor.decode_runtime.config.page_table_layout is executor.page_table_layout
587
+ assert executor.warmup.config.page_table_layout is executor.page_table_layout
588
+
589
+ def fake_allocate():
590
+ assert executor._runtime_configuration_sealed
591
+ assert executor.warmup._configuration_sealed
592
+ return ["allocated"]
593
+
594
+ monkeypatch.setattr(executor.kv_cache_manager, "allocate", fake_allocate)
595
+ assert executor.allocate_kv_cache() == ["allocated"]
596
+
597
+
598
+ def test_model_reconfigures_construction_and_live_attention_without_allocating():
599
+ construction_paged = llama_model.Llama31_8BPagedAttentionConfig(block_size=32, max_num_blocks=132)
600
+ live_paged = llama_model.Llama31_8BPagedAttentionConfig(block_size=32, max_num_blocks=132)
601
+ construction_attention = SimpleNamespace(
602
+ use_vllm_paged_kv_cache=True,
603
+ paged_attention_config=construction_paged,
604
+ )
605
+ live_attention = SimpleNamespace(
606
+ use_vllm_paged_kv_cache=True,
607
+ paged_attention_config=live_paged,
608
+ kv_cache=None,
609
+ )
610
+ model = object.__new__(llama_model.Llama3Transformer1D)
611
+ model.config = SimpleNamespace(
612
+ block_configs=(SimpleNamespace(attention_config=construction_attention),),
613
+ )
614
+ model.layers = (SimpleNamespace(attention=SimpleNamespace(config=live_attention, kv_cache=None)),)
615
+
616
+ model.configure_paged_attention(block_size=16, max_num_blocks=200)
617
+
618
+ assert construction_attention.paged_attention_config.block_size == 16
619
+ assert construction_attention.paged_attention_config.max_num_blocks == 200
620
+ assert live_attention.paged_attention_config.block_size == 16
621
+ assert live_attention.paged_attention_config.max_num_blocks == 200
622
+ assert live_attention.kv_cache is None
623
+ assert model.layers[0].attention.kv_cache is None
624
+
625
+
626
+ def test_late_vllm_capacity_resolution_fails_before_mutating_kv_configuration(expect_error):
627
+ executor = llama_executor.Llama3Executor(_model(), _runtime_config(), _config())
628
+ executor._seal_runtime_configuration()
629
+ unresolved = executor.kv_cache_manager.config
630
+
631
+ with expect_error(RuntimeError, "runtime configuration is sealed"):
632
+ executor.configure_paged_kv_cache(
633
+ PagedKVCacheConfig(
634
+ block_size=32,
635
+ max_num_blocks=132,
636
+ dtype=ttnn.bfloat8_b,
637
+ num_blocks=64,
638
+ )
639
+ )
640
+
641
+ assert executor.kv_cache_manager.config is unresolved
642
+ assert not executor.kv_cache_manager.config.is_resolved()
643
+
644
+
645
+ def test_direct_demo_resolves_physical_capacity_to_configured_maximum(monkeypatch):
646
+ attention = SimpleNamespace(
647
+ paged_attention_config=SimpleNamespace(block_size=32, max_num_blocks=128),
648
+ kv_cache_dtype=ttnn.bfloat8_b,
649
+ )
650
+ llm = SimpleNamespace(
651
+ model=SimpleNamespace(config=SimpleNamespace(block_configs=(SimpleNamespace(attention_config=attention),)))
652
+ )
653
+ captured = []
654
+ monkeypatch.setattr(
655
+ llama_demo, "build_llama3_executor", lambda product, config: captured.append(config) or object()
656
+ )
657
+
658
+ llama_demo._build_demo_executor(llm, trace_mode="all", device_sampling_enabled=False)
659
+
660
+ assert captured[0].paged_kv_cache.num_blocks == captured[0].paged_kv_cache.max_num_blocks == 128
661
+
662
+
663
+ def test_model_owned_cleanup_is_ordered_best_effort_retryable_and_idempotent(expect_error):
664
+ calls = []
665
+ failures = {"reader", "trace"}
666
+
667
+ class _Owner:
668
+ def __init__(self, name):
669
+ self.name = name
670
+
671
+ def cleanup(self, *args):
672
+ calls.append(self.name)
673
+ if self.name in failures:
674
+ raise RuntimeError(self.name)
675
+
676
+ drain = cleanup
677
+ drain_external_outputs = cleanup
678
+ cleanup_transients = cleanup
679
+ release = cleanup
680
+
681
+ executor = object.__new__(llama_executor.Llama3Executor)
682
+ executor._terminal = False
683
+ executor._cleaned_up = False
684
+ executor.decode_runtime = _Owner("decode-external")
685
+ executor.output_reader = _Owner("reader")
686
+ executor.prefill_runtime = _Owner("prefill")
687
+ executor.trace_compiler = _Owner("trace")
688
+ executor.program_compiler = _Owner("program")
689
+ executor.config = SimpleNamespace(device_sampling_enabled=True)
690
+ executor.model = SimpleNamespace(sampling=_Owner("sampling"))
691
+ executor.sampling_state_controller = _Owner("sampling-state")
692
+ executor.sampling_state = object()
693
+ executor.kv_cache_manager = _Owner("kv")
694
+
695
+ with expect_error(RuntimeError, "reader") as raised:
696
+ executor.cleanup()
697
+
698
+ expected_order = [
699
+ "decode-external",
700
+ "reader",
701
+ "prefill",
702
+ "decode-external",
703
+ "trace",
704
+ "program",
705
+ "sampling-state",
706
+ "sampling",
707
+ "kv",
708
+ ]
709
+ assert calls == expected_order
710
+ assert tuple(error.args[0] for error in raised.value.cleanup_failures) == ("trace",)
711
+ assert executor.terminal
712
+ assert not executor._cleaned_up
713
+
714
+ failures.clear()
715
+ executor.cleanup()
716
+ assert calls == expected_order * 2
717
+ assert executor._cleaned_up
718
+
719
+ executor.cleanup()
720
+ assert calls == expected_order * 2
721
+
722
+
723
+ def test_build_llama3_executor_uses_prebuilt_product(monkeypatch):
724
+ sentinel = object()
725
+ calls = []
726
+ llm = SimpleNamespace(model=object(), runtime_config=object())
727
+ monkeypatch.setattr(
728
+ llama_executor,
729
+ "Llama3Executor",
730
+ lambda model, runtime_config, config: calls.append((model, runtime_config, config)) or sentinel,
731
+ )
732
+
733
+ config = object()
734
+ assert llama_executor.build_llama3_executor(llm, config) is sentinel
735
+ assert calls == [(llm.model, llm.runtime_config, config)]
736
+
737
+
738
+ def test_configured_path_has_no_legacy_or_common_aggregate_surface():
739
+ source = inspect.getsource(llama_executor)
740
+ assert not hasattr(llama_executor, "EagerLlamaExecutor")
741
+ assert not hasattr(llama_executor, "TracedLlamaExecutor")
742
+ assert "models.common.models.executor" not in source
743
+ assert "llm_runtime.executor" not in source
744
+ assert "class LLMExecutor" not in source
745
+ assert EagerExecutor not in TracedExecutor.__mro__
746
+
747
+
748
+ def test_generator_explicitly_accepts_static_trace_mode():
749
+ assert llama_generator.Llama3Generator.model_capabilities["accepts_trace_mode"] is True
750
+
751
+
752
+ def test_initialize_vllm_model_threads_policy(monkeypatch):
753
+ captured = []
754
+ sentinel = object()
755
+ monkeypatch.setattr(
756
+ llama_generator,
757
+ "build_llama3_generator",
758
+ lambda config: captured.append(config) or sentinel,
759
+ )
760
+
761
+ result = llama_generator.Llama3Generator.initialize_vllm_model(
762
+ SimpleNamespace(_name_or_path="meta-llama/Llama-3.1-8B-Instruct"),
763
+ object(),
764
+ 8,
765
+ 4096,
766
+ n_layers=3,
767
+ tt_data_parallel=2,
768
+ optimizations="accuracy",
769
+ trace_mode="decode_only",
770
+ device_sampling_enabled=True,
771
+ )
772
+
773
+ assert result is sentinel
774
+ assert captured[0].tt_data_parallel == 2
775
+ assert captured[0].trace_mode == "decode_only"
776
+ assert captured[0].device_sampling_enabled is True
777
+
778
+
779
+ class _FakeLane:
780
+ requires_prefill_trace_warmup = True
781
+
782
+ def __init__(self, llm, config):
783
+ self.model = llm.model
784
+ self.model_args = llm.runtime_config
785
+ self.mesh_device = llm.model.config.mesh_device
786
+ self.cache_path = llm.runtime_config.model_cache_path
787
+ self._request_state_fields = ("prompt_tokens", "output_tokens", "slot_remap")
788
+ self.config = config
789
+ self.paged_kv_cache_config = config.paged_kv_cache
790
+ self.already_warmed_up_prefill = False
791
+ self.cleanup_calls = 0
792
+
793
+ def cleanup(self):
794
+ self.cleanup_calls += 1
795
+
796
+
797
+ @pytest.mark.parametrize("device_sampling_enabled", (False, True), ids=("sampling-off", "sampling-on"))
798
+ def test_generator_constructs_model_owned_lane_configs_with_exact_decode_coverage(
799
+ monkeypatch,
800
+ device_sampling_enabled,
801
+ ):
802
+ executor_calls = []
803
+ monkeypatch.setattr(llama_generator, "_create_submeshes", lambda mesh, dp: [_Mesh(), _Mesh()])
804
+
805
+ def fake_from_pretrained(
806
+ mesh_device,
807
+ *,
808
+ hf_model: str | None = None,
809
+ instruct: bool | None = None,
810
+ max_batch_size: int,
811
+ max_seq_len: int,
812
+ optimizations="performance",
813
+ n_layers: int | None = None,
814
+ dtype=ttnn.bfloat8_b,
815
+ paged_attention_config=None,
816
+ ):
817
+ return SimpleNamespace(model=_model(max_batch_size=max_batch_size), runtime_config=_runtime_config())
818
+
819
+ def fake_build_executor(llm, config):
820
+ executor_calls.append((llm, config))
821
+ return _FakeLane(llm, config)
822
+
823
+ monkeypatch.setattr(llama_generator, "from_pretrained", fake_from_pretrained)
824
+ monkeypatch.setattr(llama_generator, "build_llama3_executor", fake_build_executor)
825
+ monkeypatch.setattr(llama_generator, "_model_kv_metadata", lambda model: ((ttnn.bfloat8_b,), 1, 8, 128))
826
+
827
+ generator = llama_generator.build_llama3_generator(
828
+ llama_generator.Llama3GeneratorConfig(
829
+ hf_model="meta-llama/Llama-3.1-8B-Instruct",
830
+ mesh_device=object(),
831
+ max_batch_size=4,
832
+ max_seq_len=4096,
833
+ n_layers=1,
834
+ tt_data_parallel=2,
835
+ trace_mode="all",
836
+ device_sampling_enabled=device_sampling_enabled,
837
+ )
838
+ )
839
+
840
+ assert isinstance(generator.target, LaneGroupExecutor)
841
+ assert len(executor_calls) == 2
842
+ assert all(isinstance(config, llama_executor.Llama3ExecutorConfig) for _, config in executor_calls)
843
+ assert all(config.warmup.include_decode_top_k is device_sampling_enabled for _, config in executor_calls)
844
+ for _, executor_config in executor_calls:
845
+ plan = _build_plan(
846
+ warmup=executor_config.warmup,
847
+ layout=PageTableLayout(block_size=32, raw_capacity_width=32, prefill_width=64, decode_width=32),
848
+ prefill_sequence_lengths=(128,),
849
+ lane_batch_size=2,
850
+ allow_force_argmax=True,
851
+ can_sample_on_device=device_sampling_enabled,
852
+ )
853
+ assert [case.sampling_path for case in plan.decode] == (
854
+ ["logits", "argmax", "topk"] if device_sampling_enabled else ["logits"]
855
+ )
856
+ assert isinstance(generator._adapter.config, llama_generator.VLLMAdapterConfig)
857
+ assert vars(generator._adapter) == {"config": generator._adapter.config}
858
+ assert generator._adapter.config.trace.mode == "all"
859
+ assert generator._adapter.config.expected_num_layers == 1
860
+ assert generator._adapter.config.expected_kv_heads_per_device == 8
861
+ assert generator._adapter.config.expected_head_dim == 128
862
+
863
+
864
+ @pytest.mark.parametrize(
865
+ ("sampling_mode", "sampling_params", "num_devices", "expected"),
866
+ [
867
+ ("on_device_topk", object(), 1, False),
868
+ ("on_device_topk", object(), 2, False),
869
+ ("on_device_topk", object(), 8, True),
870
+ ("on_device", object(), 8, False),
871
+ ("on_device_topk", None, 8, False),
872
+ ],
873
+ )
874
+ def test_direct_demo_forces_decode_top_k_only_on_t3k(sampling_mode, sampling_params, num_devices, expected):
875
+ assert llama_demo._force_decode_top_k(sampling_mode, sampling_params, num_devices) is expected
876
+
877
+
878
+ class _RecordingTarget:
879
+ model = SimpleNamespace(config=SimpleNamespace(max_batch_size=4))
880
+ model_args = object()
881
+ mesh_device = object()
882
+ cache_path = "cache"
883
+ already_warmed_up_prefill = False
884
+ eager_execution = object()
885
+ traced_prefill_execution = object()
886
+ traced_decode_execution = object()
887
+
888
+ def __init__(self, *, traceable_prefill=True):
889
+ self.calls = []
890
+ self.traceable_prefill = traceable_prefill
891
+ self._request_state_fields = ("prompt_tokens", "output_tokens", "slot_remap")
892
+
893
+ def _record(self, name, arguments):
894
+ arguments = {key: value for key, value in arguments.items() if key != "self"}
895
+ self.calls.append((name, (), arguments))
896
+ return name
897
+
898
+ def can_trace_prefill(
899
+ self,
900
+ *,
901
+ tokens: torch.Tensor, # ↓ Core request
902
+ prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata
903
+ start_pos: torch.Tensor | None = None,
904
+ empty_slots: Sequence[int] | None = None, # ↓ Lane routing
905
+ ) -> bool:
906
+ self._record("can_trace_prefill", locals())
907
+ return self.traceable_prefill
908
+
909
+ def prefill_forward(
910
+ self,
911
+ tokens: torch.Tensor,
912
+ page_table: torch.Tensor,
913
+ *,
914
+ prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata
915
+ start_pos: torch.Tensor | None = None,
916
+ empty_slots: Sequence[int] | None = None, # ↓ Lane routing
917
+ kv_cache: Any = None, # ↓ Borrowed resources
918
+ sampling_params: Any = None, # ↓ Sampling
919
+ prompt_tokens: Any = None,
920
+ output_tokens: Any = None,
921
+ slot_remap: Any = None,
922
+ execution: EagerExecutor | TracedExecutor | None = None, # ↓ Internal dispatch
923
+ ) -> str:
924
+ return self._record("prefill_forward", locals())
925
+
926
+ def decode_forward(
927
+ self,
928
+ tokens: torch.Tensor,
929
+ start_pos: torch.Tensor,
930
+ page_table: torch.Tensor,
931
+ *,
932
+ kv_cache: Any = None, # ↓ Borrowed resources
933
+ sampling_params: Any = None, # ↓ Sampling
934
+ prompt_tokens: Any = None,
935
+ output_tokens: Any = None,
936
+ slot_remap: Any = None,
937
+ reset_batch: bool = False, # ↓ State transition
938
+ read_from_device: bool = True, # ↓ Output policy
939
+ execution: EagerExecutor | TracedExecutor | None = None, # ↓ Internal dispatch
940
+ ) -> str:
941
+ return self._record("decode_forward", locals())
942
+
943
+ def cleanup(self) -> str:
944
+ self.calls.append(("cleanup", (), {}))
945
+ return "cleanup"
946
+
947
+
948
+ def _recording_generator(target):
949
+ target.model = _model()
950
+ target.config = _config("all")
951
+ return llama_generator.Llama3Generator(target, adapter=llama_generator._build_vllm_adapter(target))
952
+
953
+
954
+ def test_generator_delegates_without_concrete_type_checks():
955
+ target = _RecordingTarget()
956
+ generator = _recording_generator(target)
957
+ tokens = torch.tensor([[1]], dtype=torch.long)
958
+ start_pos = torch.tensor([0], dtype=torch.long)
959
+ page_table = torch.tensor([[0]], dtype=torch.int32)
960
+
961
+ assert generator.prefill_forward(tokens, page_table, enable_trace=True) == "prefill_forward"
962
+ assert generator.decode_forward(tokens, start_pos, page_table, enable_trace=True) == "decode_forward"
963
+ assert generator.cleanup() == "cleanup"
964
+ assert [name for name, _, _ in target.calls] == [
965
+ "prefill_forward",
966
+ "decode_forward",
967
+ "cleanup",
968
+ ]
969
+ assert target.calls[0][2]["execution"] is target.traced_prefill_execution
970
+ assert target.calls[1][2]["execution"] is target.traced_decode_execution
971
+
972
+
973
+ def test_generator_preserves_required_trace_intent_for_ineligible_prefill():
974
+ target = _RecordingTarget(traceable_prefill=False)
975
+ generator = _recording_generator(target)
976
+ tokens = torch.tensor([[1]], dtype=torch.long)
977
+ page_table = torch.tensor([[0]], dtype=torch.int32)
978
+
979
+ assert generator.prefill_forward(tokens, page_table, enable_trace=True) == "prefill_forward"
980
+ assert [name for name, _, _ in target.calls] == ["prefill_forward"]
981
+ assert target.calls[0][2]["execution"] is target.traced_prefill_execution
982
+
983
+
984
+ def test_generator_rejects_unavailable_traced_execution(expect_error):
985
+ target = _RecordingTarget()
986
+ target.traced_decode_execution = None
987
+ generator = _recording_generator(target)
988
+ tokens = torch.tensor([1], dtype=torch.long)
989
+ start_pos = torch.tensor([0], dtype=torch.long)
990
+ page_table = torch.tensor([[0]], dtype=torch.int32)
991
+
992
+ with expect_error(RuntimeError, "unavailable traced decode execution"):
993
+ generator.decode_forward(tokens, start_pos, page_table, enable_trace=True)
994
+
995
+ assert target.calls == []
996
+
997
+
998
+ def test_demo_uses_model_owned_config_and_order_independent_warmup(monkeypatch):
999
+ attention = SimpleNamespace(
1000
+ paged_attention_config=SimpleNamespace(block_size=32, max_num_blocks=128),
1001
+ kv_cache_dtype=ttnn.bfloat8_b,
1002
+ )
1003
+ llm = SimpleNamespace(
1004
+ model=SimpleNamespace(config=SimpleNamespace(block_configs=(SimpleNamespace(attention_config=attention),)))
1005
+ )
1006
+ captured = []
1007
+ monkeypatch.setattr(
1008
+ llama_demo, "build_llama3_executor", lambda product, config: captured.append(config) or object()
1009
+ )
1010
+
1011
+ llama_demo._build_demo_executor(llm, trace_mode="all", device_sampling_enabled=False)
1012
+ assert isinstance(captured[0], llama_executor.Llama3ExecutorConfig)
1013
+
1014
+ calls = []
1015
+
1016
+ def fake_warmup_model_prefill(
1017
+ *,
1018
+ kv_cache: Any, # ↓ Borrowed resources
1019
+ can_sample_on_device: bool, # ↓ Execution policy
1020
+ enable_trace: bool,
1021
+ ) -> None:
1022
+ calls.append(("prefill", kv_cache, can_sample_on_device, enable_trace))
1023
+
1024
+ def fake_warmup_model_decode(
1025
+ *,
1026
+ kv_cache: Any, # ↓ Borrowed resources
1027
+ max_batch_size: int, # ↓ Coverage dimensions
1028
+ num_blocks: int,
1029
+ can_sample_on_device: bool, # ↓ Execution policy
1030
+ enable_trace: bool,
1031
+ ) -> None:
1032
+ calls.append(("decode", kv_cache, max_batch_size, num_blocks, can_sample_on_device, enable_trace))
1033
+
1034
+ executor = SimpleNamespace(
1035
+ config=SimpleNamespace(
1036
+ trace=TraceConfig("all"),
1037
+ device_sampling_enabled=False,
1038
+ ),
1039
+ model=SimpleNamespace(config=SimpleNamespace(max_batch_size=4)),
1040
+ warmup_model_prefill=fake_warmup_model_prefill,
1041
+ warmup_model_decode=fake_warmup_model_decode,
1042
+ )
1043
+ kv_cache = object()
1044
+ llama_demo._warmup_demo_executor(executor, kv_cache=kv_cache, page_table=SimpleNamespace(shape=(4, 8)))
1045
+ assert calls == [
1046
+ ("prefill", kv_cache, False, False),
1047
+ ("decode", kv_cache, 4, 8, False, False),
1048
+ ("prefill", kv_cache, False, True),
1049
+ ("decode", kv_cache, 4, 8, False, True),
1050
+ ]
1051
+
1052
+ calls.clear()
1053
+ executor.config.device_sampling_enabled = True
1054
+ llama_demo._warmup_demo_executor(
1055
+ executor,
1056
+ kv_cache=kv_cache,
1057
+ page_table=SimpleNamespace(shape=(4, 8)),
1058
+ prefill_can_sample_on_device=False,
1059
+ )
1060
+ assert calls == [
1061
+ ("prefill", kv_cache, False, False),
1062
+ ("decode", kv_cache, 4, 8, True, False),
1063
+ ("prefill", kv_cache, False, True),
1064
+ ("decode", kv_cache, 4, 8, True, True),
1065
+ ]
code/models/common/tests/llm_runtime/test_llama3_8b_model_contract.py ADDED
@@ -0,0 +1,349 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+ from unittest.mock import MagicMock, call
6
+
7
+ import pytest
8
+
9
+ from models.common.models.llama3_8b import model as llama_model
10
+ from models.common.models.llama3_8b.model import Llama3Transformer1D, Llama3Transformer1DConfig, TransformerBlock1D
11
+
12
+
13
+ def _layer():
14
+ return SimpleNamespace(
15
+ attention_norm=object(),
16
+ attention=SimpleNamespace(config=SimpleNamespace(kv_cache=None), kv_cache=None),
17
+ ff_norm=object(),
18
+ feed_forward=object(),
19
+ )
20
+
21
+
22
+ def test_transformer_config_exposes_hidden_width_for_runtime_slices():
23
+ config = Llama3Transformer1DConfig(
24
+ n_layers=1,
25
+ vocab_size=128256,
26
+ max_batch_size=1,
27
+ max_seq_len=1024,
28
+ dim=4096,
29
+ num_devices=1,
30
+ mesh_device=object(),
31
+ embedding_config=object(),
32
+ rope_config=object(),
33
+ block_configs=[object()],
34
+ norm_config=object(),
35
+ lm_head_config=object(),
36
+ )
37
+
38
+ assert config.dim == 4096
39
+
40
+
41
+ def test_iter_executor_named_modules_preserves_names_and_order():
42
+ layers = [_layer(), _layer()]
43
+ model = SimpleNamespace(layers=layers, norm=object(), lm_head=object())
44
+
45
+ named_modules = list(Llama3Transformer1D.iter_executor_named_modules(model))
46
+
47
+ assert named_modules == [
48
+ ("layer[0].attn_norm", layers[0].attention_norm),
49
+ ("layer[0].attention", layers[0].attention),
50
+ ("layer[0].ff_norm", layers[0].ff_norm),
51
+ ("layer[0].mlp", layers[0].feed_forward),
52
+ ("layer[1].attn_norm", layers[1].attention_norm),
53
+ ("layer[1].attention", layers[1].attention),
54
+ ("layer[1].ff_norm", layers[1].ff_norm),
55
+ ("layer[1].mlp", layers[1].feed_forward),
56
+ ("final_norm", model.norm),
57
+ ("lm_head", model.lm_head),
58
+ ]
59
+
60
+
61
+ def test_iter_executor_named_modules_without_layers_yields_nothing():
62
+ model = SimpleNamespace(norm=object(), lm_head=object())
63
+
64
+ assert list(Llama3Transformer1D.iter_executor_named_modules(model)) == []
65
+
66
+
67
+ def test_set_kv_cache_binds_and_unbinds_config_and_runtime_references():
68
+ layers = [_layer(), _layer()]
69
+ model = SimpleNamespace(layers=layers)
70
+ kv_cache = [[object(), object()], [object(), object()]]
71
+
72
+ Llama3Transformer1D.set_kv_cache(model, kv_cache)
73
+
74
+ for layer, expected in zip(layers, kv_cache):
75
+ bound = layer.attention.config.kv_cache
76
+ assert bound == tuple(expected)
77
+ assert bound[0] is expected[0]
78
+ assert bound[1] is expected[1]
79
+ assert layer.attention.kv_cache is bound
80
+
81
+ Llama3Transformer1D.set_kv_cache(model, None)
82
+ Llama3Transformer1D.set_kv_cache(model, None)
83
+
84
+ for layer in layers:
85
+ assert layer.attention.config.kv_cache is None
86
+ assert layer.attention.kv_cache is None
87
+
88
+
89
+ def test_set_kv_cache_rejects_wrong_layer_count_before_binding(expect_error):
90
+ layers = [_layer(), _layer()]
91
+ model = SimpleNamespace(layers=layers)
92
+
93
+ with expect_error(ValueError, "model has 2 layers"):
94
+ Llama3Transformer1D.set_kv_cache(model, [[object(), object()]])
95
+
96
+ assert all(layer.attention.config.kv_cache is None for layer in layers)
97
+ assert all(layer.attention.kv_cache is None for layer in layers)
98
+
99
+
100
+ @pytest.mark.parametrize("bad_pair", [[object()], object()])
101
+ def test_set_kv_cache_validates_every_pair_before_binding(bad_pair, expect_error):
102
+ layers = [_layer(), _layer()]
103
+ model = SimpleNamespace(layers=layers)
104
+
105
+ with expect_error((TypeError, ValueError), "layer 1.*K/V tensor"):
106
+ Llama3Transformer1D.set_kv_cache(
107
+ model,
108
+ [[object(), object()], bad_pair],
109
+ )
110
+
111
+ assert all(layer.attention.config.kv_cache is None for layer in layers)
112
+ assert all(layer.attention.kv_cache is None for layer in layers)
113
+
114
+
115
+ class _FakeTensor:
116
+ def __init__(self, shape=(1, 1, 96, 4096), dtype=None):
117
+ self.shape = shape
118
+ self.dtype = dtype
119
+ self.deallocate = MagicMock()
120
+
121
+
122
+ def _identity_all_gather(_norm, tensor, *, memory_config=None):
123
+ return tensor
124
+
125
+
126
+ def _model_ttnn(
127
+ *, slice_result=None, split_result=None, typecast_result=None, embedding_result=None, unsqueeze_result=None
128
+ ):
129
+ slice_mock = MagicMock(return_value=slice_result)
130
+ return SimpleNamespace(
131
+ DRAM_MEMORY_CONFIG=object(),
132
+ TILE_LAYOUT=object(),
133
+ bfloat16=object(),
134
+ add=MagicMock(side_effect=lambda lhs, _rhs, **_kwargs: lhs),
135
+ concat=MagicMock(return_value=slice_result),
136
+ deallocate=MagicMock(),
137
+ embedding=MagicMock(return_value=embedding_result),
138
+ interleaved_to_sharded=MagicMock(side_effect=lambda tensor, _memcfg: tensor),
139
+ reshape=MagicMock(side_effect=lambda tensor, *_args, **_kwargs: tensor),
140
+ slice=slice_mock,
141
+ split=MagicMock(return_value=split_result),
142
+ to_memory_config=MagicMock(side_effect=lambda tensor, *_args, **_kwargs: tensor),
143
+ typecast=MagicMock(return_value=typecast_result),
144
+ unsqueeze_to_4D=MagicMock(return_value=unsqueeze_result),
145
+ )
146
+
147
+
148
+ @pytest.mark.parametrize("chunk_start_idx_tensor", [None, object()])
149
+ def test_transformer_block_prefill_forwards_scalar_and_tensor_chunk_start(monkeypatch, chunk_start_idx_tensor):
150
+ x = _FakeTensor()
151
+ attn_output = _FakeTensor()
152
+ attention = SimpleNamespace(prefill_forward=MagicMock(return_value=attn_output))
153
+ block = SimpleNamespace(
154
+ activation_dtype=None,
155
+ attention=attention,
156
+ attention_norm=SimpleNamespace(prefill_forward=MagicMock(return_value=x)),
157
+ feed_forward=SimpleNamespace(prefill_forward=MagicMock(side_effect=lambda tensor: tensor)),
158
+ ff_norm=SimpleNamespace(prefill_forward=MagicMock(side_effect=lambda tensor: tensor)),
159
+ prefill_residual_memcfg=object(),
160
+ )
161
+ fake_ttnn = _model_ttnn()
162
+ monkeypatch.setattr(llama_model, "ttnn", fake_ttnn)
163
+ monkeypatch.setattr(llama_model, "_all_gather_rmsnorm_tensor", _identity_all_gather)
164
+
165
+ TransformerBlock1D.prefill_forward(
166
+ block,
167
+ x,
168
+ ("cos", "sin"),
169
+ user_id=3,
170
+ page_table="page-table",
171
+ chunk_page_table="chunk-page-table",
172
+ chunk_start_idx=96,
173
+ chunk_start_idx_tensor=chunk_start_idx_tensor,
174
+ )
175
+
176
+ attention.prefill_forward.assert_called_once_with(
177
+ x,
178
+ ("cos", "sin"),
179
+ user_id=3,
180
+ page_table="page-table",
181
+ chunk_page_table="chunk-page-table",
182
+ chunk_start_idx=96,
183
+ chunk_start_idx_tensor=chunk_start_idx_tensor,
184
+ )
185
+
186
+
187
+ def test_prefill_uses_runtime_block_slice_then_embeds_exact_row(monkeypatch):
188
+ hidden_states = _FakeTensor(shape=(1, 1, 96, 4096))
189
+ sliced = _FakeTensor(shape=(1, 1, 32, 4096))
190
+ converted = _FakeTensor(shape=(1, 1, 32, 4096))
191
+ embedded = _FakeTensor(shape=(1, 1, 4096))
192
+ selected = _FakeTensor(shape=(1, 1, 1, 4096))
193
+ start_tensor = object()
194
+ end_tensor = object()
195
+ row_index = object()
196
+ fake_ttnn = _model_ttnn(
197
+ slice_result=sliced,
198
+ typecast_result=converted,
199
+ embedding_result=embedded,
200
+ unsqueeze_result=selected,
201
+ )
202
+ model = SimpleNamespace(
203
+ activation_dtypes=[],
204
+ layers=[],
205
+ lm_head=SimpleNamespace(
206
+ config=SimpleNamespace(input_memcfg=None),
207
+ forward=MagicMock(side_effect=lambda tensor: tensor),
208
+ ),
209
+ norm=SimpleNamespace(prefill_forward=MagicMock(side_effect=lambda tensor: tensor)),
210
+ )
211
+ monkeypatch.setattr(llama_model, "ttnn", fake_ttnn)
212
+ monkeypatch.setattr(llama_model, "_all_gather_rmsnorm_tensor", _identity_all_gather)
213
+
214
+ Llama3Transformer1D.prefill_forward(
215
+ model,
216
+ hidden_states,
217
+ ("cos", "sin"),
218
+ get_last_token=95,
219
+ last_token_slice=(start_tensor, end_tensor),
220
+ last_token_index=row_index,
221
+ )
222
+
223
+ fake_ttnn.slice.assert_called_once_with(
224
+ hidden_states,
225
+ start_tensor,
226
+ end_tensor,
227
+ slice_dim=2,
228
+ num_devices=3,
229
+ )
230
+ fake_ttnn.typecast.assert_called_once_with(sliced, fake_ttnn.bfloat16)
231
+ fake_ttnn.embedding.assert_called_once_with(row_index, converted, layout=fake_ttnn.TILE_LAYOUT)
232
+ fake_ttnn.unsqueeze_to_4D.assert_called_once_with(embedded)
233
+ assert fake_ttnn.deallocate.call_args_list == [call(hidden_states), call(sliced), call(converted)]
234
+ model.norm.prefill_forward.assert_called_once_with(selected)
235
+
236
+
237
+ def test_post_process_prefill_uses_runtime_bounds_for_aligned_block_slice(monkeypatch):
238
+ hidden_states = _FakeTensor(shape=(1, 1, 96, 4096))
239
+ sliced = _FakeTensor(shape=(1, 1, 32, 4096))
240
+ start_tensor = object()
241
+ end_tensor = object()
242
+ fake_ttnn = _model_ttnn(slice_result=sliced)
243
+ model = SimpleNamespace(
244
+ lm_head=SimpleNamespace(
245
+ config=SimpleNamespace(input_memcfg=None),
246
+ forward=MagicMock(side_effect=lambda tensor: tensor),
247
+ ),
248
+ norm=SimpleNamespace(prefill_forward=MagicMock(side_effect=lambda tensor: tensor)),
249
+ )
250
+ monkeypatch.setattr(llama_model, "ttnn", fake_ttnn)
251
+ monkeypatch.setattr(llama_model, "_all_gather_rmsnorm_tensor", _identity_all_gather)
252
+
253
+ Llama3Transformer1D.post_process_prefill_output(
254
+ model,
255
+ hidden_states,
256
+ last_token_idx=95,
257
+ last_token_slice=(start_tensor, end_tensor),
258
+ )
259
+
260
+ fake_ttnn.slice.assert_called_once_with(
261
+ hidden_states,
262
+ start_tensor,
263
+ end_tensor,
264
+ slice_dim=2,
265
+ num_devices=3,
266
+ )
267
+
268
+
269
+ def test_batched_post_process_uses_runtime_bounds_for_each_active_slot(monkeypatch):
270
+ hidden_states = _FakeTensor(shape=(1, 1, 256, 4096))
271
+ user_states = [_FakeTensor(shape=(1, 1, 128, 4096)) for _ in range(2)]
272
+ block = _FakeTensor(shape=(1, 1, 32, 4096))
273
+ embedded = _FakeTensor(shape=(1, 1, 4096))
274
+ selected = _FakeTensor(shape=(1, 1, 1, 4096))
275
+ start_tensor = object()
276
+ end_tensor = object()
277
+ row_index = object()
278
+ fake_ttnn = _model_ttnn(
279
+ slice_result=block,
280
+ split_result=user_states,
281
+ embedding_result=embedded,
282
+ unsqueeze_result=selected,
283
+ )
284
+ model = SimpleNamespace(
285
+ lm_head=SimpleNamespace(
286
+ config=SimpleNamespace(input_memcfg=None),
287
+ forward=MagicMock(side_effect=lambda tensor: tensor),
288
+ ),
289
+ norm=SimpleNamespace(prefill_forward=MagicMock(side_effect=lambda tensor: tensor)),
290
+ )
291
+ monkeypatch.setattr(llama_model, "ttnn", fake_ttnn)
292
+ monkeypatch.setattr(llama_model, "_all_gather_rmsnorm_tensor", _identity_all_gather)
293
+
294
+ Llama3Transformer1D.post_process_batched_prefill_output(
295
+ model,
296
+ hidden_states,
297
+ last_token_idx_list=[4, 17],
298
+ padded_batch=2,
299
+ prefill_seq_len=128,
300
+ last_token_slice=(start_tensor, end_tensor),
301
+ last_token_index=row_index,
302
+ )
303
+
304
+ assert fake_ttnn.slice.call_args_list == [
305
+ call(user_states[0], start_tensor, end_tensor, slice_dim=2, num_devices=4),
306
+ call(user_states[1], start_tensor, end_tensor, slice_dim=2, num_devices=4),
307
+ ]
308
+ assert fake_ttnn.embedding.call_args_list == [
309
+ call(row_index, block, layout=fake_ttnn.TILE_LAYOUT),
310
+ call(row_index, block, layout=fake_ttnn.TILE_LAYOUT),
311
+ ]
312
+ assert fake_ttnn.unsqueeze_to_4D.call_args_list == [call(embedded), call(embedded)]
313
+ assert fake_ttnn.deallocate.call_args_list == [call(block), call(block)]
314
+
315
+
316
+ def test_prepare_prefill_rot_mats_gathers_runtime_positions_without_slicing(monkeypatch):
317
+ position_indices = object()
318
+ cos_matrix = object()
319
+ sin_matrix = object()
320
+ cos_rows = object()
321
+ sin_rows = object()
322
+ cos_4d = object()
323
+ sin_4d = object()
324
+ embedding = MagicMock(side_effect=[cos_rows, sin_rows])
325
+ unsqueeze_to_4d = MagicMock(side_effect=[cos_4d, sin_4d])
326
+ fake_ttnn = SimpleNamespace(
327
+ TILE_LAYOUT=object(),
328
+ embedding=embedding,
329
+ unsqueeze_to_4D=unsqueeze_to_4d,
330
+ )
331
+ rope_setup = SimpleNamespace(
332
+ cos_matrix=cos_matrix,
333
+ sin_matrix=sin_matrix,
334
+ load_device_weights=MagicMock(),
335
+ )
336
+ monkeypatch.setattr(llama_model, "ttnn", fake_ttnn)
337
+
338
+ result = Llama3Transformer1D.prepare_prefill_rot_mats(
339
+ SimpleNamespace(rope_setup=rope_setup),
340
+ position_indices,
341
+ )
342
+
343
+ rope_setup.load_device_weights.assert_called_once_with()
344
+ assert embedding.call_args_list == [
345
+ call(position_indices, cos_matrix, layout=fake_ttnn.TILE_LAYOUT),
346
+ call(position_indices, sin_matrix, layout=fake_ttnn.TILE_LAYOUT),
347
+ ]
348
+ assert unsqueeze_to_4d.call_args_list == [call(cos_rows), call(sin_rows)]
349
+ assert result == (cos_4d, sin_4d)
code/models/common/tests/llm_runtime/test_model_contract.py ADDED
@@ -0,0 +1,972 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Generic host-only tensor-model contract for migrated sibling models."""
5
+
6
+ import inspect
7
+ from dataclasses import dataclass, fields
8
+ from types import SimpleNamespace
9
+ from unittest.mock import MagicMock, call
10
+
11
+ import pytest
12
+
13
+ from models.common.models.deepseek_r1_distill_qwen_14b import model as deepseek_model
14
+ from models.common.models.llama32_1b import model as llama32_model
15
+ from models.common.models.llama32_3b import model as llama32_3b_model
16
+ from models.common.models.llama33_70b import model as llama33_70b_model
17
+ from models.common.models.mistral_7b import model as mistral_model
18
+ from models.common.models.phi4 import model as phi4_model
19
+ from models.common.models.qwen2_7b import model as qwen2_model
20
+ from models.common.models.qwen3_32b import model as qwen3_32b_model
21
+ from models.common.models.qwen25_7b import model as qwen25_model
22
+ from models.common.models.qwen25_72b import model as qwen25_72b_model
23
+ from models.common.models.qwen25_coder_32b import model as qwen25_coder_32b_model
24
+
25
+ MODEL_CONTRACTS = {
26
+ "llama32_1b": SimpleNamespace(
27
+ module=llama32_model,
28
+ model_class=llama32_model.Llama32_1BTransformer1D,
29
+ config_class=llama32_model.Llama32_1BTransformer1DConfig,
30
+ attention_config_class=llama32_model.Attention1DConfig,
31
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
32
+ make_config=lambda **kwargs: _make_llama32_config(**kwargs),
33
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
34
+ construct_model=lambda monkeypatch: _construct_llama32_model(monkeypatch),
35
+ expected_module_names=(
36
+ "layer[0].attn_norm",
37
+ "layer[0].attention",
38
+ "layer[0].ff_norm",
39
+ "layer[0].mlp",
40
+ "layer[1].attn_norm",
41
+ "layer[1].attention",
42
+ "layer[1].ff_norm",
43
+ "layer[1].mlp",
44
+ "final_norm",
45
+ "lm_head",
46
+ ),
47
+ ),
48
+ "llama32_3b": SimpleNamespace(
49
+ module=llama32_3b_model,
50
+ model_class=llama32_3b_model.Llama32_3BTransformer1D,
51
+ config_class=llama32_3b_model.Llama32_3BTransformer1DConfig,
52
+ attention_config_class=llama32_3b_model.Attention1DConfig,
53
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
54
+ make_config=lambda **kwargs: _make_llama32_config(module=llama32_3b_model, **kwargs),
55
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
56
+ construct_model=lambda monkeypatch: _construct_llama32_model(monkeypatch, module=llama32_3b_model),
57
+ expected_module_names=(
58
+ "layer[0].attn_norm",
59
+ "layer[0].attention",
60
+ "layer[0].ff_norm",
61
+ "layer[0].mlp",
62
+ "layer[1].attn_norm",
63
+ "layer[1].attention",
64
+ "layer[1].ff_norm",
65
+ "layer[1].mlp",
66
+ "final_norm",
67
+ "lm_head",
68
+ ),
69
+ ),
70
+ "llama33_70b": SimpleNamespace(
71
+ module=llama33_70b_model,
72
+ model_class=llama33_70b_model.Llama33_70BTransformer1D,
73
+ config_class=llama33_70b_model.Llama33_70BTransformer1DConfig,
74
+ attention_config_class=llama33_70b_model.Attention1DConfig,
75
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
76
+ make_config=lambda **kwargs: _make_llama32_config(module=llama33_70b_model, **kwargs),
77
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
78
+ construct_model=lambda monkeypatch: _construct_llama32_model(monkeypatch, module=llama33_70b_model),
79
+ expected_module_names=(
80
+ "layer[0].attn_norm",
81
+ "layer[0].attention",
82
+ "layer[0].ff_norm",
83
+ "layer[0].mlp",
84
+ "layer[1].attn_norm",
85
+ "layer[1].attention",
86
+ "layer[1].ff_norm",
87
+ "layer[1].mlp",
88
+ "final_norm",
89
+ "lm_head",
90
+ ),
91
+ ),
92
+ "qwen2_7b": SimpleNamespace(
93
+ module=qwen2_model,
94
+ model_class=qwen2_model.Qwen2_7B,
95
+ config_class=qwen2_model.Qwen2_7BTransformerConfig,
96
+ attention_config_class=qwen2_model.Attention1DConfig,
97
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
98
+ make_config=lambda **kwargs: _make_qwen2_config(**kwargs),
99
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
100
+ construct_model=lambda monkeypatch: _construct_qwen2_model(monkeypatch),
101
+ expected_module_names=(
102
+ "layer[0].attn_norm",
103
+ "layer[0].attention",
104
+ "layer[0].ff_norm",
105
+ "layer[0].mlp",
106
+ "layer[1].attn_norm",
107
+ "layer[1].attention",
108
+ "layer[1].ff_norm",
109
+ "layer[1].mlp",
110
+ "final_norm",
111
+ "lm_head",
112
+ ),
113
+ ),
114
+ "qwen25_7b": SimpleNamespace(
115
+ module=qwen25_model,
116
+ model_class=qwen25_model.Qwen25_7B,
117
+ config_class=qwen25_model.Qwen25_7BTransformerConfig,
118
+ attention_config_class=qwen25_model.Attention1DConfig,
119
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
120
+ make_config=lambda **kwargs: _make_qwen2_config(module=qwen25_model, **kwargs),
121
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
122
+ construct_model=lambda monkeypatch: _construct_qwen2_model(monkeypatch, module=qwen25_model),
123
+ expected_module_names=(
124
+ "layer[0].attn_norm",
125
+ "layer[0].attention",
126
+ "layer[0].ff_norm",
127
+ "layer[0].mlp",
128
+ "layer[1].attn_norm",
129
+ "layer[1].attention",
130
+ "layer[1].ff_norm",
131
+ "layer[1].mlp",
132
+ "final_norm",
133
+ "lm_head",
134
+ ),
135
+ ),
136
+ "qwen25_72b": SimpleNamespace(
137
+ module=qwen25_72b_model,
138
+ model_class=qwen25_72b_model.Qwen25_72B,
139
+ config_class=qwen25_72b_model.Qwen25_72BConfig,
140
+ attention_config_class=qwen25_72b_model.Attention1DConfig,
141
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
142
+ make_config=lambda **kwargs: _make_qwen25_72b_config(**kwargs),
143
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
144
+ construct_model=lambda monkeypatch: _construct_qwen25_72b_model(monkeypatch),
145
+ expected_module_names=(
146
+ "layer[0].attn_norm",
147
+ "layer[0].attention",
148
+ "layer[0].ff_norm",
149
+ "layer[0].mlp",
150
+ "layer[1].attn_norm",
151
+ "layer[1].attention",
152
+ "layer[1].ff_norm",
153
+ "layer[1].mlp",
154
+ "final_norm",
155
+ "lm_head",
156
+ ),
157
+ ),
158
+ "qwen25_coder_32b": SimpleNamespace(
159
+ module=qwen25_coder_32b_model,
160
+ model_class=qwen25_coder_32b_model.Qwen25Coder32B,
161
+ config_class=qwen25_coder_32b_model.Qwen25Coder32BConfig,
162
+ attention_config_class=qwen25_coder_32b_model.Attention1DConfig,
163
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
164
+ make_config=lambda **kwargs: _make_qwen25_coder_32b_config(**kwargs),
165
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
166
+ construct_model=lambda monkeypatch: _construct_qwen25_coder_32b_model(monkeypatch),
167
+ expected_module_names=(
168
+ "layer[0].attn_norm",
169
+ "layer[0].attention",
170
+ "layer[0].ff_norm",
171
+ "layer[0].mlp",
172
+ "layer[1].attn_norm",
173
+ "layer[1].attention",
174
+ "layer[1].ff_norm",
175
+ "layer[1].mlp",
176
+ "final_norm",
177
+ "lm_head",
178
+ ),
179
+ ),
180
+ "qwen3_32b": SimpleNamespace(
181
+ module=qwen3_32b_model,
182
+ model_class=qwen3_32b_model.Qwen3_32B,
183
+ config_class=qwen3_32b_model.Qwen3_32BConfig,
184
+ attention_config_class=qwen3_32b_model.Attention1DConfig,
185
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
186
+ make_config=lambda **kwargs: _make_qwen3_32b_config(**kwargs),
187
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
188
+ construct_model=lambda monkeypatch: _construct_qwen3_32b_model(monkeypatch),
189
+ expected_module_names=(
190
+ "layer[0].attn_norm",
191
+ "layer[0].attention",
192
+ "layer[0].ff_norm",
193
+ "layer[0].mlp",
194
+ "layer[1].attn_norm",
195
+ "layer[1].attention",
196
+ "layer[1].ff_norm",
197
+ "layer[1].mlp",
198
+ "final_norm",
199
+ "lm_head",
200
+ ),
201
+ ),
202
+ "deepseek_r1_distill_qwen_14b": SimpleNamespace(
203
+ module=deepseek_model,
204
+ model_class=deepseek_model.DeepSeekR1Qwen14B,
205
+ config_class=deepseek_model.DeepSeekR1Qwen14BTransformerConfig,
206
+ attention_config_class=deepseek_model.Attention1DConfig,
207
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
208
+ make_config=lambda **kwargs: _make_qwen2_config(module=deepseek_model, **kwargs),
209
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
210
+ construct_model=lambda monkeypatch: _construct_qwen2_model(monkeypatch, module=deepseek_model),
211
+ expected_module_names=(
212
+ "layer[0].attn_norm",
213
+ "layer[0].attention",
214
+ "layer[0].ff_norm",
215
+ "layer[0].mlp",
216
+ "layer[1].attn_norm",
217
+ "layer[1].attention",
218
+ "layer[1].ff_norm",
219
+ "layer[1].mlp",
220
+ "final_norm",
221
+ "lm_head",
222
+ ),
223
+ ),
224
+ "mistral_7b": SimpleNamespace(
225
+ module=mistral_model,
226
+ model_class=mistral_model.Mistral7B,
227
+ config_class=mistral_model.Mistral7BTransformerConfig,
228
+ attention_config_class=mistral_model.Attention1DConfig,
229
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
230
+ make_config=lambda **kwargs: _make_mistral_config(**kwargs),
231
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
232
+ construct_model=lambda monkeypatch: _construct_mistral_model(monkeypatch),
233
+ expected_module_names=(
234
+ "layer[0].attn_norm",
235
+ "layer[0].attention",
236
+ "layer[0].ff_norm",
237
+ "layer[0].mlp",
238
+ "layer[1].attn_norm",
239
+ "layer[1].attention",
240
+ "layer[1].ff_norm",
241
+ "layer[1].mlp",
242
+ "final_norm",
243
+ "lm_head",
244
+ ),
245
+ ),
246
+ "phi4": SimpleNamespace(
247
+ module=phi4_model,
248
+ model_class=phi4_model.Phi4Transformer,
249
+ config_class=phi4_model.Phi4TransformerConfig,
250
+ attention_config_class=phi4_model.Attention1DConfig,
251
+ make_attention_config=lambda **kwargs: _make_llama32_attention_config(**kwargs),
252
+ make_config=lambda **kwargs: _make_phi4_config(**kwargs),
253
+ make_layer=lambda attention_config=None: _make_llama32_layer(attention_config),
254
+ construct_model=lambda monkeypatch: _construct_phi4_model(monkeypatch),
255
+ expected_module_names=(
256
+ "layer[0].attn_norm",
257
+ "layer[0].attention",
258
+ "layer[0].ff_norm",
259
+ "layer[0].mlp",
260
+ "layer[1].attn_norm",
261
+ "layer[1].attention",
262
+ "layer[1].ff_norm",
263
+ "layer[1].mlp",
264
+ "final_norm",
265
+ "lm_head",
266
+ ),
267
+ ),
268
+ }
269
+
270
+
271
+ @pytest.fixture(params=MODEL_CONTRACTS.items(), ids=lambda item: item[0])
272
+ def contract(request):
273
+ return request.param[1]
274
+
275
+
276
+ @dataclass(frozen=True)
277
+ class _PagedAttentionConfig:
278
+ block_size: int
279
+ max_num_blocks: int
280
+
281
+
282
+ def _make_llama32_attention_config(*, n_kv_heads=8, kv_cache=None):
283
+ return SimpleNamespace(
284
+ n_kv_heads=n_kv_heads,
285
+ use_vllm_paged_kv_cache=True,
286
+ paged_attention_config=_PagedAttentionConfig(block_size=32, max_num_blocks=128),
287
+ kv_cache=kv_cache,
288
+ )
289
+
290
+
291
+ def _make_llama32_layer(attention_config=None):
292
+ attention_config = attention_config or _make_llama32_attention_config()
293
+ attention = SimpleNamespace(config=attention_config, kv_cache=attention_config.kv_cache)
294
+ return SimpleNamespace(
295
+ attention_norm=object(),
296
+ attention=attention,
297
+ self_attn=attention,
298
+ ff_norm=object(),
299
+ feed_forward=object(),
300
+ )
301
+
302
+
303
+ def _make_llama32_config(*, module=llama32_model, n_layers=1, num_devices=2, n_kv_heads=8, sampling_config=None):
304
+ block_configs = [
305
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
306
+ ]
307
+ config_class = next(
308
+ getattr(module, name)
309
+ for name in (
310
+ "Llama32_1BTransformer1DConfig",
311
+ "Llama32_3BTransformer1DConfig",
312
+ "Llama33_70BTransformer1DConfig",
313
+ )
314
+ if hasattr(module, name)
315
+ )
316
+ return config_class(
317
+ n_layers=n_layers,
318
+ vocab_size=128256,
319
+ max_batch_size=4,
320
+ max_seq_len=4096,
321
+ dim=2048,
322
+ num_devices=num_devices,
323
+ mesh_device=SimpleNamespace(get_num_devices=lambda: num_devices),
324
+ embedding_config=object(),
325
+ rope_config=object(),
326
+ block_configs=block_configs,
327
+ norm_config=object(),
328
+ lm_head_config=object(),
329
+ sampling_config=sampling_config,
330
+ )
331
+
332
+
333
+ def _construct_llama32_model(monkeypatch, *, module=llama32_model):
334
+ sentinels = {
335
+ "embedding": object(),
336
+ "rope_setup": object(),
337
+ "layer": _make_llama32_layer(),
338
+ "norm": object(),
339
+ "lm_head": object(),
340
+ "sampling": object(),
341
+ }
342
+ for owner_name, sentinel_name in (
343
+ ("Embedding1D", "embedding"),
344
+ ("RotarySetup1D", "rope_setup"),
345
+ ("TransformerBlock1D", "layer"),
346
+ ("RMSNorm1D", "norm"),
347
+ ("LMHead1D", "lm_head"),
348
+ ("Sampling1D", "sampling"),
349
+ ):
350
+ monkeypatch.setattr(
351
+ getattr(module, owner_name),
352
+ "from_config",
353
+ MagicMock(return_value=sentinels[sentinel_name]),
354
+ )
355
+ config = _make_llama32_config(module=module, num_devices=1, sampling_config=object())
356
+ model_class = next(
357
+ getattr(module, name)
358
+ for name in (
359
+ "Llama32_1BTransformer1D",
360
+ "Llama32_3BTransformer1D",
361
+ "Llama33_70BTransformer1D",
362
+ )
363
+ if hasattr(module, name)
364
+ )
365
+ return model_class(config), config, sentinels
366
+
367
+
368
+ def _make_qwen2_config(*, module=qwen2_model, n_layers=1, num_devices=2, n_kv_heads=4, sampling_config=None):
369
+ block_configs = [
370
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
371
+ ]
372
+ config_class = (
373
+ getattr(module, "Qwen2_7BTransformerConfig", None)
374
+ or getattr(module, "Qwen25_7BTransformerConfig", None)
375
+ or module.DeepSeekR1Qwen14BTransformerConfig
376
+ )
377
+ return config_class(
378
+ n_layers=n_layers,
379
+ vocab_size=152064,
380
+ max_batch_size=4,
381
+ max_seq_len=4096,
382
+ dim=3584,
383
+ num_devices=num_devices,
384
+ mesh_device=SimpleNamespace(get_num_devices=lambda: num_devices),
385
+ embedding_config=object(),
386
+ rope_config=object(),
387
+ block_configs=block_configs,
388
+ norm_config=object(),
389
+ lm_head_config=object(),
390
+ sampling_config=sampling_config,
391
+ tt_ccl=object(),
392
+ )
393
+
394
+
395
+ def _construct_qwen2_model(monkeypatch, *, module=qwen2_model):
396
+ sentinels = {
397
+ "embedding": object(),
398
+ "rope_setup": object(),
399
+ "layer": _make_llama32_layer(),
400
+ "norm": object(),
401
+ "lm_head": object(),
402
+ "sampling": object(),
403
+ }
404
+ layer_class = (
405
+ "Qwen2_7BDecoderLayer"
406
+ if module is qwen2_model
407
+ else ("Qwen25_7BDecoderLayer" if module is qwen25_model else "DeepSeekR1Qwen14BDecoderLayer")
408
+ )
409
+ for owner_name, sentinel_name in (
410
+ ("Embedding1D", "embedding"),
411
+ ("RotarySetup1D", "rope_setup"),
412
+ (layer_class, "layer"),
413
+ ("RMSNorm1D", "norm"),
414
+ ("LMHead1D", "lm_head"),
415
+ ("Sampling1D", "sampling"),
416
+ ):
417
+ monkeypatch.setattr(
418
+ getattr(module, owner_name),
419
+ "from_config",
420
+ MagicMock(return_value=sentinels[sentinel_name]),
421
+ )
422
+ config = _make_qwen2_config(module=module, sampling_config=object())
423
+ model_class = getattr(module, "Qwen2_7B", None) or getattr(module, "Qwen25_7B", None) or module.DeepSeekR1Qwen14B
424
+ return model_class(config), config, sentinels
425
+
426
+
427
+ def _make_qwen25_72b_config(*, n_layers=1, num_devices=8, n_kv_heads=8, sampling_config=None):
428
+ del sampling_config
429
+ block_configs = [
430
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
431
+ ]
432
+ return qwen25_72b_model.Qwen25_72BConfig(
433
+ hf_model_id="Qwen/Qwen2.5-72B-Instruct",
434
+ dim=8192,
435
+ n_heads=64,
436
+ n_kv_heads=n_kv_heads,
437
+ head_dim=128,
438
+ hidden_dim=29568,
439
+ vocab_size=152064,
440
+ rms_norm_eps=1e-6,
441
+ rope_theta=1_000_000.0,
442
+ num_hidden_layers=n_layers,
443
+ max_batch_size=4,
444
+ max_seq_len=4096,
445
+ rope_table_len=8192,
446
+ num_devices=num_devices,
447
+ mesh_device=SimpleNamespace(get_num_devices=lambda: num_devices),
448
+ block_configs=block_configs,
449
+ )
450
+
451
+
452
+ def _construct_qwen25_72b_model(monkeypatch):
453
+ sentinels = {
454
+ "embedding": object(),
455
+ "rope_setup": object(),
456
+ "layer": _make_llama32_layer(),
457
+ "norm": object(),
458
+ "lm_head": object(),
459
+ "sampling": object(),
460
+ }
461
+ monkeypatch.setattr(qwen25_72b_model, "get_tt_ccl", MagicMock(return_value=object()))
462
+ monkeypatch.setattr(qwen25_72b_model, "Sampling1D", MagicMock(return_value=sentinels["sampling"]))
463
+ config = _make_qwen25_72b_config()
464
+ return (
465
+ qwen25_72b_model.Qwen25_72B(
466
+ config,
467
+ sentinels["embedding"],
468
+ sentinels["rope_setup"],
469
+ [sentinels["layer"]],
470
+ sentinels["norm"],
471
+ sentinels["lm_head"],
472
+ config.mesh_device,
473
+ ),
474
+ config,
475
+ sentinels,
476
+ )
477
+
478
+
479
+ def _make_qwen25_coder_32b_config(*, n_layers=1, num_devices=8, n_kv_heads=8, sampling_config=None):
480
+ del sampling_config
481
+ block_configs = [
482
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
483
+ ]
484
+ return qwen25_coder_32b_model.Qwen25Coder32BConfig(
485
+ hf_model_id="Qwen/Qwen2.5-Coder-32B-Instruct",
486
+ dim=5120,
487
+ n_heads=40,
488
+ n_kv_heads=n_kv_heads,
489
+ head_dim=128,
490
+ hidden_dim=27648,
491
+ vocab_size=152064,
492
+ rms_norm_eps=1e-6,
493
+ rope_theta=1_000_000.0,
494
+ num_hidden_layers=n_layers,
495
+ max_batch_size=4,
496
+ max_seq_len=4096,
497
+ rope_table_len=8192,
498
+ num_devices=num_devices,
499
+ mesh_device=SimpleNamespace(get_num_devices=lambda: num_devices),
500
+ block_configs=block_configs,
501
+ )
502
+
503
+
504
+ def _construct_qwen25_coder_32b_model(monkeypatch):
505
+ sentinels = {
506
+ "embedding": object(),
507
+ "rope_setup": object(),
508
+ "layer": _make_llama32_layer(),
509
+ "norm": object(),
510
+ "lm_head": object(),
511
+ "sampling": object(),
512
+ }
513
+ monkeypatch.setattr(qwen25_coder_32b_model, "get_tt_ccl", MagicMock(return_value=object()))
514
+ monkeypatch.setattr(qwen25_coder_32b_model, "Sampling1D", MagicMock(return_value=sentinels["sampling"]))
515
+ config = _make_qwen25_coder_32b_config()
516
+ return (
517
+ qwen25_coder_32b_model.Qwen25Coder32B(
518
+ config,
519
+ sentinels["embedding"],
520
+ sentinels["rope_setup"],
521
+ [sentinels["layer"]],
522
+ sentinels["norm"],
523
+ sentinels["lm_head"],
524
+ config.mesh_device,
525
+ ),
526
+ config,
527
+ sentinels,
528
+ )
529
+
530
+
531
+ def _make_qwen3_32b_config(*, n_layers=1, num_devices=8, n_kv_heads=8, sampling_config=None):
532
+ del sampling_config
533
+ block_configs = [
534
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
535
+ ]
536
+ return qwen3_32b_model.Qwen3_32BConfig(
537
+ hf_model_id="Qwen/Qwen3-32B",
538
+ dim=5120,
539
+ n_heads=64,
540
+ n_kv_heads=n_kv_heads,
541
+ head_dim=128,
542
+ hidden_dim=27648,
543
+ vocab_size=151936,
544
+ rms_norm_eps=1e-6,
545
+ rope_theta=1_000_000.0,
546
+ num_hidden_layers=n_layers,
547
+ max_batch_size=4,
548
+ max_seq_len=4096,
549
+ rope_table_len=8192,
550
+ num_devices=num_devices,
551
+ mesh_device=SimpleNamespace(get_num_devices=lambda: num_devices),
552
+ block_configs=block_configs,
553
+ )
554
+
555
+
556
+ def _construct_qwen3_32b_model(monkeypatch):
557
+ sentinels = {
558
+ "embedding": object(),
559
+ "rope_setup": object(),
560
+ "layer": _make_llama32_layer(),
561
+ "norm": object(),
562
+ "lm_head": object(),
563
+ "sampling": object(),
564
+ }
565
+ monkeypatch.setattr(qwen3_32b_model, "get_tt_ccl", MagicMock(return_value=object()))
566
+ monkeypatch.setattr(qwen3_32b_model, "Sampling1D", MagicMock(return_value=sentinels["sampling"]))
567
+ monkeypatch.setattr(qwen3_32b_model.weight_utils, "lm_head_padded_vocab_size", MagicMock(return_value=152064))
568
+ config = _make_qwen3_32b_config()
569
+ return (
570
+ qwen3_32b_model.Qwen3_32B(
571
+ config,
572
+ sentinels["embedding"],
573
+ sentinels["rope_setup"],
574
+ [sentinels["layer"]],
575
+ sentinels["norm"],
576
+ sentinels["lm_head"],
577
+ config.mesh_device,
578
+ ),
579
+ config,
580
+ sentinels,
581
+ )
582
+
583
+
584
+ def _make_mistral_config(*, n_layers=1, num_devices=2, n_kv_heads=8, sampling_config=None):
585
+ block_configs = [
586
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
587
+ ]
588
+ return mistral_model.Mistral7BTransformerConfig(
589
+ n_layers=n_layers,
590
+ vocab_size=32768,
591
+ max_batch_size=4,
592
+ max_seq_len=4096,
593
+ dim=4096,
594
+ num_devices=num_devices,
595
+ mesh_device=object(),
596
+ embedding_config=object(),
597
+ rope_config=object(),
598
+ block_configs=block_configs,
599
+ norm_config=object(),
600
+ lm_head_config=object(),
601
+ sampling_config=sampling_config,
602
+ tt_ccl=object(),
603
+ )
604
+
605
+
606
+ def _construct_mistral_model(monkeypatch):
607
+ sentinels = {
608
+ "embedding": object(),
609
+ "rope_setup": object(),
610
+ "layer": _make_llama32_layer(),
611
+ "norm": object(),
612
+ "lm_head": object(),
613
+ "sampling": object(),
614
+ }
615
+ for owner_name, sentinel_name in (
616
+ ("Embedding1D", "embedding"),
617
+ ("RotarySetup1D", "rope_setup"),
618
+ ("Mistral7BDecoderLayer", "layer"),
619
+ ("RMSNorm1D", "norm"),
620
+ ("LMHead1D", "lm_head"),
621
+ ("Sampling1D", "sampling"),
622
+ ):
623
+ monkeypatch.setattr(
624
+ getattr(mistral_model, owner_name),
625
+ "from_config",
626
+ MagicMock(return_value=sentinels[sentinel_name]),
627
+ )
628
+ config = _make_mistral_config(sampling_config=object())
629
+ return mistral_model.Mistral7B(config), config, sentinels
630
+
631
+
632
+ def _make_phi4_config(*, n_layers=1, num_devices=2, n_kv_heads=10, sampling_config=None):
633
+ block_configs = [
634
+ SimpleNamespace(attention_config=_make_llama32_attention_config(n_kv_heads=n_kv_heads)) for _ in range(n_layers)
635
+ ]
636
+ return phi4_model.Phi4TransformerConfig(
637
+ n_layers=n_layers,
638
+ vocab_size=100352,
639
+ max_batch_size=4,
640
+ max_seq_len=4096,
641
+ dim=5120,
642
+ num_devices=num_devices,
643
+ mesh_device=object(),
644
+ embedding_config=object(),
645
+ rope_config=object(),
646
+ block_configs=block_configs,
647
+ norm_config=object(),
648
+ lm_head_config=object(),
649
+ sampling_config=sampling_config,
650
+ tt_ccl=object(),
651
+ )
652
+
653
+
654
+ def _construct_phi4_model(monkeypatch):
655
+ sentinels = {
656
+ "embedding": object(),
657
+ "rope_setup": object(),
658
+ "layer": _make_llama32_layer(),
659
+ "norm": object(),
660
+ "lm_head": object(),
661
+ "sampling": object(),
662
+ }
663
+ for owner_name, sentinel_name in (
664
+ ("Embedding1D", "embedding"),
665
+ ("RotarySetup1D", "rope_setup"),
666
+ ("Phi4DecoderLayer", "layer"),
667
+ ("RMSNorm1D", "norm"),
668
+ ("LMHead1D", "lm_head"),
669
+ ("Sampling1D", "sampling"),
670
+ ):
671
+ monkeypatch.setattr(
672
+ getattr(phi4_model, owner_name),
673
+ "from_config",
674
+ MagicMock(return_value=sentinels[sentinel_name]),
675
+ )
676
+ config = _make_phi4_config(sampling_config=object())
677
+ return phi4_model.Phi4Transformer(config), config, sentinels
678
+
679
+
680
+ def test_config_exposes_complete_runtime_metadata(contract):
681
+ names = {field.name for field in fields(contract.config_class)}
682
+ assert {
683
+ "dim",
684
+ "mesh_device",
685
+ "num_devices",
686
+ "n_layers",
687
+ "max_batch_size",
688
+ "max_seq_len",
689
+ "block_configs",
690
+ } <= names
691
+ assert "n_kv_heads" in {field.name for field in fields(contract.attention_config_class)}
692
+
693
+ config = contract.make_config(n_layers=2, num_devices=2, n_kv_heads=8)
694
+ assert len(config.block_configs) == config.n_layers
695
+ for block in config.block_configs:
696
+ assert hasattr(block.attention_config, "n_kv_heads")
697
+ assert block.attention_config.n_kv_heads % config.num_devices == 0
698
+
699
+
700
+ @pytest.mark.parametrize(
701
+ "method,names,keyword_only,defaults",
702
+ [
703
+ (
704
+ "prefill_forward",
705
+ (
706
+ "self",
707
+ "x_embed",
708
+ "rot_mats",
709
+ "user_id",
710
+ "page_table",
711
+ "chunk_page_table",
712
+ "chunk_start_idx",
713
+ "get_last_token",
714
+ "batch_size",
715
+ "chunk_start_idx_tensor",
716
+ "last_token_slice",
717
+ "last_token_index",
718
+ ),
719
+ (),
720
+ {
721
+ "user_id": 0,
722
+ "page_table": None,
723
+ "chunk_page_table": None,
724
+ "chunk_start_idx": None,
725
+ "get_last_token": -1,
726
+ "batch_size": 1,
727
+ "chunk_start_idx_tensor": None,
728
+ "last_token_slice": None,
729
+ "last_token_index": None,
730
+ },
731
+ ),
732
+ (
733
+ "post_process_prefill_output",
734
+ ("self", "hidden_states", "last_token_idx", "last_token_slice", "last_token_index"),
735
+ (),
736
+ {"last_token_slice": None, "last_token_index": None},
737
+ ),
738
+ (
739
+ "post_process_batched_prefill_output",
740
+ (
741
+ "self",
742
+ "hidden_states",
743
+ "last_token_idx_list",
744
+ "padded_batch",
745
+ "prefill_seq_len",
746
+ "last_token_slice",
747
+ "last_token_index",
748
+ ),
749
+ (),
750
+ {"last_token_slice": None, "last_token_index": None},
751
+ ),
752
+ ("set_kv_cache", ("self", "kv_cache"), (), {}),
753
+ (
754
+ "configure_paged_attention",
755
+ ("self", "block_size", "max_num_blocks"),
756
+ ("block_size", "max_num_blocks"),
757
+ {},
758
+ ),
759
+ ("prepare_prefill_rot_mats", ("self", "position_indices"), (), {}),
760
+ ("iter_executor_named_modules", ("self",), (), {}),
761
+ ],
762
+ )
763
+ def test_method_signatures_are_exact(contract, method, names, keyword_only, defaults):
764
+ signature = inspect.signature(getattr(contract.model_class, method))
765
+ parameters = signature.parameters
766
+ assert tuple(parameters) == names
767
+ for name, parameter in parameters.items():
768
+ expected_kind = (
769
+ inspect.Parameter.KEYWORD_ONLY if name in keyword_only else inspect.Parameter.POSITIONAL_OR_KEYWORD
770
+ )
771
+ assert parameter.kind is expected_kind
772
+ expected_default = defaults.get(name, inspect.Parameter.empty)
773
+ assert parameter.default == expected_default
774
+
775
+ annotations = {
776
+ "prefill_forward": (
777
+ (
778
+ inspect.Parameter.empty,
779
+ "ttnn.Tensor",
780
+ "tuple[ttnn.Tensor, ttnn.Tensor]",
781
+ "int",
782
+ "ttnn.Tensor | None",
783
+ "ttnn.Tensor | None",
784
+ "int | None",
785
+ "int",
786
+ "int",
787
+ "ttnn.Tensor | None",
788
+ "tuple[ttnn.Tensor, ttnn.Tensor] | None",
789
+ "ttnn.Tensor | None",
790
+ ),
791
+ "ttnn.Tensor",
792
+ ),
793
+ "post_process_prefill_output": (
794
+ (
795
+ inspect.Parameter.empty,
796
+ "ttnn.Tensor",
797
+ "int",
798
+ "tuple[ttnn.Tensor, ttnn.Tensor] | None",
799
+ "ttnn.Tensor | None",
800
+ ),
801
+ "ttnn.Tensor",
802
+ ),
803
+ "post_process_batched_prefill_output": (
804
+ (
805
+ inspect.Parameter.empty,
806
+ "ttnn.Tensor",
807
+ "list[int]",
808
+ "int",
809
+ "int",
810
+ "tuple[ttnn.Tensor, ttnn.Tensor] | None",
811
+ "ttnn.Tensor | None",
812
+ ),
813
+ "ttnn.Tensor",
814
+ ),
815
+ "set_kv_cache": ((inspect.Parameter.empty, "list | None"), "None"),
816
+ "configure_paged_attention": ((inspect.Parameter.empty, "int", "int"), "None"),
817
+ "prepare_prefill_rot_mats": (
818
+ (inspect.Parameter.empty, "ttnn.Tensor"),
819
+ "tuple[ttnn.Tensor, ttnn.Tensor]",
820
+ ),
821
+ "iter_executor_named_modules": (
822
+ (inspect.Parameter.empty,),
823
+ inspect.Signature.empty,
824
+ ),
825
+ }
826
+ expected_parameter_annotations, expected_return_annotation = annotations[method]
827
+ assert tuple(parameter.annotation for parameter in parameters.values()) == expected_parameter_annotations
828
+ assert signature.return_annotation == expected_return_annotation
829
+
830
+
831
+ def test_required_methods_exist(contract):
832
+ for method in (
833
+ "set_kv_cache",
834
+ "configure_paged_attention",
835
+ "prepare_prefill_rot_mats",
836
+ "iter_executor_named_modules",
837
+ "embed_decode",
838
+ "embed_prefill",
839
+ "gather_and_untilize_logits",
840
+ "increment_positions",
841
+ ):
842
+ assert callable(getattr(contract.model_class, method))
843
+
844
+
845
+ def test_constructed_model_resolves_runtime_surface(contract, monkeypatch):
846
+ model, config, sentinels = contract.construct_model(monkeypatch)
847
+
848
+ assert model.config is config
849
+ assert model.embedding is sentinels["embedding"]
850
+ assert model.rope_setup is sentinels["rope_setup"]
851
+ assert model.layers == [sentinels["layer"]]
852
+ assert model.norm is sentinels["norm"]
853
+ assert model.lm_head is sentinels["lm_head"]
854
+ assert model.sampling is sentinels["sampling"]
855
+ assert model.supports_on_device_sampling
856
+ assert model.mesh_device is config.mesh_device
857
+ assert model.vocab_size == config.vocab_size
858
+ assert model.n_layers == config.n_layers
859
+ assert model.num_devices == config.num_devices
860
+ assert model.model_args is None
861
+
862
+
863
+ def test_named_modules_are_complete_unique_and_ordered(contract):
864
+ layers = [contract.make_layer(), contract.make_layer()]
865
+ model = SimpleNamespace(layers=layers, norm=object(), lm_head=object())
866
+ named = list(contract.model_class.iter_executor_named_modules(model))
867
+
868
+ assert tuple(name for name, _ in named) == contract.expected_module_names
869
+ assert len(named) == 4 * len(layers) + 2
870
+ assert len({name for name, _ in named}) == len(named)
871
+ assert tuple(module for _, module in named) == (
872
+ layers[0].attention_norm,
873
+ layers[0].attention,
874
+ layers[0].ff_norm,
875
+ layers[0].feed_forward,
876
+ layers[1].attention_norm,
877
+ layers[1].attention,
878
+ layers[1].ff_norm,
879
+ layers[1].feed_forward,
880
+ model.norm,
881
+ model.lm_head,
882
+ )
883
+
884
+
885
+ def test_named_modules_without_layers_yields_nothing(contract):
886
+ assert list(contract.model_class.iter_executor_named_modules(SimpleNamespace())) == []
887
+
888
+
889
+ def test_set_kv_cache_binds_identity_and_unbinds_idempotently(contract):
890
+ layers = [contract.make_layer(), contract.make_layer()]
891
+ model = SimpleNamespace(layers=layers)
892
+ cache = [[object(), object()], [object(), object()]]
893
+
894
+ contract.model_class.set_kv_cache(model, cache)
895
+ for layer, expected in zip(layers, cache):
896
+ bound = layer.attention.config.kv_cache
897
+ assert bound == tuple(expected)
898
+ assert bound[0] is expected[0]
899
+ assert bound[1] is expected[1]
900
+ assert layer.attention.kv_cache is bound
901
+
902
+ contract.model_class.set_kv_cache(model, None)
903
+ contract.model_class.set_kv_cache(model, None)
904
+ assert all(layer.attention.config.kv_cache is None for layer in layers)
905
+ assert all(layer.attention.kv_cache is None for layer in layers)
906
+
907
+
908
+ def test_set_kv_cache_rejects_wrong_layer_count_before_binding(contract, expect_error):
909
+ layers = [contract.make_layer(), contract.make_layer()]
910
+ model = SimpleNamespace(layers=layers)
911
+
912
+ with expect_error(ValueError, "model has 2 layers"):
913
+ contract.model_class.set_kv_cache(model, [[object(), object()]])
914
+
915
+ assert all(layer.attention.config.kv_cache is None for layer in layers)
916
+ assert all(layer.attention.kv_cache is None for layer in layers)
917
+
918
+
919
+ @pytest.mark.parametrize("bad_pair", [[object()], object()])
920
+ def test_set_kv_cache_validates_all_pairs_before_binding(contract, bad_pair, expect_error):
921
+ layers = [contract.make_layer(), contract.make_layer()]
922
+ model = SimpleNamespace(layers=layers)
923
+
924
+ with expect_error((TypeError, ValueError), "layer 1.*K/V tensor"):
925
+ contract.model_class.set_kv_cache(model, [[object(), object()], bad_pair])
926
+
927
+ assert all(layer.attention.config.kv_cache is None for layer in layers)
928
+ assert all(layer.attention.kv_cache is None for layer in layers)
929
+
930
+
931
+ def test_configure_paged_attention_updates_construction_and_live_configs_and_rejects_bound_cache(
932
+ contract, expect_error
933
+ ):
934
+ construction = contract.make_attention_config()
935
+ live = contract.make_attention_config()
936
+ model = SimpleNamespace(
937
+ config=SimpleNamespace(block_configs=(SimpleNamespace(attention_config=construction),)),
938
+ layers=(contract.make_layer(live),),
939
+ )
940
+
941
+ contract.model_class.configure_paged_attention(model, block_size=16, max_num_blocks=200)
942
+ assert construction.paged_attention_config.block_size == live.paged_attention_config.block_size == 16
943
+ assert construction.paged_attention_config.max_num_blocks == live.paged_attention_config.max_num_blocks == 200
944
+
945
+ bound = object()
946
+ live.kv_cache = (bound, bound)
947
+ with expect_error(RuntimeError, "already has a bound KV cache"):
948
+ contract.model_class.configure_paged_attention(model, block_size=32, max_num_blocks=128)
949
+ assert construction.paged_attention_config.block_size == live.paged_attention_config.block_size == 16
950
+
951
+
952
+ def test_prepare_prefill_rot_mats_gathers_device_rows(contract, monkeypatch):
953
+ position_indices = object()
954
+ cos_matrix, sin_matrix = object(), object()
955
+ cos_rows, sin_rows = object(), object()
956
+ cos_4d, sin_4d = object(), object()
957
+ fake_ttnn = SimpleNamespace(
958
+ TILE_LAYOUT=object(),
959
+ embedding=MagicMock(side_effect=[cos_rows, sin_rows]),
960
+ unsqueeze_to_4D=MagicMock(side_effect=[cos_4d, sin_4d]),
961
+ )
962
+ rope = SimpleNamespace(cos_matrix=cos_matrix, sin_matrix=sin_matrix, load_device_weights=MagicMock())
963
+ monkeypatch.setattr(contract.module, "ttnn", fake_ttnn)
964
+
965
+ result = contract.model_class.prepare_prefill_rot_mats(SimpleNamespace(rope_setup=rope), position_indices)
966
+
967
+ rope.load_device_weights.assert_called_once_with()
968
+ assert fake_ttnn.embedding.call_args_list == [
969
+ call(position_indices, cos_matrix, layout=fake_ttnn.TILE_LAYOUT),
970
+ call(position_indices, sin_matrix, layout=fake_ttnn.TILE_LAYOUT),
971
+ ]
972
+ assert result == (cos_4d, sin_4d)
code/models/common/tests/llm_runtime/test_model_executor.py ADDED
@@ -0,0 +1,291 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ """Focused contracts for the family-neutral model composition root."""
5
+
6
+ import ast
7
+ from pathlib import Path
8
+ from types import SimpleNamespace
9
+ from unittest.mock import MagicMock
10
+
11
+ import pytest
12
+
13
+ import ttnn
14
+ from models.common.llm_runtime.config import PagedKVCacheConfig, TraceConfig, WarmupConfig
15
+ from models.common.models import executor as executor_module
16
+ from models.common.models.executor import ModelExecutor, ModelExecutorConfig
17
+
18
+ _EXECUTOR_PATH = Path(__file__).parents[2] / "models" / "executor.py"
19
+ _MODELS_ROOT = _EXECUTOR_PATH.parent
20
+
21
+
22
+ def _config(*, device_sampling_enabled: bool = False) -> ModelExecutorConfig:
23
+ return ModelExecutorConfig(
24
+ trace=TraceConfig(mode="none"),
25
+ warmup=WarmupConfig(),
26
+ paged_kv_cache=PagedKVCacheConfig(
27
+ block_size=32,
28
+ max_num_blocks=128,
29
+ num_blocks=128,
30
+ dtype=ttnn.bfloat8_b,
31
+ ),
32
+ device_sampling_enabled=device_sampling_enabled,
33
+ )
34
+
35
+
36
+ def test_common_executor_has_no_concrete_model_dependencies_or_dispatch() -> None:
37
+ tree = ast.parse(_EXECUTOR_PATH.read_text())
38
+ imports = {node.module for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) and node.module is not None}
39
+ imported_names = {alias.name for node in ast.walk(tree) if isinstance(node, ast.Import) for alias in node.names}
40
+ assert not any(module.startswith("models.common.models.") for module in imports)
41
+ assert not any(name.startswith("models.common.models.") for name in imported_names)
42
+
43
+ control_flow = [
44
+ node.test for node in ast.walk(tree) if isinstance(node, (ast.If, ast.IfExp, ast.While, ast.Assert))
45
+ ]
46
+ dispatch_names = {"model_name", "model_id", "provider_id", "checkpoint_path", "model_version"}
47
+ assert not any(
48
+ isinstance(node, ast.Name) and node.id in dispatch_names
49
+ for expression in control_flow
50
+ for node in ast.walk(expression)
51
+ )
52
+
53
+ torch_calls = [
54
+ node
55
+ for node in ast.walk(tree)
56
+ if isinstance(node, ast.Call)
57
+ and isinstance(node.func, ast.Attribute)
58
+ and isinstance(node.func.value, ast.Name)
59
+ and node.func.value.id == "torch"
60
+ ]
61
+ assert torch_calls == []
62
+
63
+
64
+ def test_model_layer_has_only_the_approved_family_modules_and_readmes() -> None:
65
+ assert (_MODELS_ROOT / "llama3_executor.py").is_file()
66
+ assert (_MODELS_ROOT / "qwen2_executor.py").is_file()
67
+ assert not (_MODELS_ROOT / "qwen3_executor.py").exists()
68
+
69
+ model_directories = sorted(
70
+ path for path in _MODELS_ROOT.iterdir() if path.is_dir() and (path / "model.py").is_file()
71
+ )
72
+ assert len(model_directories) == 12
73
+ assert all((path / "README.md").is_file() for path in model_directories)
74
+
75
+
76
+ @pytest.mark.parametrize(
77
+ "relative_path",
78
+ (
79
+ "deepseek_r1_distill_qwen_14b/executor.py",
80
+ "mistral_7b/executor.py",
81
+ "phi4/executor.py",
82
+ ),
83
+ )
84
+ def test_direct_composition_examples_do_not_depend_on_shared_family_executors(relative_path: str) -> None:
85
+ tree = ast.parse((_MODELS_ROOT / relative_path).read_text())
86
+ imports = {node.module for node in ast.walk(tree) if isinstance(node, ast.ImportFrom)}
87
+ assert "models.common.models.executor" not in imports
88
+ assert "models.common.models.llama3_executor" not in imports
89
+ assert "models.common.models.qwen2_executor" not in imports
90
+
91
+
92
+ def test_common_config_is_frozen_and_rejects_non_exact_nested_config_types(expect_error) -> None:
93
+ config = _config()
94
+ with expect_error(AttributeError, "cannot assign to field"):
95
+ config.device_sampling_enabled = True
96
+
97
+ class TraceConfigSubclass(TraceConfig):
98
+ pass
99
+
100
+ with expect_error(TypeError, "trace must be exactly TraceConfig"):
101
+ ModelExecutorConfig(
102
+ trace=TraceConfigSubclass(mode="none"),
103
+ warmup=config.warmup,
104
+ paged_kv_cache=config.paged_kv_cache,
105
+ device_sampling_enabled=False,
106
+ )
107
+
108
+
109
+ def test_sampling_state_inputs_are_an_optional_owned_pair(expect_error) -> None:
110
+ with expect_error(ValueError, "must be supplied together"):
111
+ ModelExecutor(
112
+ None,
113
+ None,
114
+ _config(device_sampling_enabled=True),
115
+ sampling_state_controller=object(),
116
+ )
117
+ with expect_error(ValueError, "requires device sampling"):
118
+ ModelExecutor(
119
+ None,
120
+ None,
121
+ _config(),
122
+ sampling_state_controller=object(),
123
+ sampling_state=object(),
124
+ )
125
+
126
+
127
+ @pytest.mark.parametrize(
128
+ ("enable_trace", "expected"),
129
+ [(True, ["prime", "coordinator"]), (False, ["coordinator", "prime"])],
130
+ )
131
+ def test_prefill_warmup_policy_controls_order_through_one_continuation(enable_trace, expected) -> None:
132
+ events = []
133
+
134
+ def policy(executor, default_warmup, *, kv_cache, can_sample_on_device, enable_trace):
135
+ assert executor is target
136
+ assert kv_cache is cache
137
+ assert can_sample_on_device
138
+ if enable_trace:
139
+ events.append("prime")
140
+ default_warmup()
141
+ if not enable_trace:
142
+ events.append("prime")
143
+
144
+ cache = object()
145
+ target = object.__new__(ModelExecutor)
146
+ target._terminal = False
147
+ target.prefill_runtime = SimpleNamespace(transient_orphan_count=0)
148
+ target.decode_runtime = SimpleNamespace(transient_orphan_count=0)
149
+ target._prefill_warmup = policy
150
+ target.warmup = SimpleNamespace(
151
+ warmup_prefill=lambda **kwargs: events.append("coordinator"),
152
+ )
153
+
154
+ target.warmup_model_prefill(
155
+ kv_cache=cache,
156
+ can_sample_on_device=True,
157
+ enable_trace=enable_trace,
158
+ )
159
+
160
+ assert events == expected
161
+
162
+
163
+ def test_request_state_and_execution_target_are_forwarded_by_identity() -> None:
164
+ target = object.__new__(ModelExecutor)
165
+ target._terminal = False
166
+ target.prefill_runtime = SimpleNamespace(transient_orphan_count=0)
167
+ target.decode_runtime = SimpleNamespace(transient_orphan_count=0)
168
+ target._validate_bound_cache = MagicMock()
169
+ target._ensure_sampling_for = MagicMock()
170
+ target._prefill_execution = MagicMock()
171
+ target._request_state_fields = ("prompt_tokens", "output_tokens", "slot_remap")
172
+
173
+ values = {name: object() for name in ("tokens", "page_table", "prompt_tokens", "output_tokens", "slot_remap")}
174
+ target.compile_prefill(
175
+ tokens=values["tokens"],
176
+ page_table=values["page_table"],
177
+ prompt_tokens=values["prompt_tokens"],
178
+ output_tokens=values["output_tokens"],
179
+ slot_remap=values["slot_remap"],
180
+ )
181
+
182
+ forwarded = target._prefill_execution.compile_prefill.call_args.kwargs
183
+ for name, value in values.items():
184
+ assert forwarded[name] is value
185
+
186
+
187
+ def test_layout_refresh_preserves_owner_and_sampling_state_identity(monkeypatch) -> None:
188
+ state = object()
189
+ layout = object()
190
+ prefill_config = SimpleNamespace(
191
+ max_batch_size=4,
192
+ max_prefill_chunk_size=2048,
193
+ device_sampling_enabled=True,
194
+ can_enable_trace=lambda *_: True,
195
+ supports_batched_prefill=True,
196
+ disable_batched_prefill=True,
197
+ max_prefill_batch_size=4,
198
+ batched_prefill_batched_extract=True,
199
+ trace_capture_prime_sequence_lengths=(128,),
200
+ sampling_state_controller=object(),
201
+ sampling_state=state,
202
+ )
203
+ decode_config = SimpleNamespace(
204
+ lane_capacity=4,
205
+ device_sampling_enabled=True,
206
+ force_greedy_top_k=True,
207
+ sampling_state_controller=prefill_config.sampling_state_controller,
208
+ sampling_state=state,
209
+ )
210
+ warmup_config = SimpleNamespace(warmup=object(), prefill_sequence_lengths=(128,))
211
+ resolved_prefill = SimpleNamespace(page_table_layout=layout, sampling_state=state)
212
+ resolved_decode = SimpleNamespace(page_table_layout=layout, sampling_state=state)
213
+ resolved_warmup = SimpleNamespace(page_table_layout=layout)
214
+ prefill_resolve = MagicMock(return_value=resolved_prefill)
215
+ decode_resolve = MagicMock(return_value=resolved_decode)
216
+ warmup_resolve = MagicMock(return_value=resolved_warmup)
217
+ monkeypatch.setattr(executor_module.PrefillRuntimeConfig, "resolve", prefill_resolve)
218
+ monkeypatch.setattr(executor_module.DecodeRuntimeConfig, "resolve", decode_resolve)
219
+ monkeypatch.setattr(executor_module.WarmupCoordinatorConfig, "resolve", warmup_resolve)
220
+
221
+ target = object.__new__(ModelExecutor)
222
+ target.model = object()
223
+ target.output_reader = object()
224
+ target.config = SimpleNamespace(trace=object())
225
+ target.prefill_runtime = SimpleNamespace(config=prefill_config)
226
+ target.decode_runtime = SimpleNamespace(config=decode_config)
227
+ target.warmup = SimpleNamespace(config=warmup_config)
228
+ target._resolve_page_table_layout = lambda: layout
229
+ owners = (target.prefill_runtime, target.decode_runtime, target.warmup)
230
+
231
+ target._refresh_page_table_layout()
232
+
233
+ assert (target.prefill_runtime, target.decode_runtime, target.warmup) == owners
234
+ assert target.page_table_layout is layout
235
+ assert all(owner.config.page_table_layout is layout for owner in owners)
236
+ assert target.prefill_runtime.config.sampling_state is state
237
+ assert target.decode_runtime.config.sampling_state is state
238
+ assert prefill_resolve.call_args.kwargs["trace_capture_prime_sequence_lengths"] == (128,)
239
+ assert prefill_resolve.call_args.kwargs["sampling_state_controller"] is prefill_config.sampling_state_controller
240
+ assert prefill_resolve.call_args.kwargs["sampling_state"] is state
241
+ assert decode_resolve.call_args.kwargs["sampling_state_controller"] is prefill_config.sampling_state_controller
242
+ assert decode_resolve.call_args.kwargs["sampling_state"] is state
243
+
244
+
245
+ def test_cleanup_is_ordered_retryable_idempotent_and_terminal(expect_error) -> None:
246
+ events = []
247
+ failing = {"reader", "trace"}
248
+
249
+ class _Owner:
250
+ def __init__(self, name):
251
+ self.name = name
252
+
253
+ def action(self, *args):
254
+ events.append(self.name)
255
+ if self.name in failing:
256
+ raise RuntimeError(self.name)
257
+
258
+ cleanup = action
259
+ drain = action
260
+ drain_external_outputs = action
261
+ cleanup_transients = action
262
+ release = action
263
+
264
+ target = object.__new__(ModelExecutor)
265
+ target._terminal = False
266
+ target._cleaned_up = False
267
+ target._owner_name = "TestExecutor"
268
+ target.decode_runtime = _Owner("decode")
269
+ target.output_reader = _Owner("reader")
270
+ target.prefill_runtime = _Owner("prefill")
271
+ target.trace_compiler = _Owner("trace")
272
+ target.program_compiler = _Owner("program")
273
+ target.config = SimpleNamespace(device_sampling_enabled=True)
274
+ target.sampling_state_controller = _Owner("sampling-state")
275
+ target.sampling_state = object()
276
+ target.model = SimpleNamespace(sampling=_Owner("sampling"))
277
+ target.kv_cache_manager = _Owner("kv")
278
+
279
+ expected = ["decode", "reader", "prefill", "decode", "trace", "program", "sampling-state", "sampling", "kv"]
280
+ with expect_error(RuntimeError, "reader") as raised:
281
+ target.cleanup()
282
+ assert events == expected
283
+ assert [str(error) for error in raised.value.cleanup_failures] == ["trace"]
284
+ assert target.terminal
285
+ assert not target._cleaned_up
286
+
287
+ failing.clear()
288
+ target.cleanup()
289
+ target.cleanup()
290
+ assert events == expected * 2
291
+ assert target._cleaned_up
code/models/common/tests/llm_runtime/test_output_reader.py ADDED
@@ -0,0 +1,221 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from __future__ import annotations
5
+
6
+ import weakref
7
+
8
+ import torch
9
+
10
+ import ttnn
11
+ from models.common.llm_runtime.output_reader import OutputReader, PendingRead
12
+
13
+
14
+ class FakeDeviceTensor:
15
+ def __init__(self, name: str):
16
+ self.name = name
17
+ self.cpu_calls = []
18
+ self.host_value = object()
19
+
20
+ def cpu(self, *, blocking: bool):
21
+ self.cpu_calls.append(blocking)
22
+ return self.host_value
23
+
24
+
25
+ class FailingDeviceTensor:
26
+ def __init__(self):
27
+ self.cpu_calls = []
28
+
29
+ def cpu(self, *, blocking: bool):
30
+ self.cpu_calls.append(blocking)
31
+ raise RuntimeError("copy failed")
32
+
33
+
34
+ def _install_events(monkeypatch):
35
+ events = []
36
+ synchronized = []
37
+
38
+ def record_event(device, queue_id):
39
+ event = object()
40
+ events.append((device, queue_id, event))
41
+ return event
42
+
43
+ monkeypatch.setattr(ttnn, "record_event", record_event)
44
+ monkeypatch.setattr(ttnn, "event_synchronize", synchronized.append)
45
+ return events, synchronized
46
+
47
+
48
+ def test_blocking_read_returns_nested_host_payload_without_retention(monkeypatch):
49
+ events, synchronized = _install_events(monkeypatch)
50
+ reader = OutputReader("mesh")
51
+ output = FakeDeviceTensor("output")
52
+ log_probs = FakeDeviceTensor("log_probs")
53
+
54
+ host = reader.read((output, log_probs), blocking=True)
55
+
56
+ assert host == (output.host_value, log_probs.host_value)
57
+ assert output.cpu_calls == [True]
58
+ assert log_probs.cpu_calls == [True]
59
+ assert events == []
60
+ assert synchronized == []
61
+
62
+
63
+ def test_synchronized_read_submits_every_copy_then_synchronizes_device_once(monkeypatch):
64
+ synchronized = []
65
+ monkeypatch.setattr(ttnn, "synchronize_device", synchronized.append)
66
+ reader = OutputReader("mesh")
67
+ output = FakeDeviceTensor("output")
68
+ log_probs = FakeDeviceTensor("log_probs")
69
+
70
+ host = reader.read_synchronized((output, log_probs))
71
+
72
+ assert host == (output.host_value, log_probs.host_value)
73
+ assert output.cpu_calls == [False]
74
+ assert log_probs.cpu_calls == [False]
75
+ assert synchronized == ["mesh"]
76
+
77
+
78
+ def test_async_read_retains_destination_and_event_until_completion(monkeypatch):
79
+ events, synchronized = _install_events(monkeypatch)
80
+ reader = OutputReader("mesh")
81
+ output = FakeDeviceTensor("output")
82
+
83
+ pending = reader.submit(output)
84
+
85
+ assert isinstance(pending, PendingRead)
86
+ assert pending.value is output.host_value
87
+ assert pending.events == (events[0][2],)
88
+ assert events[0][:2] == ("mesh", 0)
89
+ assert output.cpu_calls == [False]
90
+
91
+ assert reader.complete(pending) is output.host_value
92
+ assert synchronized == [events[0][2]]
93
+
94
+ assert reader.complete(pending) is output.host_value
95
+ assert synchronized == [events[0][2]]
96
+
97
+
98
+ def test_async_read_can_be_completed_by_exact_unwrapped_value(monkeypatch):
99
+ events, synchronized = _install_events(monkeypatch)
100
+ reader = OutputReader("mesh")
101
+ output = FakeDeviceTensor("output")
102
+ pending = reader.submit(output)
103
+
104
+ assert reader.complete(pending.value) is pending.value
105
+ assert synchronized == [events[0][2]]
106
+
107
+
108
+ def test_reader_rejects_another_readers_pending_handle(monkeypatch, expect_error):
109
+ _install_events(monkeypatch)
110
+ first_reader = OutputReader("first_mesh")
111
+ second_reader = OutputReader("second_mesh")
112
+ foreign_pending = first_reader.submit(FakeDeviceTensor("first"))
113
+ second_reader.submit(FakeDeviceTensor("second"))
114
+
115
+ with expect_error(ValueError, "not owned"):
116
+ second_reader.complete(foreign_pending)
117
+
118
+ first_reader.drain()
119
+ second_reader.drain()
120
+
121
+
122
+ def test_reader_rejects_another_readers_completed_pending_handle(monkeypatch, expect_error):
123
+ _install_events(monkeypatch)
124
+ first_reader = OutputReader("first_mesh")
125
+ second_reader = OutputReader("second_mesh")
126
+ foreign_pending = first_reader.submit(torch.tensor([1]))
127
+
128
+ with expect_error(ValueError, "not owned"):
129
+ second_reader.complete(foreign_pending)
130
+
131
+
132
+ def test_drain_completes_every_pending_read_and_is_idempotent(monkeypatch, expect_error):
133
+ events, synchronized = _install_events(monkeypatch)
134
+ reader = OutputReader("mesh")
135
+ first = reader.submit(FakeDeviceTensor("first"))
136
+ second = reader.submit(FakeDeviceTensor("second"))
137
+
138
+ reader.drain()
139
+
140
+ assert synchronized == [events[0][2], events[1][2]]
141
+ assert reader.complete(first) is first.value
142
+ assert reader.complete(second) is second.value
143
+
144
+ reader.drain()
145
+ assert synchronized == [events[0][2], events[1][2]]
146
+
147
+
148
+ def test_async_host_only_payload_is_already_complete(monkeypatch):
149
+ events, synchronized = _install_events(monkeypatch)
150
+ reader = OutputReader("mesh")
151
+ host_tensor = torch.tensor([1, 2])
152
+
153
+ pending = reader.submit((host_tensor, None))
154
+
155
+ assert torch.equal(pending.value[0], host_tensor)
156
+ assert pending.value[1] is None
157
+ assert pending.events == ()
158
+ assert reader.complete(pending) is pending.value
159
+ assert events == []
160
+ assert synchronized == []
161
+
162
+
163
+ def test_nested_dict_and_list_preserve_shape(monkeypatch):
164
+ _install_events(monkeypatch)
165
+ reader = OutputReader("mesh")
166
+ output = FakeDeviceTensor("output")
167
+ log_probs = FakeDeviceTensor("log_probs")
168
+
169
+ pending = reader.submit({"outputs": [output], "log_probs": log_probs})
170
+
171
+ assert pending.value == {"outputs": [output.host_value], "log_probs": log_probs.host_value}
172
+ reader.complete(pending)
173
+
174
+
175
+ def test_record_event_failure_synchronizes_device(monkeypatch):
176
+ device_synchronizations = []
177
+
178
+ # record_event is an external nanobind API with optional backend arguments.
179
+ def fail_record_event(*_args, **_kwargs):
180
+ raise RuntimeError("record failed")
181
+
182
+ monkeypatch.setattr(ttnn, "record_event", fail_record_event)
183
+ monkeypatch.setattr(ttnn, "synchronize_device", device_synchronizations.append)
184
+ reader = OutputReader("mesh")
185
+
186
+ try:
187
+ reader.submit(FakeDeviceTensor("output"))
188
+ except RuntimeError as error:
189
+ assert str(error) == "record failed"
190
+ else:
191
+ raise AssertionError("record_event failure was not propagated")
192
+
193
+ assert device_synchronizations == ["mesh"]
194
+
195
+
196
+ def test_partial_nested_async_copy_failure_synchronizes_while_destinations_are_retained(monkeypatch, expect_error):
197
+ destination_refs = []
198
+ device_synchronizations = []
199
+
200
+ class HostDestination:
201
+ pass
202
+
203
+ class EphemeralDeviceTensor:
204
+ def cpu(self, *, blocking: bool):
205
+ destination = HostDestination()
206
+ destination_refs.append(weakref.ref(destination))
207
+ return destination
208
+
209
+ def synchronize_device(mesh_device):
210
+ assert destination_refs[0]() is not None
211
+ device_synchronizations.append(mesh_device)
212
+
213
+ monkeypatch.setattr(ttnn, "synchronize_device", synchronize_device)
214
+ reader = OutputReader("mesh")
215
+ failing = FailingDeviceTensor()
216
+
217
+ with expect_error(RuntimeError, "copy failed"):
218
+ reader.submit((EphemeralDeviceTensor(), failing))
219
+
220
+ assert failing.cpu_calls == [False]
221
+ assert device_synchronizations == ["mesh"]
code/models/common/tests/llm_runtime/test_paged_kv_cache.py ADDED
@@ -0,0 +1,415 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from dataclasses import replace
5
+ from types import SimpleNamespace
6
+
7
+ import pytest
8
+ import torch
9
+
10
+ import ttnn
11
+ from models.common.llm_runtime.config import PagedKVCacheConfig
12
+ from models.common.llm_runtime.paged_kv_cache import PagedKVCacheManager, torch_dtype_for_ttnn
13
+
14
+
15
+ class FakeMesh:
16
+ def __init__(self, num_devices=2):
17
+ self._num_devices = num_devices
18
+
19
+ def get_num_devices(self):
20
+ return self._num_devices
21
+
22
+
23
+ class FakeTensor:
24
+ def __init__(self, shape, dtype):
25
+ self.shape = tuple(shape)
26
+ self.dtype = dtype
27
+
28
+ def volume(self):
29
+ result = 1
30
+ for dimension in self.shape:
31
+ result *= dimension
32
+ return result
33
+
34
+ def element_size(self):
35
+ return 1 if self.dtype in (ttnn.bfloat8_b, ttnn.bfloat4_b) else 2
36
+
37
+
38
+ class FakeModel:
39
+ def __init__(self, dtypes=(ttnn.bfloat8_b, ttnn.bfloat8_b), *, bind_fails=False):
40
+ mesh = FakeMesh(2)
41
+ blocks = []
42
+ for dtype in dtypes:
43
+ attention = SimpleNamespace(
44
+ n_kv_heads=8,
45
+ head_dim=16,
46
+ kv_cache_dtype=dtype,
47
+ paged_attention_config=SimpleNamespace(block_size=32, max_num_blocks=8),
48
+ )
49
+ blocks.append(SimpleNamespace(attention_config=attention))
50
+ self.config = SimpleNamespace(
51
+ block_configs=blocks,
52
+ n_layers=len(blocks),
53
+ num_devices=2,
54
+ mesh_device=mesh,
55
+ )
56
+ self.set_calls = []
57
+ self.bound_cache = None
58
+ self.bind_fails = bind_fails
59
+
60
+ def set_kv_cache(self, cache):
61
+ self.set_calls.append(cache)
62
+ if cache is not None and self.bind_fails:
63
+ raise RuntimeError("bind failed")
64
+ self.bound_cache = cache
65
+
66
+
67
+ def cache_config(**overrides):
68
+ values = {
69
+ "block_size": 32,
70
+ "max_num_blocks": 8,
71
+ "dtype": ttnn.bfloat8_b,
72
+ }
73
+ values.update(overrides)
74
+ return PagedKVCacheConfig(**values)
75
+
76
+
77
+ @pytest.fixture
78
+ def fake_allocator(monkeypatch):
79
+ allocated = []
80
+ deallocated = []
81
+
82
+ # Non-failure TTNN fakes retain overloaded backend keyword options for assertions.
83
+ def as_tensor(host_tensor, **kwargs):
84
+ tensor = FakeTensor(host_tensor.shape, kwargs["dtype"])
85
+ allocated.append((tensor, host_tensor, kwargs))
86
+ return tensor
87
+
88
+ monkeypatch.setattr(ttnn, "as_tensor", as_tensor)
89
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh: ("replicate", mesh))
90
+ monkeypatch.setattr(ttnn, "deallocate", lambda tensor: deallocated.append(tensor))
91
+ return allocated, deallocated
92
+
93
+
94
+ def test_unresolved_config_accepts_one_resolved_replacement(expect_error):
95
+ manager = PagedKVCacheManager(FakeModel(), cache_config())
96
+
97
+ resolved = replace(manager.config, num_blocks=4)
98
+ manager.configure(resolved)
99
+
100
+ assert manager.config is resolved
101
+ assert manager.config.num_blocks == 4
102
+ with expect_error(RuntimeError, "only once"):
103
+ manager.configure(replace(resolved, num_blocks=5))
104
+
105
+
106
+ def test_resolved_replacement_accepts_model_aligned_external_geometry(expect_error):
107
+ model = FakeModel()
108
+ manager = PagedKVCacheManager(model, cache_config())
109
+
110
+ for block in model.config.block_configs:
111
+ block.attention_config.paged_attention_config.block_size = 16
112
+ block.attention_config.paged_attention_config.max_num_blocks = 12
113
+ manager.configure(cache_config(block_size=16, max_num_blocks=12, num_blocks=12))
114
+
115
+ assert manager.config.block_size == 16
116
+ assert manager.config.max_num_blocks == manager.config.num_blocks == 12
117
+
118
+ other = PagedKVCacheManager(FakeModel(), cache_config())
119
+ with expect_error(ValueError, "dtype"):
120
+ other.configure(cache_config(dtype=ttnn.bfloat16, num_blocks=4))
121
+
122
+
123
+ def test_pre_resolved_config_starts_configured_and_cannot_be_replaced(expect_error):
124
+ manager = PagedKVCacheManager(FakeModel(), cache_config(num_blocks=4))
125
+
126
+ with expect_error(RuntimeError, "only once"):
127
+ manager.configure(cache_config(num_blocks=5))
128
+
129
+
130
+ def test_model_paged_attention_policy_and_uniform_dtype_are_validated(expect_error):
131
+ model = FakeModel()
132
+ model.config.block_configs[1].attention_config.paged_attention_config.block_size = 16
133
+ with expect_error(ValueError, "block_size"):
134
+ PagedKVCacheManager(model, cache_config())
135
+
136
+ with expect_error(ValueError, "model-owned dtype"):
137
+ PagedKVCacheManager(FakeModel(), cache_config(dtype=ttnn.bfloat16))
138
+
139
+
140
+ def test_vllm_torch_dtype_mapping_includes_quantized_surrogate(expect_error):
141
+ manager = PagedKVCacheManager(FakeModel(), cache_config())
142
+
143
+ assert torch_dtype_for_ttnn(ttnn.bfloat8_b) == torch.bfloat16
144
+ manager.validate_vllm_cache_spec(block_size=32, dtype=torch.bfloat16, num_blocks=8)
145
+ with expect_error(ValueError, "incompatible"):
146
+ manager.validate_vllm_cache_spec(block_size=32, dtype=torch.float32)
147
+ with expect_error(ValueError, "block_size"):
148
+ manager.validate_vllm_cache_spec(block_size=16, dtype=torch.bfloat16)
149
+ with expect_error(ValueError, "exceeds configured maximum"):
150
+ manager.validate_vllm_cache_spec(block_size=32, dtype=torch.bfloat16, num_blocks=9)
151
+
152
+
153
+ def test_nonuniform_model_dtypes_remain_per_layer():
154
+ manager = PagedKVCacheManager(
155
+ FakeModel(dtypes=(ttnn.bfloat8_b, ttnn.bfloat16)),
156
+ cache_config(dtype=ttnn.bfloat8_b),
157
+ )
158
+
159
+ assert manager.per_layer_dtypes == (ttnn.bfloat8_b, ttnn.bfloat16)
160
+ manager.validate_vllm_cache_spec(block_size=32, dtype=torch.bfloat16)
161
+
162
+
163
+ def test_allocate_derives_shapes_and_dtypes_binds_exact_borrowed_handle(fake_allocator, expect_error):
164
+ allocated, _ = fake_allocator
165
+ model = FakeModel(dtypes=(ttnn.bfloat8_b, ttnn.bfloat16))
166
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
167
+
168
+ cache = manager.allocate()
169
+
170
+ assert cache is model.bound_cache
171
+ assert manager.cache_shapes == ((4, 4, 32, 16), (4, 4, 32, 16))
172
+ assert [entry[0].dtype for entry in allocated] == [
173
+ ttnn.bfloat8_b,
174
+ ttnn.bfloat8_b,
175
+ ttnn.bfloat16,
176
+ ttnn.bfloat16,
177
+ ]
178
+ assert all(tuple(entry[1].shape) == (4, 4, 32, 16) for entry in allocated)
179
+ assert len({id(entry[1]) for entry in allocated}) == 1
180
+ assert all(entry[2]["memory_config"] == ttnn.DRAM_MEMORY_CONFIG for entry in allocated)
181
+ assert all(entry[2]["cache_file_name"] is None for entry in allocated)
182
+ manager.validate_borrowed_handle(cache)
183
+ with expect_error(ValueError, "exact manager-owned"):
184
+ manager.validate_borrowed_handle([pair[:] for pair in cache])
185
+ assert manager.bound_context.config is manager.config
186
+ assert manager.bound_context.tensors[0][0] is cache[0][0]
187
+
188
+
189
+ def test_allocate_reuses_legacy_cache_files_when_dtype_is_unambiguous(fake_allocator, tmp_path):
190
+ allocated, _ = fake_allocator
191
+ model = FakeModel()
192
+ model.model_args = SimpleNamespace(model_cache_path=tmp_path)
193
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
194
+
195
+ manager.allocate()
196
+
197
+ shape = (4, 4, 32, 16)
198
+ expected = [
199
+ tmp_path / f"empty_kcache_paged_attention{shape}",
200
+ tmp_path / f"empty_vcache_paged_attention{shape}",
201
+ ]
202
+ assert [entry[2]["cache_file_name"] for entry in allocated] == expected * 2
203
+ assert len({id(entry[1]) for entry in allocated}) == 1
204
+
205
+
206
+ def test_allocate_avoids_cache_file_collision_for_nonuniform_device_dtypes(fake_allocator, tmp_path):
207
+ allocated, _ = fake_allocator
208
+ model = FakeModel(dtypes=(ttnn.bfloat8_b, ttnn.bfloat16))
209
+ model.model_args = SimpleNamespace(model_cache_path=tmp_path)
210
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
211
+
212
+ manager.allocate()
213
+
214
+ assert all(entry[2]["cache_file_name"] is None for entry in allocated)
215
+
216
+
217
+ def test_allocation_requires_resolved_capacity_and_happens_once(fake_allocator, expect_error):
218
+ model = FakeModel()
219
+ manager = PagedKVCacheManager(model, cache_config())
220
+ with expect_error(RuntimeError, "must be resolved"):
221
+ manager.allocate()
222
+
223
+ manager.configure(replace(manager.config, num_blocks=4))
224
+ cache = manager.allocate()
225
+ with expect_error(RuntimeError, "already been allocated"):
226
+ manager.allocate()
227
+ assert cache is model.bound_cache
228
+
229
+
230
+ def test_partial_allocation_failure_deallocates_created_tensors(monkeypatch, expect_error):
231
+ first = FakeTensor((4, 4, 32, 16), ttnn.bfloat8_b)
232
+ calls = 0
233
+ deallocated = []
234
+
235
+ def fail_second_allocation(
236
+ host_tensor,
237
+ *,
238
+ device,
239
+ mesh_mapper,
240
+ layout,
241
+ memory_config,
242
+ dtype,
243
+ cache_file_name,
244
+ ):
245
+ nonlocal calls
246
+ calls += 1
247
+ if calls == 2:
248
+ raise RuntimeError("allocation failed")
249
+ return first
250
+
251
+ monkeypatch.setattr(ttnn, "as_tensor", fail_second_allocation)
252
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh: None)
253
+ monkeypatch.setattr(ttnn, "deallocate", lambda tensor: deallocated.append(tensor))
254
+ model = FakeModel()
255
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
256
+
257
+ with expect_error(RuntimeError, "allocation failed"):
258
+ manager.allocate()
259
+
260
+ assert deallocated == [first]
261
+ assert model.set_calls == []
262
+
263
+
264
+ def test_bind_failure_unbinds_then_deallocates_all_tensors(fake_allocator, expect_error):
265
+ allocated, deallocated = fake_allocator
266
+ model = FakeModel(bind_fails=True)
267
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
268
+
269
+ with expect_error(RuntimeError, "bind failed"):
270
+ manager.allocate()
271
+
272
+ assert model.set_calls[-1] is None
273
+ assert deallocated == [entry[0] for entry in reversed(allocated)]
274
+ assert model.bound_cache is None
275
+
276
+
277
+ def test_release_unbinds_before_deallocating_and_is_idempotent(monkeypatch, expect_error):
278
+ operations = []
279
+ model = FakeModel()
280
+ original_set = model.set_kv_cache
281
+
282
+ def set_kv_cache(cache):
283
+ operations.append(("bind", cache))
284
+ original_set(cache)
285
+
286
+ model.set_kv_cache = set_kv_cache
287
+ monkeypatch.setattr(
288
+ ttnn,
289
+ "as_tensor",
290
+ lambda host_tensor, **kwargs: FakeTensor(host_tensor.shape, kwargs["dtype"]),
291
+ )
292
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh: None)
293
+ monkeypatch.setattr(ttnn, "deallocate", lambda tensor: operations.append(("deallocate", tensor)))
294
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
295
+ cache = manager.allocate()
296
+ operations.clear()
297
+
298
+ manager.release()
299
+
300
+ assert operations[0] == ("bind", None)
301
+ assert [operation[1] for operation in operations[1:]] == [tensor for pair in cache for tensor in pair]
302
+ assert manager.bound_context is None
303
+
304
+ manager.release()
305
+ assert len(operations) == 5
306
+ with expect_error(RuntimeError, "terminal"):
307
+ manager.allocate()
308
+
309
+
310
+ def test_borrowed_handle_mutation_cannot_redirect_owned_tensor_release(fake_allocator, expect_error):
311
+ allocated, deallocated = fake_allocator
312
+ model = FakeModel()
313
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
314
+ cache = manager.allocate()
315
+ owned_tensors = [entry[0] for entry in allocated]
316
+ replacement = FakeTensor(cache[0][0].shape, cache[0][0].dtype)
317
+
318
+ cache[0][0] = replacement
319
+
320
+ with expect_error(ValueError, "exact manager-owned K/V tensors"):
321
+ manager.validate_borrowed_handle(cache)
322
+
323
+ manager.release()
324
+
325
+ assert deallocated == owned_tensors
326
+ assert replacement not in deallocated
327
+
328
+
329
+ def test_release_failure_retains_only_failed_tensor_and_retries(monkeypatch, expect_error):
330
+ tensors = []
331
+
332
+ def as_tensor(host_tensor, **kwargs):
333
+ tensor = FakeTensor(host_tensor.shape, kwargs["dtype"])
334
+ tensors.append(tensor)
335
+ return tensor
336
+
337
+ monkeypatch.setattr(ttnn, "as_tensor", as_tensor)
338
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh: None)
339
+ model = FakeModel()
340
+ manager = PagedKVCacheManager(model, cache_config(num_blocks=4))
341
+ manager.allocate()
342
+ failed_tensor = tensors[0]
343
+ attempts = []
344
+
345
+ def fail_once(tensor):
346
+ attempts.append(tensor)
347
+ if tensor is failed_tensor and attempts.count(tensor) == 1:
348
+ raise RuntimeError("deallocate failed")
349
+
350
+ monkeypatch.setattr(ttnn, "deallocate", fail_once)
351
+
352
+ with expect_error(RuntimeError, "Failed to deallocate 1"):
353
+ manager.release()
354
+
355
+ assert manager.bound_context is None
356
+ assert model.bound_cache is None
357
+ assert attempts == tensors
358
+
359
+ manager.release()
360
+
361
+ assert attempts.count(failed_tensor) == 2
362
+ assert all(attempts.count(tensor) == 1 for tensor in tensors[1:])
363
+
364
+
365
+ def test_partial_allocation_cleanup_failure_preserves_tensor_for_release_retry(monkeypatch, expect_error):
366
+ tensor = FakeTensor((4, 4, 32, 16), ttnn.bfloat8_b)
367
+ allocation_calls = 0
368
+ deallocation_calls = 0
369
+
370
+ def fail_second_allocation(
371
+ host_tensor,
372
+ *,
373
+ device,
374
+ mesh_mapper,
375
+ layout,
376
+ memory_config,
377
+ dtype,
378
+ cache_file_name,
379
+ ):
380
+ nonlocal allocation_calls
381
+ allocation_calls += 1
382
+ if allocation_calls == 2:
383
+ raise RuntimeError("allocation failed")
384
+ return tensor
385
+
386
+ def fail_first_deallocation(value):
387
+ nonlocal deallocation_calls
388
+ deallocation_calls += 1
389
+ if deallocation_calls == 1:
390
+ raise RuntimeError("cleanup failed")
391
+
392
+ monkeypatch.setattr(ttnn, "as_tensor", fail_second_allocation)
393
+ monkeypatch.setattr(ttnn, "ReplicateTensorToMesh", lambda mesh: None)
394
+ monkeypatch.setattr(ttnn, "deallocate", fail_first_deallocation)
395
+ manager = PagedKVCacheManager(FakeModel(), cache_config(num_blocks=4))
396
+
397
+ with expect_error(RuntimeError, "allocation failed") as exc_info:
398
+ manager.allocate()
399
+
400
+ assert [str(error) for error in exc_info.value.cleanup_failures] == ["cleanup failed"]
401
+ manager.release()
402
+
403
+ assert deallocation_calls == 2
404
+ with expect_error(RuntimeError, "terminal"):
405
+ manager.allocate()
406
+
407
+
408
+ def test_release_before_allocation_is_terminal_and_idempotent(expect_error):
409
+ manager = PagedKVCacheManager(FakeModel(), cache_config())
410
+
411
+ manager.release()
412
+ manager.release()
413
+
414
+ with expect_error(RuntimeError, "terminal"):
415
+ manager.allocate()
code/models/common/tests/llm_runtime/test_prefill_inputs.py ADDED
@@ -0,0 +1,173 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+
6
+ import pytest
7
+ import torch
8
+
9
+ import models.common.llm_runtime.prefill.inputs as inputs_module
10
+ from models.common.llm_runtime.prefill.inputs import (
11
+ PrefillDeviceInputs,
12
+ PrefillHostInputs,
13
+ PrefillInputStager,
14
+ allocate_device_tensors,
15
+ copy_into_device_tensors,
16
+ )
17
+
18
+
19
+ class _Model:
20
+ def __init__(self, *, rotary_capacity=8, rotary_outputs=("cos", "sin")):
21
+ self.config = SimpleNamespace(dim=64)
22
+ self.rope_setup = SimpleNamespace(
23
+ cos_matrix=torch.zeros(1, 1, rotary_capacity),
24
+ load_device_weights=lambda: None,
25
+ )
26
+ self.rotary_outputs = rotary_outputs
27
+
28
+ def prepare_prefill_rot_mats(self, position_indices):
29
+ del position_indices
30
+ return self.rotary_outputs
31
+
32
+
33
+ def _stager(*, model=None, released=None):
34
+ released = [] if released is None else released
35
+ return PrefillInputStager(
36
+ model=_Model() if model is None else model,
37
+ mesh_device="mesh",
38
+ release_transient=lambda values: released.append(values) or [],
39
+ )
40
+
41
+
42
+ def _patch_host_conversion(monkeypatch):
43
+ converted = []
44
+ monkeypatch.setattr(inputs_module.ttnn, "ReplicateTensorToMesh", lambda mesh: ("mapper", mesh))
45
+ monkeypatch.setattr(
46
+ inputs_module.ttnn,
47
+ "from_torch",
48
+ lambda value, **kwargs: converted.append((value.clone(), kwargs)) or value.clone(),
49
+ )
50
+ return converted
51
+
52
+
53
+ @pytest.mark.parametrize("shape", [(8,), (1, 2, 8)])
54
+ def test_prepare_host_inputs_rejects_non_matrix_tokens_before_conversion(monkeypatch, expect_error, shape):
55
+ monkeypatch.setattr(inputs_module.ttnn, "from_torch", lambda *args, **kwargs: pytest.fail("converted"))
56
+
57
+ with expect_error(ValueError, "rank 2"):
58
+ _stager().prepare_host_inputs(torch.zeros(shape), torch.zeros(1, 1, dtype=torch.int32))
59
+
60
+
61
+ def test_prepare_host_inputs_rejects_negative_start_and_last_token_beyond_rotary_capacity(
62
+ monkeypatch,
63
+ expect_error,
64
+ ):
65
+ _patch_host_conversion(monkeypatch)
66
+ stager = _stager(model=_Model(rotary_capacity=8))
67
+ tokens = torch.zeros(1, 4, dtype=torch.long)
68
+ page_table = torch.zeros(1, 1, dtype=torch.int32)
69
+
70
+ with expect_error(ValueError, "start position must be nonnegative"):
71
+ stager.prepare_host_inputs(tokens, page_table, start_pos=-1)
72
+ with expect_error(ValueError, "exceeds rotary capacity 8"):
73
+ stager.prepare_host_inputs(tokens, page_table, last_token_idx=8)
74
+
75
+
76
+ def test_prepare_host_inputs_clamps_padded_positions_to_last_rotary_entry(monkeypatch):
77
+ converted = _patch_host_conversion(monkeypatch)
78
+
79
+ _stager(model=_Model(rotary_capacity=8)).prepare_host_inputs(
80
+ torch.zeros(1, 4, dtype=torch.long),
81
+ torch.zeros(1, 1, dtype=torch.int32),
82
+ start_pos=6,
83
+ )
84
+
85
+ position_indices = converted[1][0]
86
+ assert position_indices.tolist() == [[6, 7, 7, 7]]
87
+
88
+
89
+ @pytest.mark.parametrize("relative_last,sequence_length", [(-1, 32), (32, 32), (0, 0)])
90
+ def test_prepare_position_inputs_rejects_positions_outside_padded_sequence(
91
+ monkeypatch,
92
+ expect_error,
93
+ relative_last,
94
+ sequence_length,
95
+ ):
96
+ monkeypatch.setattr(inputs_module.ttnn, "from_torch", lambda *args, **kwargs: pytest.fail("converted"))
97
+
98
+ with expect_error(ValueError, "last-token position"):
99
+ _stager().prepare_position_inputs_host(relative_last, sequence_length)
100
+
101
+
102
+ def test_allocate_device_tensors_releases_partial_allocation_on_failure(monkeypatch, expect_error):
103
+ first_device = object()
104
+ calls = []
105
+
106
+ def to_device(host_tensor, *, device):
107
+ calls.append((host_tensor, device))
108
+ if len(calls) == 2:
109
+ raise RuntimeError("allocation failed")
110
+ return first_device
111
+
112
+ released = []
113
+ monkeypatch.setattr(inputs_module.ttnn, "to_device", to_device)
114
+ monkeypatch.setattr(
115
+ inputs_module,
116
+ "best_effort_deallocate_owned_tensors",
117
+ lambda values: released.append(tuple(values)) or [],
118
+ )
119
+
120
+ with expect_error(RuntimeError, "allocation failed"):
121
+ allocate_device_tensors(("host-0", "host-1"), mesh_device="mesh")
122
+
123
+ assert released == [(first_device,)]
124
+
125
+
126
+ def test_stage_device_inputs_releases_raw_and_malformed_rotary_outputs(monkeypatch, expect_error):
127
+ raw = ["tokens", "positions", "page", None, None]
128
+ released = []
129
+ model = _Model(rotary_outputs=("cos-only",))
130
+ monkeypatch.setattr(inputs_module, "allocate_device_tensors", lambda values, *, mesh_device: raw)
131
+
132
+ host = PrefillHostInputs("host-tokens", "host-positions", "host-page", None, None)
133
+ with expect_error(ValueError, "cosine and sine"):
134
+ _stager(model=model, released=released).stage_device_inputs(host)
135
+
136
+ assert released == [(model.rotary_outputs, raw)]
137
+
138
+
139
+ def test_copy_rotary_inputs_rejects_malformed_output_count_and_releases_it(monkeypatch, expect_error):
140
+ released = []
141
+ model = _Model(rotary_outputs=("cos-only",))
142
+ device = PrefillDeviceInputs("tokens", "cos", "sin", "page", None, "positions", None)
143
+ monkeypatch.setattr(inputs_module.ttnn, "copy", lambda **kwargs: pytest.fail("copied"))
144
+
145
+ with expect_error(ValueError, "cosine and sine"):
146
+ _stager(model=model, released=released).copy_rotary_inputs(device)
147
+
148
+ assert released == [model.rotary_outputs]
149
+
150
+
151
+ @pytest.mark.parametrize(
152
+ ("host", "device"),
153
+ [
154
+ ((None,), ("device",)),
155
+ (("host",), (None,)),
156
+ (("host", None), ("device", "unexpected-device")),
157
+ (("host", None), ("device",)),
158
+ ],
159
+ )
160
+ def test_copy_into_device_tensors_rejects_structure_changes_before_copy(
161
+ monkeypatch,
162
+ expect_error,
163
+ host,
164
+ device,
165
+ ):
166
+ monkeypatch.setattr(
167
+ inputs_module.ttnn,
168
+ "copy_host_to_device_tensor",
169
+ lambda *args: pytest.fail("copied"),
170
+ )
171
+
172
+ with expect_error(ValueError, "host/device"):
173
+ copy_into_device_tensors(host, device)
code/models/common/tests/llm_runtime/test_prefill_runtime.py ADDED
The diff for this file is too large to render. See raw diff
 
code/models/common/tests/llm_runtime/test_program_compiler.py ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from dataclasses import dataclass, field
5
+
6
+ import pytest
7
+ import torch
8
+
9
+ import models.common.llm_runtime.program_compiler as program_compiler_module
10
+ import ttnn
11
+ from models.common.llm_runtime.program_compiler import OutputSpec, ProgramCompiler, ProgramKey
12
+
13
+
14
+ @dataclass(frozen=True)
15
+ class _Signature:
16
+ mode: str
17
+ batch: int
18
+ optional: int | None = None
19
+ runtime_token: int = field(default=0, compare=False)
20
+
21
+ @property
22
+ def key_material(self):
23
+ return (("mode", self.mode), ("batch", self.batch), ("optional", self.optional))
24
+
25
+
26
+ def _compiler(context=object()):
27
+ return ProgramCompiler("mesh", lambda: context)
28
+
29
+
30
+ def _patch_sync(monkeypatch, events=None):
31
+ events = [] if events is None else events
32
+ monkeypatch.setattr(ttnn, "synchronize_device", lambda mesh: events.append(("sync", mesh)))
33
+ return events
34
+
35
+
36
+ def test_program_key_has_stable_golden_digest_and_separate_trace_domain():
37
+ from models.common.llm_runtime.trace_compiler import TraceKey
38
+
39
+ signature = _Signature("decode", 32)
40
+
41
+ assert ProgramKey.from_signature(signature).digest == (
42
+ "c807e5574b34f3f83d386fc441580d1ae499aa534c8a426541dbeafdefa2ffdb"
43
+ )
44
+ assert TraceKey.from_signature(signature).digest == (
45
+ "27bea60de6b500f448e680d3423ad62efbd01294c254524e37c114ff8b5dff35"
46
+ )
47
+ assert ProgramKey.from_signature(signature).digest != TraceKey.from_signature(signature).digest
48
+
49
+
50
+ def test_program_key_changes_for_each_material_field_but_not_runtime_values():
51
+ base = _Signature("decode", 32)
52
+
53
+ assert ProgramKey.from_signature(base) != ProgramKey.from_signature(_Signature("prefill", 32))
54
+ assert ProgramKey.from_signature(base) != ProgramKey.from_signature(_Signature("decode", 16))
55
+ assert ProgramKey.from_signature(base) != ProgramKey.from_signature(_Signature("decode", 32, 1))
56
+ assert ProgramKey.from_signature(base) == ProgramKey.from_signature(_Signature("decode", 32, runtime_token=99))
57
+
58
+
59
+ def test_compiler_memoizes_program_key_by_canonical_material(monkeypatch):
60
+ calls = []
61
+ digest = program_compiler_module.signature_digest
62
+ monkeypatch.setattr(
63
+ program_compiler_module,
64
+ "signature_digest",
65
+ lambda domain, schema_version, signature: calls.append((domain, schema_version, signature))
66
+ or digest(domain, schema_version, signature),
67
+ )
68
+ compiler = _compiler()
69
+
70
+ first = compiler.key_for(_Signature("decode", 32))
71
+ second = compiler.key_for(_Signature("decode", 32, runtime_token=99))
72
+
73
+ assert second is first
74
+ assert len(calls) == 1
75
+
76
+
77
+ def test_compiler_program_key_memo_tracks_mutated_material(monkeypatch):
78
+ class MutableSignature:
79
+ def __init__(self):
80
+ self.value = 1
81
+
82
+ def key_material(self):
83
+ return (("value", self.value),)
84
+
85
+ calls = []
86
+ digest = program_compiler_module.signature_digest
87
+ monkeypatch.setattr(
88
+ program_compiler_module,
89
+ "signature_digest",
90
+ lambda domain, schema_version, signature: calls.append((domain, schema_version, signature))
91
+ or digest(domain, schema_version, signature),
92
+ )
93
+ compiler = _compiler()
94
+ signature = MutableSignature()
95
+
96
+ first = compiler.key_for(signature)
97
+ signature.value = 2
98
+ second = compiler.key_for(signature)
99
+
100
+ assert second != first
101
+ assert compiler.key_for(signature) is second
102
+ assert len(calls) == 2
103
+
104
+
105
+ def test_compiler_program_key_memo_preserves_tagged_type_identity():
106
+ compiler = _compiler()
107
+
108
+ keys = {
109
+ compiler.key_for(type("Signature", (), {"key_material": (("value", value),)})())
110
+ for value in (True, 1, 1.0, "1")
111
+ }
112
+
113
+ assert len(keys) == 4
114
+
115
+
116
+ @pytest.mark.parametrize("material", [{"unordered": "mapping"}, ["list"], torch.zeros(1)])
117
+ def test_program_key_rejects_noncanonical_material(material, expect_error):
118
+ class InvalidSignature:
119
+ key_material = (("invalid", material),)
120
+
121
+ with expect_error(TypeError, "key material"):
122
+ ProgramKey.from_signature(InvalidSignature())
123
+
124
+
125
+ def test_same_digest_with_different_retained_signature_is_rejected(monkeypatch, expect_error):
126
+ _patch_sync(monkeypatch)
127
+ compiler = _compiler()
128
+ forced_key = ProgramKey("0" * 64)
129
+ monkeypatch.setattr(compiler, "key_for", lambda signature: forced_key)
130
+ compiler.compile(_Signature("decode", 32), lambda context: torch.zeros(1))
131
+
132
+ with expect_error(RuntimeError, "collision"):
133
+ compiler.compile(_Signature("prefill", 1), lambda context: torch.zeros(1))
134
+
135
+
136
+ def test_compile_receives_exact_context_deduplicates_and_checks_output_contract(monkeypatch, expect_error):
137
+ events = _patch_sync(monkeypatch)
138
+ context = object()
139
+ compiler = _compiler(context)
140
+ calls = []
141
+ expected = OutputSpec((2, 3), torch.bfloat16)
142
+
143
+ program = compiler.compile(
144
+ _Signature("decode", 32),
145
+ lambda supplied: calls.append(supplied) or torch.zeros(2, 3, dtype=torch.bfloat16),
146
+ expected_output_spec=expected,
147
+ )
148
+ duplicate = compiler.compile(
149
+ _Signature("decode", 32),
150
+ lambda supplied: (_ for _ in ()).throw(AssertionError("duplicate invoked")),
151
+ expected_output_spec=expected,
152
+ )
153
+
154
+ assert program is duplicate
155
+ assert calls == [context]
156
+ assert program.signature == _Signature("decode", 32)
157
+ assert program.output_spec == expected
158
+ assert events == [("sync", "mesh"), ("sync", "mesh")]
159
+ with expect_error(ValueError, "different output contract"):
160
+ compiler.compile(
161
+ _Signature("decode", 32),
162
+ lambda supplied: torch.zeros(1),
163
+ expected_output_spec=OutputSpec((1,), torch.float32),
164
+ )
165
+
166
+
167
+ def test_compiled_program_snapshot_is_immutable_and_registry_authoritative(monkeypatch):
168
+ _patch_sync(monkeypatch)
169
+ compiler = _compiler()
170
+
171
+ first = compiler.compile(_Signature("prefill", 1), lambda _: torch.zeros(1))
172
+ snapshot = compiler.compiled_programs
173
+ second = compiler.compile(_Signature("decode", 32), lambda _: torch.zeros(1))
174
+
175
+ assert snapshot == (first,)
176
+ assert compiler.compiled_programs == (first, second)
177
+
178
+
179
+ def test_compile_uses_explicit_result_value_and_owned_release_selector(monkeypatch):
180
+ _patch_sync(monkeypatch)
181
+
182
+ class OwnedTensor:
183
+ pass
184
+
185
+ @dataclass
186
+ class InvocationResult:
187
+ value: torch.Tensor
188
+ owned: object
189
+
190
+ released = []
191
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
192
+ monkeypatch.setattr(ttnn, "deallocate", released.append)
193
+ owned = OwnedTensor()
194
+ compiler = _compiler()
195
+
196
+ program = compiler.compile(
197
+ _Signature("prefill", 1),
198
+ lambda context: InvocationResult(torch.zeros(1, 2), owned),
199
+ output_spec=lambda result: OutputSpec.from_value(result.value),
200
+ release_output=lambda result: result.owned,
201
+ )
202
+
203
+ assert program.output_spec.shape == (1, 2)
204
+ assert released == [owned]
205
+
206
+
207
+ def test_compile_release_failure_is_retryable_and_blocks_further_compile(monkeypatch, expect_error):
208
+ _patch_sync(monkeypatch)
209
+
210
+ class OwnedTensor:
211
+ pass
212
+
213
+ first = OwnedTensor()
214
+ retry = OwnedTensor()
215
+ attempts = []
216
+ failure = RuntimeError("release once")
217
+
218
+ def deallocate(value):
219
+ attempts.append(value)
220
+ if value is retry and attempts.count(retry) == 1:
221
+ raise failure
222
+
223
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
224
+ monkeypatch.setattr(ttnn, "deallocate", deallocate)
225
+ compiler = _compiler()
226
+
227
+ with expect_error(RuntimeError, "Failed to deallocate") as caught:
228
+ compiler.compile(
229
+ _Signature("decode", 32),
230
+ lambda context: (first, retry),
231
+ output_spec=lambda output: OutputSpec((1,), "dtype"),
232
+ )
233
+
234
+ assert caught.value.cleanup_failures == (failure,)
235
+ assert compiler.compile_orphan_count == 1
236
+ with expect_error(RuntimeError, "unreleased compile outputs"):
237
+ compiler.compile(_Signature("decode", 16), lambda context: torch.zeros(1))
238
+
239
+ compiler.cleanup()
240
+ assert attempts.count(first) == 1
241
+ assert attempts.count(retry) == 2
242
+
243
+
244
+ def test_compile_gate_distinguishes_capture_from_activation(monkeypatch, expect_error):
245
+ _patch_sync(monkeypatch)
246
+ compiler = _compiler()
247
+ compiler.set_trace_capture_in_progress(True)
248
+ with expect_error(RuntimeError, "capture is in progress"):
249
+ compiler.compile(_Signature("decode", 32), lambda context: torch.zeros(1))
250
+ assert compiler.post_activation_compile_rejections == 0
251
+
252
+ compiler.set_trace_capture_in_progress(False)
253
+ compiler.set_trace_active(True)
254
+ with expect_error(RuntimeError, "after trace activation"):
255
+ compiler.compile(_Signature("decode", 32), lambda context: torch.zeros(1))
256
+ assert compiler.post_activation_compile_rejections == 1
257
+
258
+
259
+ def test_cleanup_terminalizes_only_program_metadata(monkeypatch, expect_error):
260
+ _patch_sync(monkeypatch)
261
+ compiler = _compiler()
262
+ program = compiler.compile(_Signature("decode", 32), lambda context: torch.zeros(1))
263
+
264
+ compiler.cleanup()
265
+ compiler.cleanup()
266
+
267
+ with expect_error(RuntimeError, "released"):
268
+ compiler.require_compiled(program.key)
code/models/common/tests/llm_runtime/test_tensor_resources.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from dataclasses import dataclass
5
+
6
+ from models.common.llm_runtime import tensor_resources
7
+ from models.common.llm_runtime.decode import DecodeDeviceInputs, DecodePersistentInputs
8
+ from models.common.llm_runtime.prefill.inputs import PrefillDeviceInputs, PrefillPositionInputs
9
+ from models.common.llm_runtime.prefill.trace import PrefillHiddenPersistentInputs, PrefillReplayState
10
+
11
+
12
+ def test_owned_runtime_containers_release_aliased_tensors_once(monkeypatch):
13
+ class FakeTensor:
14
+ pass
15
+
16
+ shared = FakeTensor()
17
+ other = FakeTensor()
18
+ released = []
19
+ monkeypatch.setattr(tensor_resources.ttnn, "Tensor", FakeTensor)
20
+ monkeypatch.setattr(tensor_resources.ttnn, "deallocate", released.append)
21
+
22
+ decode_inputs = DecodeDeviceInputs(shared, other, shared, None)
23
+ prefill_inputs = PrefillDeviceInputs(shared, other, None, shared, None, other, None)
24
+ positions = PrefillPositionInputs(other, shared, other)
25
+ values = (
26
+ decode_inputs,
27
+ DecodePersistentInputs(decode_inputs, (shared, other, shared)),
28
+ prefill_inputs,
29
+ positions,
30
+ PrefillHiddenPersistentInputs(prefill_inputs),
31
+ PrefillReplayState(positions, (shared, other, shared), shared),
32
+ )
33
+
34
+ assert tensor_resources.best_effort_deallocate_owned_tensors(values) == []
35
+ assert released == [shared, other]
36
+
37
+
38
+ def test_arbitrary_dataclass_is_not_treated_as_an_ownership_projection(monkeypatch):
39
+ class FakeTensor:
40
+ pass
41
+
42
+ @dataclass
43
+ class BorrowedValue:
44
+ tensor: FakeTensor
45
+
46
+ released = []
47
+ monkeypatch.setattr(tensor_resources.ttnn, "Tensor", FakeTensor)
48
+ monkeypatch.setattr(tensor_resources.ttnn, "deallocate", released.append)
49
+
50
+ assert tensor_resources.best_effort_deallocate_owned_tensors(BorrowedValue(FakeTensor())) == []
51
+ assert released == []
52
+
53
+
54
+ def test_raise_cleanup_failures_preserves_primary_and_attaches_rest(expect_error):
55
+ primary = RuntimeError("primary")
56
+ secondary = RuntimeError("secondary")
57
+
58
+ with expect_error(RuntimeError, "primary") as raised:
59
+ tensor_resources.raise_cleanup_failures((primary, secondary))
60
+
61
+ assert raised.value is primary
62
+ assert primary.cleanup_failures == (secondary,)
code/models/common/tests/llm_runtime/test_trace_compiler.py ADDED
@@ -0,0 +1,576 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from dataclasses import dataclass
5
+ from itertools import permutations
6
+
7
+ import pytest
8
+ import torch
9
+
10
+ import models.common.llm_runtime.trace_compiler as trace_compiler_module
11
+ import ttnn
12
+ from models.common.llm_runtime.decode import DecodeDeviceInputs, DecodePersistentInputs
13
+ from models.common.llm_runtime.prefill.inputs import PrefillDeviceInputs, PrefillPositionInputs
14
+ from models.common.llm_runtime.prefill.trace import PrefillHiddenPersistentInputs, PrefillReplayState
15
+ from models.common.llm_runtime.program_compiler import ProgramCompiler
16
+ from models.common.llm_runtime.trace_compiler import (
17
+ InputRefreshPolicy,
18
+ PersistentInputs,
19
+ TraceCapturePlan,
20
+ TraceCompiler,
21
+ )
22
+
23
+
24
+ @dataclass(frozen=True)
25
+ class _Signature:
26
+ kind: str
27
+ variant: int
28
+
29
+ @property
30
+ def key_material(self):
31
+ return (("kind", self.kind), ("variant", self.variant))
32
+
33
+
34
+ def _patch_backend(monkeypatch, events):
35
+ next_trace_id = iter(range(100, 200))
36
+ monkeypatch.setattr(ttnn, "synchronize_device", lambda mesh: events.append(("sync", mesh)))
37
+ monkeypatch.setattr(
38
+ ttnn,
39
+ "begin_trace_capture",
40
+ lambda mesh, cq_id: events.append(("begin", mesh, cq_id)) or next(next_trace_id),
41
+ )
42
+ monkeypatch.setattr(
43
+ ttnn,
44
+ "end_trace_capture",
45
+ lambda mesh, trace_id, cq_id: events.append(("end", trace_id, cq_id)),
46
+ )
47
+ monkeypatch.setattr(
48
+ ttnn,
49
+ "execute_trace",
50
+ lambda mesh, trace_id, cq_id, blocking: events.append(("execute", trace_id, cq_id, blocking)),
51
+ )
52
+ monkeypatch.setattr(ttnn, "release_trace", lambda mesh, trace_id: events.append(("release", trace_id)))
53
+ monkeypatch.setattr(trace_compiler_module, "_trim_host_allocator", lambda: events.append(("trim",)))
54
+
55
+
56
+ def _compiled_program(program_compiler, monkeypatch, variant):
57
+ monkeypatch.setattr(ttnn, "synchronize_device", lambda mesh: None)
58
+ return program_compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1))
59
+
60
+
61
+ def _plan(program, variant, events, *, operation="decode", policy=InputRefreshPolicy()):
62
+ return TraceCapturePlan(
63
+ program_key=program.key,
64
+ trace_signature=_Signature("trace", variant),
65
+ operation=operation,
66
+ prepare_inputs=lambda: events.append(("prepare", variant)) or (),
67
+ capture=lambda persistent: events.append(("capture", variant)) or torch.zeros(1),
68
+ refresh_policy=policy,
69
+ )
70
+
71
+
72
+ def test_trace_compiler_retains_exact_program_compiler_and_separate_registries(monkeypatch):
73
+ compiler = ProgramCompiler("mesh", lambda: object())
74
+ program = _compiled_program(compiler, monkeypatch, 1)
75
+ trace = TraceCompiler(compiler)
76
+ trace_key = trace.register_capture_plan(_plan(program, 1, []))
77
+
78
+ assert trace.program_compiler is compiler
79
+ assert trace.mesh_device is compiler.mesh_device
80
+ assert trace.trace_key_for_program(program.key) == trace_key
81
+ assert trace.get(trace_key) is not None
82
+ assert compiler.require_compiled(program.key) is program
83
+ assert not hasattr(program, "artifact")
84
+ assert not hasattr(trace, "programs")
85
+
86
+
87
+ def test_trace_aliases_share_one_artifact_without_copying_program_records(monkeypatch):
88
+ events = []
89
+ _patch_backend(monkeypatch, events)
90
+ compiler = ProgramCompiler("mesh", lambda: object())
91
+ first = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1))
92
+ second = compiler.compile(_Signature("program", 2), lambda context: torch.zeros(1))
93
+ trace = TraceCompiler(compiler)
94
+ shared_signature = _Signature("trace", 1)
95
+ first_key = trace.register_capture_plan(
96
+ TraceCapturePlan(first.key, shared_signature, "decode", lambda: (), lambda persistent: torch.zeros(1))
97
+ )
98
+ second_key = trace.register_capture_plan(
99
+ TraceCapturePlan(second.key, shared_signature, "decode", lambda: (), lambda persistent: torch.zeros(1))
100
+ )
101
+
102
+ trace.capture_all()
103
+
104
+ assert first_key == second_key
105
+ assert trace.get(first_key) is not None
106
+ assert trace.trace_key_for_program(first.key) == first_key
107
+ assert trace.trace_key_for_program(second.key) == first_key
108
+ assert compiler.require_compiled(first.key) is first
109
+ assert compiler.require_compiled(second.key) is second
110
+ assert [event[0] for event in events].count("begin") == 1
111
+ assert trace.trace_count == 1
112
+ assert trace.trace_association_count == 2
113
+ assert trace.registered_coverage("decode") == ((first_key, shared_signature),)
114
+ assert trace.registered_coverage("prefill") == ()
115
+
116
+
117
+ @pytest.mark.parametrize("order", tuple(permutations(("logits", "argmax", "topk"))))
118
+ def test_hidden_trace_alias_workspaces_are_registration_order_independent(monkeypatch, order):
119
+ events = []
120
+ _patch_backend(monkeypatch, events)
121
+ compiler = ProgramCompiler("mesh", lambda: object())
122
+ programs = {
123
+ path: compiler.compile(_Signature("program", index), lambda context: torch.zeros(1))
124
+ for index, path in enumerate(("logits", "argmax", "topk"), 1)
125
+ }
126
+ trace = TraceCompiler(compiler)
127
+ signature = _Signature("shared-hidden", 1)
128
+
129
+ for path in order:
130
+ trace.register_capture_plan(
131
+ TraceCapturePlan(
132
+ programs[path].key,
133
+ signature,
134
+ "prefill",
135
+ lambda: "hidden-inputs",
136
+ lambda persistent: "hidden-output",
137
+ schema_fingerprint=("hidden-v2",),
138
+ prepare_workspace=lambda path=path: f"{path}-workspace",
139
+ workspace_fingerprint=("postprocess-v1", path),
140
+ )
141
+ )
142
+
143
+ trace.capture_all()
144
+
145
+ assert [event[0] for event in events].count("begin") == 1
146
+ assert {path: trace.workspace_for_program(program.key) for path, program in programs.items()} == {
147
+ "logits": "logits-workspace",
148
+ "argmax": "argmax-workspace",
149
+ "topk": "topk-workspace",
150
+ }
151
+
152
+
153
+ def test_trace_alias_rejects_a_mismatched_persistent_schema(monkeypatch, expect_error):
154
+ compiler = ProgramCompiler("mesh", lambda: object())
155
+ first = _compiled_program(compiler, monkeypatch, 1)
156
+ second = _compiled_program(compiler, monkeypatch, 2)
157
+ trace = TraceCompiler(compiler)
158
+ signature = _Signature("trace", 1)
159
+ trace.register_capture_plan(
160
+ TraceCapturePlan(first.key, signature, "prefill", lambda: (), lambda _: (), schema_fingerprint=("v1",))
161
+ )
162
+
163
+ with expect_error(ValueError, "different schema fingerprint"):
164
+ trace.register_capture_plan(
165
+ TraceCapturePlan(second.key, signature, "prefill", lambda: (), lambda _: (), schema_fingerprint=("v2",))
166
+ )
167
+
168
+
169
+ def test_capture_allocates_every_input_before_capture_and_coordinates_gates(monkeypatch, expect_error):
170
+ events = []
171
+ _patch_backend(monkeypatch, events)
172
+ compiler = ProgramCompiler("mesh", lambda: object())
173
+ programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)]
174
+ trace = TraceCompiler(compiler)
175
+
176
+ def capture_with_gate(persistent):
177
+ events.append(("capture", 1))
178
+ with expect_error(RuntimeError, "capture is in progress"):
179
+ compiler.compile(_Signature("program", 3), lambda context: torch.zeros(1))
180
+ return torch.zeros(1)
181
+
182
+ trace.register_capture_plan(
183
+ TraceCapturePlan(
184
+ programs[0].key,
185
+ _Signature("trace", 1),
186
+ "decode",
187
+ lambda: events.append(("prepare", 1)) or (),
188
+ capture_with_gate,
189
+ )
190
+ )
191
+ trace.register_capture_plan(_plan(programs[1], 2, events))
192
+ events.clear()
193
+
194
+ trace.capture_all()
195
+
196
+ first_begin = next(index for index, event in enumerate(events) if event[0] == "begin")
197
+ assert events[:first_begin] == [("prepare", 1), ("prepare", 2)]
198
+ assert trace.trace_active and compiler.trace_active
199
+ assert not compiler.trace_capture_in_progress
200
+ with expect_error(RuntimeError, "after trace activation"):
201
+ compiler.compile(_Signature("program", 3), lambda context: torch.zeros(1))
202
+
203
+
204
+ def test_opt_in_capture_prime_runs_after_allocations_and_before_capture_gate(monkeypatch):
205
+ events = []
206
+ _patch_backend(monkeypatch, events)
207
+ compiler = ProgramCompiler("mesh", lambda: object())
208
+ program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1))
209
+ trace = TraceCompiler(compiler)
210
+
211
+ def prime(persistent):
212
+ assert not compiler.trace_capture_in_progress
213
+ events.append(("prime", persistent))
214
+ return "prime-output"
215
+
216
+ def release_prime_output(output):
217
+ events.append(("release-prime", output))
218
+ return []
219
+
220
+ trace.register_capture_plan(
221
+ TraceCapturePlan(
222
+ program.key,
223
+ _Signature("trace", 1),
224
+ "prefill",
225
+ lambda: events.append(("prepare", 1)) or "persistent",
226
+ lambda persistent: events.append(("capture", persistent)) or torch.zeros(1),
227
+ prepare_workspace=lambda: events.append(("workspace", 1)) or "workspace",
228
+ workspace_fingerprint=("workspace",),
229
+ prime=prime,
230
+ release_prime_output=release_prime_output,
231
+ )
232
+ )
233
+ events.clear()
234
+
235
+ trace.capture_all()
236
+
237
+ first_begin = next(index for index, event in enumerate(events) if event[0] == "begin")
238
+ assert events[:first_begin] == [
239
+ ("prepare", 1),
240
+ ("workspace", 1),
241
+ ("prime", PersistentInputs("persistent")),
242
+ ("sync", "mesh"),
243
+ ("release-prime", "prime-output"),
244
+ ("sync", "mesh"),
245
+ ]
246
+ assert ("capture", PersistentInputs("persistent")) in events[first_begin:]
247
+
248
+
249
+ def test_capture_prime_failure_rolls_back_without_beginning_trace(monkeypatch, expect_error):
250
+ events = []
251
+ _patch_backend(monkeypatch, events)
252
+
253
+ class OwnedTensor:
254
+ pass
255
+
256
+ released = []
257
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
258
+ monkeypatch.setattr(ttnn, "deallocate", released.append)
259
+ compiler = ProgramCompiler("mesh", lambda: object())
260
+ program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1))
261
+ trace = TraceCompiler(compiler)
262
+ persistent = OwnedTensor()
263
+ trace.register_capture_plan(
264
+ TraceCapturePlan(
265
+ program.key,
266
+ _Signature("trace", 1),
267
+ "prefill",
268
+ lambda: persistent,
269
+ lambda _: pytest.fail("capture must not begin after prime failure"),
270
+ prime=lambda _: (_ for _ in ()).throw(RuntimeError("prime failed")),
271
+ release_prime_output=lambda _: [],
272
+ )
273
+ )
274
+
275
+ with expect_error(RuntimeError, "prime failed"):
276
+ trace.capture_all()
277
+
278
+ assert not any(event[0] == "begin" for event in events)
279
+ assert released == [persistent]
280
+ assert not trace.trace_active and not compiler.trace_active
281
+
282
+
283
+ def test_capture_prime_release_failure_still_synchronizes_before_rollback(monkeypatch, expect_error):
284
+ events = []
285
+ _patch_backend(monkeypatch, events)
286
+ compiler = ProgramCompiler("mesh", lambda: object())
287
+ program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1))
288
+ trace = TraceCompiler(compiler)
289
+ release_error = RuntimeError("prime release failed")
290
+ trace.register_capture_plan(
291
+ TraceCapturePlan(
292
+ program.key,
293
+ _Signature("trace", 1),
294
+ "prefill",
295
+ lambda: (),
296
+ lambda _: pytest.fail("capture must not begin after prime release failure"),
297
+ prime=lambda _: events.append(("prime",)) or "prime-output",
298
+ release_prime_output=lambda output: events.append(("release-prime", output)) or [release_error],
299
+ )
300
+ )
301
+ events.clear()
302
+
303
+ with expect_error(RuntimeError, "prime release failed") as caught:
304
+ trace.capture_all()
305
+
306
+ assert caught.value is release_error
307
+ assert events[:4] == [
308
+ ("prime",),
309
+ ("sync", "mesh"),
310
+ ("release-prime", "prime-output"),
311
+ ("sync", "mesh"),
312
+ ]
313
+ assert not any(event[0] == "begin" for event in events)
314
+
315
+
316
+ def test_capture_orders_decode_before_prefill_after_allocating_every_input(monkeypatch):
317
+ events = []
318
+ _patch_backend(monkeypatch, events)
319
+ compiler = ProgramCompiler("mesh", lambda: object())
320
+ programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)]
321
+ trace = TraceCompiler(compiler)
322
+ trace.register_capture_plan(_plan(programs[0], 1, events, operation="prefill"))
323
+ trace.register_capture_plan(_plan(programs[1], 2, events, operation="decode"))
324
+ events.clear()
325
+
326
+ trace.capture_all()
327
+
328
+ first_begin = next(index for index, event in enumerate(events) if event[0] == "begin")
329
+ assert events[:first_begin] == [("prepare", 1), ("prepare", 2)]
330
+ assert [event for event in events if event[0] == "capture"] == [("capture", 2), ("capture", 1)]
331
+
332
+
333
+ def test_capture_failure_rolls_back_traces_and_uncaptured_inputs(monkeypatch, expect_error):
334
+ events = []
335
+ _patch_backend(monkeypatch, events)
336
+
337
+ class OwnedTensor:
338
+ pass
339
+
340
+ released = []
341
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
342
+ monkeypatch.setattr(ttnn, "deallocate", released.append)
343
+ compiler = ProgramCompiler("mesh", lambda: object())
344
+ programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)]
345
+ trace = TraceCompiler(compiler)
346
+ first_input, first_output, second_input = OwnedTensor(), OwnedTensor(), OwnedTensor()
347
+ trace.register_capture_plan(
348
+ TraceCapturePlan(
349
+ programs[0].key,
350
+ _Signature("trace", 1),
351
+ "decode",
352
+ lambda: first_input,
353
+ lambda persistent: first_output,
354
+ )
355
+ )
356
+ primary = RuntimeError("second capture failed")
357
+ trace.register_capture_plan(
358
+ TraceCapturePlan(
359
+ programs[1].key,
360
+ _Signature("trace", 2),
361
+ "decode",
362
+ lambda: second_input,
363
+ lambda persistent: (_ for _ in ()).throw(primary),
364
+ )
365
+ )
366
+
367
+ with expect_error(RuntimeError, "second capture failed") as caught:
368
+ trace.capture_all()
369
+
370
+ assert caught.value is primary
371
+ assert released.count(first_input) == 1
372
+ assert released.count(first_output) == 1
373
+ assert released.count(second_input) == 1
374
+ assert not trace.trace_active and not compiler.trace_active
375
+ for program in programs:
376
+ trace_key = trace.trace_key_for_program(program.key)
377
+ assert trace_key is not None
378
+ assert trace.get(trace_key).artifact is None
379
+
380
+
381
+ def test_incomplete_capture_rollback_keeps_program_gate_closed_until_cleanup(monkeypatch, expect_error):
382
+ events = []
383
+ _patch_backend(monkeypatch, events)
384
+
385
+ class OwnedTensor:
386
+ pass
387
+
388
+ retry = OwnedTensor()
389
+ attempts = []
390
+
391
+ def deallocate(value):
392
+ attempts.append(value)
393
+ if value is retry and attempts.count(retry) == 1:
394
+ raise RuntimeError("release once")
395
+
396
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
397
+ monkeypatch.setattr(ttnn, "deallocate", deallocate)
398
+ compiler = ProgramCompiler("mesh", lambda: object())
399
+ programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)]
400
+ trace = TraceCompiler(compiler)
401
+ trace.register_capture_plan(
402
+ TraceCapturePlan(programs[0].key, _Signature("trace", 1), "decode", lambda: retry, lambda _: torch.zeros(1))
403
+ )
404
+ trace.register_capture_plan(
405
+ TraceCapturePlan(
406
+ programs[1].key,
407
+ _Signature("trace", 2),
408
+ "decode",
409
+ lambda: (_ for _ in ()).throw(RuntimeError("prepare failed")),
410
+ lambda _: torch.zeros(1),
411
+ )
412
+ )
413
+
414
+ with expect_error(RuntimeError, "prepare failed"):
415
+ trace.capture_all()
416
+
417
+ assert trace.trace_active and compiler.trace_active
418
+ with expect_error(RuntimeError, "after trace activation"):
419
+ compiler.compile(_Signature("program", 3), lambda context: torch.zeros(1))
420
+
421
+ trace.cleanup()
422
+ assert attempts.count(retry) == 2
423
+ assert not compiler.trace_active
424
+
425
+
426
+ def test_replay_refresh_decisions_cover_first_replay_page_change_feedback_and_switch(monkeypatch):
427
+ events = []
428
+ _patch_backend(monkeypatch, events)
429
+ compiler = ProgramCompiler("mesh", lambda: object())
430
+ programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)]
431
+ policy = InputRefreshPolicy(every_replay=("position", "sampling"))
432
+ trace = TraceCompiler(compiler)
433
+ for variant, program in enumerate(programs, 1):
434
+ trace.register_capture_plan(_plan(program, variant, events, policy=policy))
435
+ trace.capture_all()
436
+ decisions = []
437
+
438
+ trace.replay(
439
+ programs[0].key,
440
+ lambda artifact, decision: decisions.append(decision),
441
+ device_feedback_enabled=True,
442
+ feedback_compatible=True,
443
+ )
444
+ trace.replay(
445
+ programs[0].key,
446
+ lambda artifact, decision: decisions.append(decision),
447
+ device_feedback_enabled=True,
448
+ feedback_compatible=True,
449
+ page_table_changed=True,
450
+ )
451
+ trace.replay(
452
+ programs[1].key,
453
+ lambda artifact, decision: decisions.append(decision),
454
+ device_feedback_enabled=True,
455
+ feedback_compatible=True,
456
+ )
457
+ trace.replay(
458
+ programs[1].key,
459
+ lambda artifact, decision: decisions.append(decision),
460
+ device_feedback_enabled=False,
461
+ )
462
+
463
+ assert [decision.full for decision in decisions] == [True, False, True, True]
464
+ assert [decision.page_table for decision in decisions] == [False, True, False, False]
465
+ assert all(decision.fields == ("position", "sampling") for decision in decisions)
466
+ assert [event[0] for event in events].count("execute") == 4
467
+ assert trace.replay_count == 4
468
+ assert trace.replay_counts == {"prefill": 0, "decode": 4}
469
+
470
+
471
+ def test_replay_counters_increment_only_after_successful_submission(monkeypatch, expect_error):
472
+ events = []
473
+ _patch_backend(monkeypatch, events)
474
+ compiler = ProgramCompiler("mesh", lambda: object())
475
+ program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1))
476
+ trace = TraceCompiler(compiler)
477
+ trace.register_capture_plan(_plan(program, 1, events, operation="prefill"))
478
+ trace.capture_all()
479
+
480
+ monkeypatch.setattr(ttnn, "execute_trace", lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("submit")))
481
+ with expect_error(RuntimeError, "submit"):
482
+ trace.replay(program.key, lambda artifact, decision: None)
483
+ assert trace.replay_count == 0
484
+ assert trace.replay_counts == {"prefill": 0, "decode": 0}
485
+
486
+
487
+ def test_cleanup_retries_trace_release_before_deallocating_and_does_not_own_programs(monkeypatch, expect_error):
488
+ events = []
489
+ _patch_backend(monkeypatch, events)
490
+
491
+ class OwnedTensor:
492
+ pass
493
+
494
+ persistent, output = OwnedTensor(), OwnedTensor()
495
+ deallocated = []
496
+ release_attempts = []
497
+
498
+ def release(mesh, trace_id):
499
+ release_attempts.append(trace_id)
500
+ if len(release_attempts) == 1:
501
+ raise RuntimeError("trace release once")
502
+
503
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
504
+ monkeypatch.setattr(ttnn, "deallocate", deallocated.append)
505
+ monkeypatch.setattr(ttnn, "release_trace", release)
506
+ compiler = ProgramCompiler("mesh", lambda: object())
507
+ program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1))
508
+ trace = TraceCompiler(compiler)
509
+ trace.register_capture_plan(
510
+ TraceCapturePlan(
511
+ program.key,
512
+ _Signature("trace", 1),
513
+ "decode",
514
+ lambda: persistent,
515
+ lambda values: output,
516
+ )
517
+ )
518
+ trace.capture_all()
519
+
520
+ with expect_error(RuntimeError, "Failed to release"):
521
+ trace.cleanup()
522
+ assert deallocated == []
523
+
524
+ trace.cleanup()
525
+ trace.cleanup()
526
+ assert release_attempts == [100, 100]
527
+ assert deallocated.count(persistent) == 1
528
+ assert deallocated.count(output) == 1
529
+ compiler.cleanup()
530
+
531
+
532
+ def test_cleanup_releases_operation_owned_persistent_dataclasses_once(monkeypatch):
533
+ events = []
534
+ _patch_backend(monkeypatch, events)
535
+
536
+ class OwnedTensor:
537
+ pass
538
+
539
+ values = [OwnedTensor() for _ in range(20)]
540
+ decode = DecodePersistentInputs(
541
+ DecodeDeviceInputs(*values[:4]),
542
+ tuple(values[4:7]),
543
+ )
544
+ prefill = PrefillHiddenPersistentInputs(PrefillDeviceInputs(*values[7:14]))
545
+ prefill_workspace = PrefillReplayState(
546
+ PrefillPositionInputs(*values[14:17]),
547
+ (values[17], values[17], values[17]),
548
+ values[18],
549
+ )
550
+ deallocated = []
551
+ monkeypatch.setattr(ttnn, "Tensor", OwnedTensor)
552
+ monkeypatch.setattr(ttnn, "deallocate", deallocated.append)
553
+
554
+ compiler = ProgramCompiler("mesh", lambda: object())
555
+ programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)]
556
+ trace = TraceCompiler(compiler)
557
+ trace.register_capture_plan(
558
+ TraceCapturePlan(programs[0].key, _Signature("trace", 1), "decode", lambda: decode, lambda _: values[0])
559
+ )
560
+ trace.register_capture_plan(
561
+ TraceCapturePlan(
562
+ programs[1].key,
563
+ _Signature("trace", 2),
564
+ "prefill",
565
+ lambda: prefill,
566
+ lambda _: values[19],
567
+ prepare_workspace=lambda: prefill_workspace,
568
+ workspace_fingerprint=("prefill-workspace",),
569
+ )
570
+ )
571
+
572
+ trace.capture_all()
573
+ trace.cleanup()
574
+
575
+ assert len(deallocated) == len(values)
576
+ assert {id(value) for value in deallocated} == {id(value) for value in values}
code/models/common/tests/llm_runtime/test_vllm_adapter.py ADDED
@@ -0,0 +1,1025 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import dataclasses
5
+ import inspect
6
+ from types import SimpleNamespace
7
+ from unittest.mock import create_autospec
8
+
9
+ import pytest
10
+ import torch
11
+
12
+ import ttnn
13
+ from models.common.llm_runtime.config import PagedKVCacheConfig, TraceConfig
14
+ from models.common.llm_runtime.vllm_adapter import (
15
+ NormalizedDecodeKwargs,
16
+ NormalizedPrefillKwargs,
17
+ VLLMAdapter,
18
+ VLLMAdapterConfig,
19
+ )
20
+ from models.common.models.llama3_8b import generator as llama3_generator_module
21
+ from models.common.models.llama3_8b.executor import Llama3Executor
22
+ from models.common.models.llama3_8b.generator import Llama3Generator
23
+ from models.common.models.llama33_70b import generator as llama70_generator_module
24
+ from models.common.models.qwen2_7b.generator import Qwen2Generator
25
+ from models.common.models.qwen3_32b import generator as qwen3_generator_module
26
+ from models.common.models.qwen25_7b.generator import Qwen25Generator
27
+ from models.common.models.qwen25_72b.generator import Qwen25_72BGenerator
28
+ from models.common.models.qwen25_coder_32b.generator import Qwen25Coder32BGenerator
29
+
30
+
31
+ def _adapter(
32
+ *,
33
+ trace=None,
34
+ paged_config=None,
35
+ model_dtype=ttnn.bfloat8_b,
36
+ request_state_fields=("prompt_tokens", "output_tokens", "slot_remap"),
37
+ ):
38
+ return VLLMAdapter(
39
+ VLLMAdapterConfig.resolve(
40
+ trace=trace or TraceConfig(mode="all"),
41
+ paged_kv_cache=paged_config or PagedKVCacheConfig(block_size=32, max_num_blocks=128, dtype=ttnn.bfloat8_b),
42
+ expected_num_layers=32,
43
+ expected_kv_heads_per_device=8,
44
+ expected_head_dim=128,
45
+ model_kv_cache_dtype=model_dtype,
46
+ request_state_fields=request_state_fields,
47
+ )
48
+ )
49
+
50
+
51
+ def test_config_resolves_canonical_static_policy_and_is_frozen(expect_error):
52
+ trace = TraceConfig(mode="all")
53
+ paged_kv_cache = PagedKVCacheConfig(block_size=32, max_num_blocks=128, dtype=ttnn.bfloat8_b)
54
+
55
+ config = VLLMAdapterConfig.resolve(
56
+ trace=trace,
57
+ paged_kv_cache=paged_kv_cache,
58
+ expected_num_layers=32.0,
59
+ expected_kv_heads_per_device=8.0,
60
+ expected_head_dim=128.0,
61
+ model_kv_cache_dtype=[ttnn.bfloat8_b] * 32,
62
+ )
63
+
64
+ assert config.trace is trace
65
+ assert config.paged_kv_cache is paged_kv_cache
66
+ assert config.expected_num_layers == 32
67
+ assert isinstance(config.expected_num_layers, int)
68
+ assert config.expected_kv_heads_per_device == 8
69
+ assert isinstance(config.expected_kv_heads_per_device, int)
70
+ assert config.expected_head_dim == 128
71
+ assert isinstance(config.expected_head_dim, int)
72
+ assert config.model_kv_cache_dtypes == (ttnn.bfloat8_b,) * 32
73
+ assert config.request_state_fields == ()
74
+ with expect_error(dataclasses.FrozenInstanceError, "cannot assign to field"):
75
+ config.expected_num_layers = 1
76
+ with expect_error(ValueError, "expected_num_layers"):
77
+ VLLMAdapterConfig(
78
+ trace=trace,
79
+ paged_kv_cache=paged_kv_cache,
80
+ expected_num_layers=0,
81
+ expected_kv_heads_per_device=8,
82
+ expected_head_dim=128,
83
+ model_kv_cache_dtypes=(ttnn.bfloat8_b,),
84
+ )
85
+ with expect_error(TypeError, "must be a tuple"):
86
+ VLLMAdapterConfig(
87
+ trace=trace,
88
+ paged_kv_cache=paged_kv_cache,
89
+ expected_num_layers=32,
90
+ expected_kv_heads_per_device=8,
91
+ expected_head_dim=128,
92
+ model_kv_cache_dtypes=[ttnn.bfloat8_b],
93
+ )
94
+
95
+
96
+ @pytest.mark.parametrize(
97
+ ("overrides", "error_type", "message"),
98
+ [
99
+ ({"trace": object()}, TypeError, "TraceConfig"),
100
+ ({"paged_kv_cache": object()}, TypeError, "PagedKVCacheConfig"),
101
+ ({"expected_num_layers": 0}, ValueError, "positive integer"),
102
+ ({"expected_num_layers": True}, ValueError, "positive integer"),
103
+ ({"expected_kv_heads_per_device": 0}, ValueError, "positive integer"),
104
+ ({"expected_kv_heads_per_device": -1}, ValueError, "positive integer"),
105
+ ({"expected_kv_heads_per_device": True}, ValueError, "positive integer"),
106
+ ({"expected_kv_heads_per_device": 8.5}, ValueError, "positive integer"),
107
+ ({"expected_head_dim": -1}, ValueError, "positive integer"),
108
+ ({"expected_head_dim": 0}, ValueError, "positive integer"),
109
+ ({"expected_head_dim": True}, ValueError, "positive integer"),
110
+ ({"expected_head_dim": 128.5}, ValueError, "positive integer"),
111
+ ({"model_kv_cache_dtype": ()}, ValueError, "cannot be empty"),
112
+ ({"model_kv_cache_dtype": (ttnn.bfloat8_b,) * 2}, ValueError, "one dtype per model layer"),
113
+ ({"model_kv_cache_dtype": None}, TypeError, "model metadata"),
114
+ ],
115
+ )
116
+ def test_config_rejects_inconsistent_static_inputs(overrides, error_type, message, expect_error):
117
+ arguments = {
118
+ "trace": TraceConfig(mode="all"),
119
+ "paged_kv_cache": PagedKVCacheConfig(block_size=32, max_num_blocks=128, dtype=ttnn.bfloat8_b),
120
+ "expected_num_layers": 32,
121
+ "expected_kv_heads_per_device": 8,
122
+ "expected_head_dim": 128,
123
+ "model_kv_cache_dtype": ttnn.bfloat8_b,
124
+ }
125
+
126
+ with expect_error(error_type, message):
127
+ VLLMAdapterConfig.resolve(**(arguments | overrides))
128
+
129
+
130
+ def test_config_requires_exact_static_policy_types(expect_error):
131
+ class TraceConfigSubclass(TraceConfig):
132
+ pass
133
+
134
+ class PagedKVCacheConfigSubclass(PagedKVCacheConfig):
135
+ pass
136
+
137
+ arguments = {
138
+ "trace": TraceConfig(mode="all"),
139
+ "paged_kv_cache": PagedKVCacheConfig(block_size=32, max_num_blocks=128, dtype=ttnn.bfloat8_b),
140
+ "expected_num_layers": 32,
141
+ "model_kv_cache_dtype": ttnn.bfloat8_b,
142
+ }
143
+
144
+ with expect_error(TypeError, "TraceConfig"):
145
+ VLLMAdapterConfig.resolve(**(arguments | {"trace": TraceConfigSubclass(mode="all")}))
146
+ with expect_error(TypeError, "PagedKVCacheConfig"):
147
+ VLLMAdapterConfig.resolve(
148
+ **(
149
+ arguments
150
+ | {
151
+ "paged_kv_cache": PagedKVCacheConfigSubclass(
152
+ block_size=32,
153
+ max_num_blocks=128,
154
+ dtype=ttnn.bfloat8_b,
155
+ )
156
+ }
157
+ )
158
+ )
159
+
160
+
161
+ def test_adapter_is_plain_orchestration_with_one_config_surface(expect_error):
162
+ adapter = _adapter()
163
+
164
+ assert tuple(inspect.signature(VLLMAdapter).parameters) == ("config",)
165
+ assert vars(adapter) == {"config": adapter.config}
166
+ with expect_error(TypeError, "VLLMAdapterConfig"):
167
+ VLLMAdapter(config=None)
168
+
169
+
170
+ @pytest.mark.parametrize(
171
+ ("method_name", "expected"),
172
+ [
173
+ (
174
+ "normalize_prefill",
175
+ [
176
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
177
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
178
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
179
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
180
+ ("prompt_lens", inspect.Parameter.KEYWORD_ONLY, None),
181
+ ("start_pos", inspect.Parameter.KEYWORD_ONLY, None),
182
+ ("empty_slots", inspect.Parameter.KEYWORD_ONLY, None),
183
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, None),
184
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
185
+ ("compatibility_kwargs", inspect.Parameter.KEYWORD_ONLY, None),
186
+ ],
187
+ ),
188
+ (
189
+ "normalize_decode",
190
+ [
191
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
192
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
193
+ ("start_pos", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
194
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
195
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
196
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, None),
197
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
198
+ ("reset_batch", inspect.Parameter.KEYWORD_ONLY, False),
199
+ ("compatibility_kwargs", inspect.Parameter.KEYWORD_ONLY, None),
200
+ ],
201
+ ),
202
+ ],
203
+ )
204
+ def test_normalizer_signatures_are_explicit(method_name, expected):
205
+ parameters = inspect.signature(getattr(VLLMAdapter, method_name)).parameters
206
+
207
+ assert [(name, parameter.kind, parameter.default) for name, parameter in parameters.items()] == expected
208
+
209
+
210
+ def test_normalized_typed_dicts_have_stable_key_order_and_optional_request_state():
211
+ assert tuple(NormalizedPrefillKwargs.__annotations__) == (
212
+ "tokens",
213
+ "page_table",
214
+ "prompt_lens",
215
+ "start_pos",
216
+ "empty_slots",
217
+ "kv_cache",
218
+ "sampling_params",
219
+ "prompt_tokens",
220
+ "output_tokens",
221
+ "slot_remap",
222
+ )
223
+ assert NormalizedPrefillKwargs.__required_keys__ == frozenset(
224
+ {"tokens", "page_table", "prompt_lens", "start_pos", "empty_slots", "kv_cache", "sampling_params"}
225
+ )
226
+ assert NormalizedPrefillKwargs.__optional_keys__ == frozenset({"prompt_tokens", "output_tokens", "slot_remap"})
227
+ assert tuple(NormalizedDecodeKwargs.__annotations__) == (
228
+ "tokens",
229
+ "start_pos",
230
+ "page_table",
231
+ "kv_cache",
232
+ "sampling_params",
233
+ "reset_batch",
234
+ "prompt_tokens",
235
+ "output_tokens",
236
+ "slot_remap",
237
+ )
238
+ assert NormalizedDecodeKwargs.__required_keys__ == frozenset(
239
+ {"tokens", "start_pos", "page_table", "kv_cache", "sampling_params", "reset_batch"}
240
+ )
241
+ assert NormalizedDecodeKwargs.__optional_keys__ == frozenset({"prompt_tokens", "output_tokens", "slot_remap"})
242
+
243
+
244
+ def test_normalize_prefill_positional_call_without_mutating_caller_kwargs():
245
+ adapter = _adapter()
246
+ compatibility_kwargs = {
247
+ "page_tables_per_layer": object(),
248
+ "prompt_tokens": object(),
249
+ "output_tokens": object(),
250
+ "slot_remap": object(),
251
+ "rope_deltas_all_users": object(),
252
+ }
253
+
254
+ normalized, enable_trace = adapter.normalize_prefill(
255
+ [[1, 2, 3, 4], [5, 6, 0, 0]],
256
+ [[0, 1], [2, 3]],
257
+ enable_trace=True,
258
+ prompt_lens=[4, 3],
259
+ start_pos=[0, 1],
260
+ sampling_params="sampling",
261
+ compatibility_kwargs=compatibility_kwargs,
262
+ )
263
+
264
+ assert tuple(normalized) == tuple(NormalizedPrefillKwargs.__annotations__)
265
+ assert normalized["tokens"].dtype == torch.long
266
+ assert normalized["page_table"].dtype == torch.int32
267
+ assert normalized["prompt_lens"].dtype == torch.long
268
+ assert normalized["start_pos"].dtype == torch.long
269
+ assert normalized["empty_slots"] is None
270
+ assert normalized["kv_cache"] is None
271
+ assert normalized["sampling_params"] == "sampling"
272
+ assert normalized["prompt_tokens"] is compatibility_kwargs["prompt_tokens"]
273
+ assert normalized["output_tokens"] is compatibility_kwargs["output_tokens"]
274
+ assert normalized["slot_remap"] is compatibility_kwargs["slot_remap"]
275
+ assert enable_trace is True
276
+ assert tuple(compatibility_kwargs) == (
277
+ "page_tables_per_layer",
278
+ "prompt_tokens",
279
+ "output_tokens",
280
+ "slot_remap",
281
+ "rope_deltas_all_users",
282
+ )
283
+
284
+
285
+ def test_normalize_decode_converts_existing_tensors_and_flattens_column_tokens():
286
+ adapter = _adapter(trace=TraceConfig(mode="decode_only"))
287
+
288
+ normalized, enable_trace = adapter.normalize_decode(
289
+ torch.tensor([[1], [2]], dtype=torch.int32),
290
+ torch.tensor([3, 4], dtype=torch.int32),
291
+ torch.tensor([[0], [1]], dtype=torch.int64),
292
+ enable_trace=True,
293
+ compatibility_kwargs={"slot_remap": [0, 1]},
294
+ )
295
+
296
+ assert tuple(normalized) == (
297
+ "tokens",
298
+ "start_pos",
299
+ "page_table",
300
+ "kv_cache",
301
+ "sampling_params",
302
+ "reset_batch",
303
+ "slot_remap",
304
+ )
305
+ assert normalized["tokens"].shape == (2,)
306
+ assert normalized["tokens"].dtype == torch.long
307
+ assert normalized["start_pos"].dtype == torch.long
308
+ assert normalized["page_table"].dtype == torch.int32
309
+ assert normalized["kv_cache"] is None
310
+ assert normalized["sampling_params"] is None
311
+ assert "prompt_tokens" not in normalized
312
+ assert "output_tokens" not in normalized
313
+ assert normalized["slot_remap"] == [0, 1]
314
+ assert normalized["reset_batch"] is False
315
+ assert enable_trace is True
316
+
317
+
318
+ @pytest.mark.parametrize("method_name", ["normalize_prefill", "normalize_decode"])
319
+ @pytest.mark.parametrize(
320
+ "compatibility_kwargs",
321
+ (None, {"prompt_tokens": None, "output_tokens": None, "slot_remap": None}),
322
+ )
323
+ def test_normalize_omits_unsupplied_or_none_request_state(method_name, compatibility_kwargs):
324
+ adapter = _adapter(trace=TraceConfig(mode="all"), request_state_fields=())
325
+ args = (
326
+ (torch.zeros((1, 1)), torch.zeros((1, 1)))
327
+ if method_name == "normalize_prefill"
328
+ else (torch.zeros(1), torch.zeros(1), torch.zeros((1, 1)))
329
+ )
330
+
331
+ normalized, _ = getattr(adapter, method_name)(
332
+ *args,
333
+ enable_trace=True,
334
+ compatibility_kwargs=compatibility_kwargs,
335
+ )
336
+
337
+ assert not ({"prompt_tokens", "output_tokens", "slot_remap"} & normalized.keys())
338
+
339
+
340
+ class _NarrowQwenRequestTarget:
341
+ def __init__(self):
342
+ self.calls = []
343
+
344
+ def prefill_forward(
345
+ self,
346
+ tokens,
347
+ page_table,
348
+ *,
349
+ prompt_lens=None,
350
+ start_pos=None,
351
+ empty_slots=None,
352
+ kv_cache=None,
353
+ sampling_params=None,
354
+ execution=None,
355
+ ):
356
+ self.calls.append(
357
+ (
358
+ "prefill",
359
+ {
360
+ "tokens": tokens,
361
+ "page_table": page_table,
362
+ "prompt_lens": prompt_lens,
363
+ "start_pos": start_pos,
364
+ "empty_slots": empty_slots,
365
+ "kv_cache": kv_cache,
366
+ "sampling_params": sampling_params,
367
+ "execution": execution,
368
+ },
369
+ )
370
+ )
371
+ return "prefill"
372
+
373
+ def decode_forward(
374
+ self,
375
+ tokens,
376
+ start_pos,
377
+ page_table,
378
+ *,
379
+ kv_cache=None,
380
+ sampling_params=None,
381
+ reset_batch=False,
382
+ read_from_device=True,
383
+ execution=None,
384
+ ):
385
+ self.calls.append(
386
+ (
387
+ "decode",
388
+ {
389
+ "tokens": tokens,
390
+ "start_pos": start_pos,
391
+ "page_table": page_table,
392
+ "kv_cache": kv_cache,
393
+ "sampling_params": sampling_params,
394
+ "reset_batch": reset_batch,
395
+ "read_from_device": read_from_device,
396
+ "execution": execution,
397
+ },
398
+ )
399
+ )
400
+ return "decode"
401
+
402
+
403
+ @pytest.mark.parametrize(
404
+ "generator_class",
405
+ (Qwen2Generator, Qwen25Generator, Qwen25_72BGenerator, Qwen25Coder32BGenerator),
406
+ )
407
+ def test_qwen_generator_dispatch_omits_absent_state_for_narrow_request_surface(generator_class):
408
+ generator = object.__new__(generator_class)
409
+ generator._adapter = _adapter(request_state_fields=())
410
+ generator.target = _NarrowQwenRequestTarget()
411
+ prefill_execution = object()
412
+ decode_execution = object()
413
+ generator._select_prefill_execution = lambda normalized, requested: prefill_execution
414
+ generator._select_execution = lambda operation, requested: decode_execution
415
+
416
+ request_state = {"prompt_tokens": object(), "output_tokens": object(), "slot_remap": [0]}
417
+ assert generator.prefill_forward([[1]], [[0]], enable_trace=False, **request_state) == "prefill"
418
+ assert generator.decode_forward([1], [0], [[0]], enable_trace=False, **request_state) == "decode"
419
+
420
+ assert [name for name, _ in generator.target.calls] == ["prefill", "decode"]
421
+ assert generator.target.calls[0][1]["execution"] is prefill_execution
422
+ assert generator.target.calls[1][1]["execution"] is decode_execution
423
+
424
+
425
+ @pytest.mark.parametrize(
426
+ "generator_module",
427
+ (llama3_generator_module, llama70_generator_module, qwen3_generator_module),
428
+ )
429
+ def test_stateful_generator_adapter_reuses_executor_request_state_allowlist(monkeypatch, generator_module):
430
+ request_state_fields = ("prompt_tokens", "output_tokens", "slot_remap")
431
+ lane = SimpleNamespace(
432
+ model=object(),
433
+ config=SimpleNamespace(
434
+ trace=TraceConfig(mode="decode_only"),
435
+ paged_kv_cache=PagedKVCacheConfig(block_size=32, max_num_blocks=128, dtype=ttnn.bfloat8_b),
436
+ ),
437
+ _request_state_fields=request_state_fields,
438
+ )
439
+ monkeypatch.setattr(
440
+ generator_module,
441
+ "_model_kv_metadata",
442
+ lambda model: ((ttnn.bfloat8_b,), 1, 1, 128),
443
+ )
444
+
445
+ adapter = generator_module._build_vllm_adapter(lane)
446
+
447
+ assert adapter.config.request_state_fields == request_state_fields
448
+
449
+
450
+ @pytest.mark.parametrize(
451
+ ("method_name", "args", "trace", "hint"),
452
+ [
453
+ (
454
+ "normalize_prefill",
455
+ (torch.zeros((1, 1)), torch.zeros((1, 1))),
456
+ TraceConfig(mode="decode_only"),
457
+ True,
458
+ ),
459
+ (
460
+ "normalize_decode",
461
+ (torch.zeros(1), torch.zeros(1), torch.zeros((1, 1))),
462
+ TraceConfig(mode="none"),
463
+ True,
464
+ ),
465
+ ],
466
+ )
467
+ def test_normalize_rejects_trace_hint_that_disagrees_with_static_policy(method_name, args, trace, hint, expect_error):
468
+ adapter = _adapter(trace=trace)
469
+
470
+ with expect_error(ValueError, "enable_trace"):
471
+ getattr(adapter, method_name)(*args, enable_trace=hint)
472
+
473
+
474
+ def test_normalize_accepts_unselected_operation_with_available_trace_target():
475
+ adapter = _adapter(trace=TraceConfig(mode="all"))
476
+
477
+ _, enable_trace = adapter.normalize_prefill(
478
+ torch.zeros((1, 1)),
479
+ torch.zeros((1, 1)),
480
+ enable_trace=False,
481
+ )
482
+
483
+ assert enable_trace is False
484
+
485
+
486
+ @pytest.mark.parametrize(
487
+ ("mode", "prefill_enabled", "decode_enabled"),
488
+ [
489
+ ("none", False, False),
490
+ ("decode_only", False, True),
491
+ ("all", True, True),
492
+ ],
493
+ )
494
+ def test_normalize_accepts_operation_specific_static_trace_policy(mode, prefill_enabled, decode_enabled):
495
+ adapter = _adapter(trace=TraceConfig(mode=mode))
496
+
497
+ _, normalized_prefill_trace = adapter.normalize_prefill(
498
+ torch.zeros((1, 1)),
499
+ torch.zeros((1, 1)),
500
+ enable_trace=prefill_enabled,
501
+ )
502
+ _, normalized_decode_trace = adapter.normalize_decode(
503
+ torch.zeros(1),
504
+ torch.zeros(1),
505
+ torch.zeros((1, 1)),
506
+ enable_trace=decode_enabled,
507
+ )
508
+
509
+ assert normalized_prefill_trace is prefill_enabled
510
+ assert normalized_decode_trace is decode_enabled
511
+
512
+
513
+ @pytest.mark.parametrize("hint", [None, "true", 1])
514
+ def test_normalize_requires_an_explicit_boolean_trace_selection(hint, expect_error):
515
+ adapter = _adapter()
516
+
517
+ with expect_error(TypeError, "enable_trace"):
518
+ if hint is None:
519
+ adapter.normalize_prefill(torch.zeros((1, 1)), torch.zeros((1, 1)))
520
+ else:
521
+ adapter.normalize_prefill(torch.zeros((1, 1)), torch.zeros((1, 1)), enable_trace=hint)
522
+
523
+
524
+ def test_normalize_rejects_duplicate_positional_and_keyword_argument(expect_error):
525
+ adapter = _adapter()
526
+
527
+ with expect_error(TypeError, "tokens"):
528
+ adapter.normalize_prefill(
529
+ torch.zeros((1, 1)),
530
+ torch.zeros((1, 1)),
531
+ tokens=torch.zeros((1, 1)),
532
+ enable_trace=True,
533
+ )
534
+
535
+
536
+ @pytest.mark.parametrize("method_name", ["normalize_prefill", "normalize_decode"])
537
+ def test_normalize_rejects_unknown_compatibility_keys(method_name, expect_error):
538
+ adapter = _adapter()
539
+ args = (
540
+ (torch.zeros((1, 1)), torch.zeros((1, 1)))
541
+ if method_name == "normalize_prefill"
542
+ else (torch.zeros(1), torch.zeros(1), torch.zeros((1, 1)))
543
+ )
544
+
545
+ with expect_error(TypeError, "unexpected keyword argument 'unknown_plugin_field'"):
546
+ getattr(adapter, method_name)(
547
+ *args,
548
+ enable_trace=True,
549
+ compatibility_kwargs={"unknown_plugin_field": object()},
550
+ )
551
+
552
+
553
+ def _signature_entries(method):
554
+ return [
555
+ (name, parameter.kind, parameter.default) for name, parameter in inspect.signature(method).parameters.items()
556
+ ]
557
+
558
+
559
+ @pytest.mark.parametrize(
560
+ ("method_name", "expected"),
561
+ [
562
+ (
563
+ "compile_prefill",
564
+ [
565
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
566
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
567
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
568
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
569
+ ("prompt_lens", inspect.Parameter.KEYWORD_ONLY, None),
570
+ ("start_pos", inspect.Parameter.KEYWORD_ONLY, None),
571
+ ("empty_slots", inspect.Parameter.KEYWORD_ONLY, None),
572
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, None),
573
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
574
+ ],
575
+ ),
576
+ (
577
+ "compile_decode",
578
+ [
579
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
580
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
581
+ ("start_pos", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
582
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
583
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
584
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, None),
585
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
586
+ ("reset_batch", inspect.Parameter.KEYWORD_ONLY, False),
587
+ ],
588
+ ),
589
+ (
590
+ "prefill_forward",
591
+ [
592
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
593
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
594
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
595
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
596
+ ("prompt_lens", inspect.Parameter.KEYWORD_ONLY, None),
597
+ ("start_pos", inspect.Parameter.KEYWORD_ONLY, None),
598
+ ("empty_slots", inspect.Parameter.KEYWORD_ONLY, None),
599
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, None),
600
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
601
+ ("compatibility_kwargs", inspect.Parameter.VAR_KEYWORD, inspect.Parameter.empty),
602
+ ],
603
+ ),
604
+ (
605
+ "decode_forward",
606
+ [
607
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
608
+ ("tokens", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
609
+ ("start_pos", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
610
+ ("page_table", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
611
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
612
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, None),
613
+ ("sampling_params", inspect.Parameter.KEYWORD_ONLY, None),
614
+ ("reset_batch", inspect.Parameter.KEYWORD_ONLY, False),
615
+ ("read_from_device", inspect.Parameter.KEYWORD_ONLY, True),
616
+ ("compatibility_kwargs", inspect.Parameter.VAR_KEYWORD, inspect.Parameter.empty),
617
+ ],
618
+ ),
619
+ (
620
+ "read_decode_output",
621
+ [
622
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
623
+ ("tt_out", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
624
+ ("async_read", inspect.Parameter.KEYWORD_ONLY, False),
625
+ ],
626
+ ),
627
+ (
628
+ "process_decode_output_host",
629
+ [
630
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
631
+ ("tt_out", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
632
+ ("is_tokens", inspect.Parameter.KEYWORD_ONLY, False),
633
+ ],
634
+ ),
635
+ (
636
+ "warmup_model_prefill",
637
+ [
638
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
639
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
640
+ ("can_sample_on_device", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
641
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
642
+ ],
643
+ ),
644
+ (
645
+ "warmup_model_decode",
646
+ [
647
+ ("self", inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.empty),
648
+ ("kv_cache", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
649
+ ("max_batch_size", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
650
+ ("num_blocks", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
651
+ ("can_sample_on_device", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
652
+ ("enable_trace", inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.empty),
653
+ ],
654
+ ),
655
+ ],
656
+ )
657
+ def test_registered_generator_signatures_are_exact(method_name, expected):
658
+ assert _signature_entries(getattr(Llama3Generator, method_name)) == expected
659
+
660
+
661
+ def test_registered_generator_sizing_signature_has_no_compatibility_bag():
662
+ assert _signature_entries(Llama3Generator.get_max_tokens_all_users) == [
663
+ ("model_name", inspect.Parameter.POSITIONAL_OR_KEYWORD, ""),
664
+ ("num_devices", inspect.Parameter.POSITIONAL_OR_KEYWORD, 1),
665
+ ("tt_data_parallel", inspect.Parameter.POSITIONAL_OR_KEYWORD, 1),
666
+ ("max_model_len", inspect.Parameter.POSITIONAL_OR_KEYWORD, 0),
667
+ ("max_num_seqs", inspect.Parameter.POSITIONAL_OR_KEYWORD, 1),
668
+ ]
669
+
670
+
671
+ def test_registered_generator_compatibility_bags_exist_only_on_forward_methods():
672
+ variadic_keyword_methods = {
673
+ method_name
674
+ for method_name, method in vars(Llama3Generator).items()
675
+ if inspect.isfunction(method)
676
+ and any(
677
+ parameter.kind is inspect.Parameter.VAR_KEYWORD
678
+ for parameter in inspect.signature(method).parameters.values()
679
+ )
680
+ }
681
+
682
+ assert variadic_keyword_methods == {"prefill_forward", "decode_forward"}
683
+
684
+
685
+ class _ExplicitGeneratorTarget:
686
+ model = SimpleNamespace()
687
+ model_args = object()
688
+ mesh_device = object()
689
+ cache_path = "cache"
690
+ already_warmed_up_prefill = False
691
+ eager_execution = object()
692
+ traced_prefill_execution = object()
693
+ traced_decode_execution = object()
694
+
695
+ def __init__(self):
696
+ self.calls = []
697
+
698
+ def _record(self, method_name, arguments):
699
+ self.calls.append(
700
+ (
701
+ method_name,
702
+ {name: value for name, value in arguments.items() if name != "self"},
703
+ )
704
+ )
705
+
706
+ def can_trace_prefill(
707
+ self,
708
+ *,
709
+ tokens, # ↓ Core request
710
+ prompt_lens=None, # ↓ Sequence metadata
711
+ start_pos=None,
712
+ empty_slots=None, # ↓ Lane routing
713
+ ):
714
+ self._record("can_trace_prefill", locals())
715
+ return True
716
+
717
+ def compile_prefill(
718
+ self,
719
+ tokens,
720
+ page_table,
721
+ *,
722
+ prompt_lens=None, # ↓ Sequence metadata
723
+ start_pos=None,
724
+ empty_slots=None, # ↓ Lane routing
725
+ kv_cache=None, # ↓ Borrowed resources
726
+ sampling_params=None, # ↓ Sampling
727
+ prompt_tokens=None, # ↓ Request-owned sampling state
728
+ output_tokens=None,
729
+ slot_remap=None,
730
+ execution=None, # ↓ Internal dispatch
731
+ ):
732
+ self._record("compile_prefill", locals())
733
+
734
+ def compile_decode(
735
+ self,
736
+ tokens,
737
+ start_pos,
738
+ page_table,
739
+ *,
740
+ kv_cache=None, # ↓ Borrowed resources
741
+ sampling_params=None, # ↓ Sampling
742
+ prompt_tokens=None, # ↓ Request-owned sampling state
743
+ output_tokens=None,
744
+ slot_remap=None,
745
+ reset_batch=False, # ↓ State transition
746
+ execution=None, # ↓ Internal dispatch
747
+ ):
748
+ self._record("compile_decode", locals())
749
+
750
+ def prefill_forward(
751
+ self,
752
+ tokens,
753
+ page_table,
754
+ *,
755
+ prompt_lens=None, # ↓ Sequence metadata
756
+ start_pos=None,
757
+ empty_slots=None, # ↓ Lane routing
758
+ kv_cache=None, # ↓ Borrowed resources
759
+ sampling_params=None, # ↓ Sampling
760
+ prompt_tokens=None, # ↓ Request-owned sampling state
761
+ output_tokens=None,
762
+ slot_remap=None,
763
+ execution=None, # ↓ Internal dispatch
764
+ ):
765
+ self._record("prefill_forward", locals())
766
+ return "prefill"
767
+
768
+ def decode_forward(
769
+ self,
770
+ tokens,
771
+ start_pos,
772
+ page_table,
773
+ *,
774
+ kv_cache=None, # ↓ Borrowed resources
775
+ sampling_params=None, # ↓ Sampling
776
+ prompt_tokens=None, # ↓ Request-owned sampling state
777
+ output_tokens=None,
778
+ slot_remap=None,
779
+ reset_batch=False, # ↓ State transition
780
+ read_from_device=True, # ↓ Output policy
781
+ execution=None, # ↓ Internal dispatch
782
+ ):
783
+ self._record("decode_forward", locals())
784
+ return "decode"
785
+
786
+
787
+ def test_registered_generator_compile_methods_normalize_and_select_execution():
788
+ target = _ExplicitGeneratorTarget()
789
+ generator = Llama3Generator(target, _adapter())
790
+ tokens = torch.tensor([[1, 2]])
791
+ page_table = torch.tensor([[0]])
792
+ prompt_lens = torch.tensor([2])
793
+ start_pos = torch.tensor([0])
794
+ kv_cache = object()
795
+ sampling_params = object()
796
+
797
+ generator.compile_prefill(
798
+ tokens,
799
+ page_table,
800
+ enable_trace=True,
801
+ prompt_lens=prompt_lens,
802
+ start_pos=start_pos,
803
+ empty_slots=[0],
804
+ kv_cache=kv_cache,
805
+ sampling_params=sampling_params,
806
+ )
807
+ assert target.calls[0][0] == "compile_prefill"
808
+ assert target.calls[0][1]["execution"] is target.traced_prefill_execution
809
+ assert set(target.calls[0][1]) == set(NormalizedPrefillKwargs.__annotations__) | {"execution"}
810
+
811
+ generator.compile_decode(
812
+ tokens[:, 0],
813
+ start_pos,
814
+ page_table,
815
+ enable_trace=True,
816
+ kv_cache=kv_cache,
817
+ sampling_params=sampling_params,
818
+ reset_batch=True,
819
+ )
820
+ assert target.calls[1][0] == "compile_decode"
821
+ assert target.calls[1][1]["execution"] is target.traced_decode_execution
822
+ assert set(target.calls[1][1]) == set(NormalizedDecodeKwargs.__annotations__) | {"execution"}
823
+
824
+
825
+ def test_registered_generator_discards_allowlisted_compatibility_and_limits_trace_classification():
826
+ target = _ExplicitGeneratorTarget()
827
+ generator = Llama3Generator(target, _adapter())
828
+ tokens = torch.tensor([[1, 2]])
829
+ page_table = torch.tensor([[0]])
830
+ prompt_lens = torch.tensor([2])
831
+ start_pos = torch.tensor([0])
832
+ kv_cache = object()
833
+ sampling_params = object()
834
+
835
+ assert (
836
+ generator.prefill_forward(
837
+ tokens,
838
+ page_table,
839
+ enable_trace=True,
840
+ prompt_lens=prompt_lens,
841
+ start_pos=start_pos,
842
+ empty_slots=[0],
843
+ kv_cache=kv_cache,
844
+ sampling_params=sampling_params,
845
+ page_tables_per_layer=object(),
846
+ )
847
+ == "prefill"
848
+ )
849
+ assert target.calls[0][0] == "prefill_forward"
850
+ assert target.calls[0][1]["execution"] is target.traced_prefill_execution
851
+ assert set(target.calls[0][1]) == set(NormalizedPrefillKwargs.__annotations__) | {"execution"}
852
+
853
+ assert (
854
+ generator.decode_forward(
855
+ tokens[:, 0],
856
+ start_pos,
857
+ page_table,
858
+ enable_trace=True,
859
+ kv_cache=kv_cache,
860
+ sampling_params=sampling_params,
861
+ reset_batch=True,
862
+ read_from_device=False,
863
+ slot_remap=[0],
864
+ )
865
+ == "decode"
866
+ )
867
+ assert target.calls[1][0] == "decode_forward"
868
+ assert target.calls[1][1]["execution"] is target.traced_decode_execution
869
+ assert target.calls[1][1]["read_from_device"] is False
870
+ assert set(target.calls[1][1]) == set(NormalizedDecodeKwargs.__annotations__) | {
871
+ "read_from_device",
872
+ "execution",
873
+ }
874
+
875
+
876
+ @pytest.mark.parametrize(
877
+ ("method_name", "args"),
878
+ [
879
+ ("prefill_forward", (torch.zeros((1, 1)), torch.zeros((1, 1)))),
880
+ ("decode_forward", (torch.zeros(1), torch.zeros(1), torch.zeros((1, 1)))),
881
+ ],
882
+ )
883
+ def test_registered_generator_rejects_unknown_compatibility_before_target_selection(
884
+ method_name,
885
+ args,
886
+ expect_error,
887
+ ):
888
+ target = _ExplicitGeneratorTarget()
889
+ generator = Llama3Generator(target, _adapter())
890
+
891
+ with expect_error(TypeError, "unexpected keyword argument 'unknown_plugin_field'"):
892
+ getattr(generator, method_name)(
893
+ *args,
894
+ enable_trace=True,
895
+ unknown_plugin_field=object(),
896
+ )
897
+
898
+ assert target.calls == []
899
+
900
+
901
+ def test_registered_generator_rejects_unknown_nonforward_keywords(expect_error):
902
+ with expect_error(TypeError, "unexpected keyword argument"):
903
+ Llama3Generator.get_max_tokens_all_users(unknown_plugin_field=True)
904
+
905
+
906
+ def test_registered_generator_forwards_output_and_warmup_arguments_by_name():
907
+ pending_output = object()
908
+ read_events = [object()]
909
+ target = create_autospec(Llama3Executor, instance=True)
910
+ target.read_decode_output.return_value = pending_output, read_events
911
+ target.process_decode_output_host.return_value = "tokens", "log-probs"
912
+ generator = Llama3Generator(target, adapter=object())
913
+ tt_out = object()
914
+ kv_cache = object()
915
+
916
+ assert generator.read_decode_output(tt_out, async_read=True) == (pending_output, read_events)
917
+ target.read_decode_output.assert_called_once_with(tt_out=tt_out, async_read=True)
918
+ assert generator.process_decode_output_host(pending_output, is_tokens=True) == ("tokens", "log-probs")
919
+ target.process_decode_output_host.assert_called_once_with(tt_out=pending_output, is_tokens=True)
920
+
921
+ generator.warmup_model_prefill(
922
+ kv_cache=kv_cache,
923
+ can_sample_on_device=True,
924
+ enable_trace=False,
925
+ )
926
+ target.warmup_model_prefill.assert_called_once_with(
927
+ kv_cache=kv_cache,
928
+ can_sample_on_device=True,
929
+ enable_trace=False,
930
+ )
931
+ generator.warmup_model_decode(
932
+ kv_cache=kv_cache,
933
+ max_batch_size=8,
934
+ num_blocks=128,
935
+ can_sample_on_device=False,
936
+ enable_trace=True,
937
+ )
938
+ target.warmup_model_decode.assert_called_once_with(
939
+ kv_cache=kv_cache,
940
+ max_batch_size=8,
941
+ num_blocks=128,
942
+ can_sample_on_device=False,
943
+ enable_trace=True,
944
+ )
945
+
946
+
947
+ def test_resolve_legacy_kv_cache_returns_new_immutable_config():
948
+ base = PagedKVCacheConfig(block_size=32, max_num_blocks=128, dtype=ttnn.bfloat8_b)
949
+ adapter = _adapter(paged_config=base)
950
+
951
+ resolved = adapter.resolve_legacy_kv_cache_config(
952
+ (129, 8, 64, 128),
953
+ torch.bfloat16,
954
+ 32,
955
+ )
956
+
957
+ assert resolved is not base
958
+ assert base.num_blocks is None
959
+ assert base.block_size == 32
960
+ assert base.max_num_blocks == 128
961
+ assert resolved.num_blocks == 129
962
+ assert resolved.block_size == 64
963
+ assert resolved.max_num_blocks == 129
964
+ assert resolved.dtype == base.dtype
965
+ assert resolved.memory_config == base.memory_config
966
+
967
+
968
+ @pytest.mark.parametrize(
969
+ ("shape", "dtype", "num_layers", "message"),
970
+ [
971
+ ((64, 8, 0, 128), torch.bfloat16, 32, "block_size"),
972
+ ((64, 4, 32, 128), torch.bfloat16, 32, "KV heads"),
973
+ ((64, 8, 32, 64), torch.bfloat16, 32, "head dimension"),
974
+ ((0, 8, 32, 128), torch.bfloat16, 32, "num_blocks"),
975
+ ((64, 8, 32, 128), torch.float32, 32, "dtype"),
976
+ ((64, 8, 32, 128), torch.bfloat16, 31, "layer count"),
977
+ ],
978
+ )
979
+ def test_resolve_legacy_kv_cache_rejects_mismatched_vllm_spec(shape, dtype, num_layers, message, expect_error):
980
+ adapter = _adapter()
981
+
982
+ with expect_error((TypeError, ValueError), message):
983
+ adapter.resolve_legacy_kv_cache_config(shape, dtype, num_layers)
984
+
985
+
986
+ def test_adapter_rejects_static_dtype_that_disagrees_with_model_owned_dtype(expect_error):
987
+ with expect_error(ValueError, "model-owned"):
988
+ _adapter(model_dtype=ttnn.bfloat16)
989
+
990
+
991
+ def test_adapter_requires_explicit_model_owned_dtype_metadata(expect_error):
992
+ with expect_error(TypeError, "must be supplied from model metadata"):
993
+ _adapter(model_dtype=None)
994
+
995
+
996
+ def test_bfloat4_model_dtype_uses_shared_bfloat16_torch_surrogate():
997
+ config = PagedKVCacheConfig(
998
+ block_size=32,
999
+ max_num_blocks=128,
1000
+ dtype=ttnn.bfloat4_b,
1001
+ )
1002
+ adapter = _adapter(paged_config=config, model_dtype=ttnn.bfloat4_b)
1003
+
1004
+ resolved = adapter.resolve_legacy_kv_cache_config(
1005
+ (64, 8, 32, 128),
1006
+ torch.bfloat16,
1007
+ 32,
1008
+ )
1009
+
1010
+ assert resolved.dtype == ttnn.bfloat4_b
1011
+ assert resolved.num_blocks == 64
1012
+
1013
+
1014
+ def test_resolve_legacy_kv_cache_rejects_replacing_resolved_capacity(expect_error):
1015
+ adapter = _adapter(
1016
+ paged_config=PagedKVCacheConfig(
1017
+ block_size=32,
1018
+ max_num_blocks=128,
1019
+ dtype=ttnn.bfloat8_b,
1020
+ num_blocks=32,
1021
+ )
1022
+ )
1023
+
1024
+ with expect_error(ValueError, "already resolved"):
1025
+ adapter.resolve_legacy_kv_cache_config((64, 8, 32, 128), torch.bfloat16, 32)
code/models/common/tests/llm_runtime/test_warmup.py ADDED
@@ -0,0 +1,1012 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from __future__ import annotations
5
+
6
+ import inspect
7
+ from dataclasses import replace
8
+ from types import SimpleNamespace
9
+ from typing import Any, Sequence
10
+
11
+ import pytest
12
+ import torch
13
+
14
+ from models.common.llm_runtime.config import PageTableLayout, TraceConfig, WarmupConfig
15
+ from models.common.llm_runtime.decode import DecodeRuntimeConfig
16
+ from models.common.llm_runtime.output_reader import OutputReader
17
+ from models.common.llm_runtime.prefill.config import PrefillRuntimeConfig
18
+ from models.common.llm_runtime.program_compiler import CompiledProgram, OutputSpec, ProgramKey
19
+ from models.common.llm_runtime.warmup import (
20
+ CoverageAlias,
21
+ WarmupCoordinator,
22
+ WarmupCoordinatorConfig,
23
+ _resolve_coverage_manifest,
24
+ )
25
+
26
+
27
+ class RecordingExecution:
28
+ def __init__(self, events=None):
29
+ self.prefill_calls = []
30
+ self.decode_calls = []
31
+ self.events = events if events is not None else []
32
+ self.fail_decode_call = None
33
+ self.prefill_replays = []
34
+
35
+ def compile_prefill(
36
+ self,
37
+ *,
38
+ tokens: torch.Tensor, # ↓ Core request
39
+ page_table: torch.Tensor,
40
+ prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata
41
+ start_pos: torch.Tensor | None = None,
42
+ empty_slots: Sequence[int] | None = None, # ↓ Lane routing
43
+ sampling_params: Any = None, # ↓ Sampling
44
+ ) -> None:
45
+ self.events.append("compile_prefill")
46
+ self.prefill_calls.append(
47
+ {
48
+ "tokens": tokens,
49
+ "page_table": page_table,
50
+ "prompt_lens": prompt_lens,
51
+ "start_pos": start_pos,
52
+ "empty_slots": empty_slots,
53
+ "sampling_params": sampling_params,
54
+ }
55
+ )
56
+
57
+ def compile_decode(
58
+ self,
59
+ *,
60
+ tokens: torch.Tensor, # ↓ Core request
61
+ start_pos: torch.Tensor,
62
+ page_table: torch.Tensor,
63
+ sampling_params: Any = None, # ↓ Sampling
64
+ reset_batch: bool = False, # ↓ State transition
65
+ ) -> None:
66
+ call = len(self.decode_calls) + 1
67
+ self.events.append("compile_decode")
68
+ if call == self.fail_decode_call:
69
+ self.fail_decode_call = None
70
+ raise RuntimeError("decode compile failed")
71
+ self.decode_calls.append(
72
+ {
73
+ "tokens": tokens,
74
+ "start_pos": start_pos,
75
+ "page_table": page_table,
76
+ "sampling_params": sampling_params,
77
+ "reset_batch": reset_batch,
78
+ }
79
+ )
80
+
81
+ def prefill_forward(
82
+ self,
83
+ *,
84
+ tokens: torch.Tensor, # ↓ Core request
85
+ page_table: torch.Tensor,
86
+ prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata
87
+ start_pos: torch.Tensor | None = None,
88
+ empty_slots: Sequence[int] | None = None, # ↓ Lane routing
89
+ sampling_params: Any = None, # ↓ Sampling
90
+ ) -> None:
91
+ self.events.append("prefill_replay")
92
+ self.prefill_replays.append(
93
+ {
94
+ "tokens": tokens,
95
+ "page_table": page_table,
96
+ "prompt_lens": prompt_lens,
97
+ "start_pos": start_pos,
98
+ "empty_slots": empty_slots,
99
+ "sampling_params": sampling_params,
100
+ }
101
+ )
102
+
103
+
104
+ class RecordingTraceCompiler:
105
+ def __init__(self, events=None):
106
+ self.calls = 0
107
+ self.events = events if events is not None else []
108
+
109
+ def capture_all(self):
110
+ self.events.append("capture")
111
+ self.calls += 1
112
+
113
+
114
+ class Mesh:
115
+ shape = (1, 1)
116
+
117
+
118
+ def make_runtime_configs(
119
+ *,
120
+ sampling=True,
121
+ lane_capacity=4,
122
+ allow_force_argmax=True,
123
+ page_table_layout=None,
124
+ sampling_config=None,
125
+ model=None,
126
+ ):
127
+ mesh = Mesh()
128
+ sampling_config = sampling_config or SimpleNamespace(
129
+ allow_force_argmax=allow_force_argmax,
130
+ max_top_k=32,
131
+ )
132
+ sampling_config.max_batch_size = lane_capacity
133
+ model = model or SimpleNamespace(
134
+ config=SimpleNamespace(max_batch_size=lane_capacity, mesh_device=mesh, num_devices=1),
135
+ sampling=SimpleNamespace(
136
+ config=sampling_config,
137
+ decode_forward=lambda logits, *, k=None, p=None, temp=None, seeds=None, tt_out_tok=None, enable_log_probs=False: None,
138
+ ),
139
+ vocab_size=128,
140
+ )
141
+ mesh = model.config.mesh_device
142
+ layout = page_table_layout or PageTableLayout(
143
+ block_size=32,
144
+ raw_capacity_width=128,
145
+ prefill_width=192,
146
+ decode_width=128,
147
+ )
148
+ output_reader = OutputReader(mesh)
149
+ return (
150
+ PrefillRuntimeConfig.resolve(
151
+ model=model,
152
+ output_reader=output_reader,
153
+ page_table_layout=layout,
154
+ max_batch_size=lane_capacity,
155
+ max_prefill_chunk_size=128,
156
+ device_sampling_enabled=sampling,
157
+ can_enable_trace=lambda _sequence_length, _batch_size: True,
158
+ ),
159
+ DecodeRuntimeConfig.resolve(
160
+ model=model,
161
+ output_reader=output_reader,
162
+ lane_capacity=lane_capacity,
163
+ page_table_layout=layout,
164
+ device_sampling_enabled=sampling,
165
+ ),
166
+ )
167
+
168
+
169
+ def make_coordinator(
170
+ *,
171
+ trace_mode="all",
172
+ sampling=True,
173
+ warmup_config=None,
174
+ sequence_lengths=(128, 1024),
175
+ lane_capacity=4,
176
+ execution=None,
177
+ trace_compiler=None,
178
+ events=None,
179
+ allow_force_argmax=True,
180
+ page_table_layout=None,
181
+ sampling_config=None,
182
+ ):
183
+ events = events if events is not None else []
184
+ execution = execution or RecordingExecution(events)
185
+ if trace_compiler is None and trace_mode != "none":
186
+ trace_compiler = RecordingTraceCompiler(events)
187
+ sampling_calls = []
188
+ bound_calls = []
189
+
190
+ def ensure_sampling():
191
+ events.append("sampling")
192
+ sampling_calls.append(True)
193
+
194
+ def validate_bound(value):
195
+ bound_calls.append(value)
196
+
197
+ layout = page_table_layout or PageTableLayout(
198
+ block_size=32,
199
+ raw_capacity_width=128,
200
+ prefill_width=192,
201
+ decode_width=128,
202
+ )
203
+ prefill_config, decode_config = make_runtime_configs(
204
+ sampling=sampling,
205
+ lane_capacity=lane_capacity,
206
+ allow_force_argmax=allow_force_argmax,
207
+ page_table_layout=layout,
208
+ sampling_config=sampling_config,
209
+ )
210
+ execution.prefill = SimpleNamespace(config=prefill_config)
211
+ execution.decode = SimpleNamespace(config=decode_config)
212
+ execution.eager_executor = execution
213
+ execution.trace_compiler = trace_compiler
214
+
215
+ coordinator = WarmupCoordinator(
216
+ config=WarmupCoordinatorConfig.resolve(
217
+ warmup=warmup_config or WarmupConfig(),
218
+ trace=TraceConfig(trace_mode),
219
+ prefill=prefill_config,
220
+ decode=decode_config,
221
+ prefill_sequence_lengths=sequence_lengths,
222
+ ),
223
+ execution=execution,
224
+ ensure_sampling_buffers=ensure_sampling,
225
+ validate_bound_cache=validate_bound,
226
+ )
227
+ return coordinator, execution, trace_compiler, sampling_calls, bound_calls, events
228
+
229
+
230
+ @pytest.mark.parametrize(
231
+ ("method_name", "parameter_names"),
232
+ [
233
+ ("warmup_prefill", ("self", "kv_cache", "can_sample_on_device", "enable_trace")),
234
+ (
235
+ "warmup_decode",
236
+ ("self", "kv_cache", "max_batch_size", "num_blocks", "can_sample_on_device", "enable_trace"),
237
+ ),
238
+ ],
239
+ )
240
+ def test_warmup_signatures_match_registered_plugin_contract(method_name, parameter_names):
241
+ parameters = inspect.signature(getattr(WarmupCoordinator, method_name)).parameters
242
+
243
+ assert tuple(parameters) == parameter_names
244
+ assert parameters["self"].kind is inspect.Parameter.POSITIONAL_OR_KEYWORD
245
+ for name in parameter_names[1:]:
246
+ assert parameters[name].kind is inspect.Parameter.KEYWORD_ONLY
247
+ assert parameters[name].default is inspect.Parameter.empty
248
+
249
+
250
+ def test_registered_plugin_warmup_calls_validate_cache_without_forwarding_it():
251
+ cache = object()
252
+ coordinator, execution, _, _, bound_calls, _ = make_coordinator(
253
+ trace_mode="none",
254
+ sampling=False,
255
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
256
+ sequence_lengths=(128,),
257
+ lane_capacity=1,
258
+ )
259
+ prefill_kwargs = {
260
+ "kv_cache": cache,
261
+ "can_sample_on_device": False,
262
+ }
263
+ decode_kwargs = {
264
+ "kv_cache": cache,
265
+ "max_batch_size": 1,
266
+ "num_blocks": 8,
267
+ "can_sample_on_device": False,
268
+ }
269
+
270
+ coordinator.warmup_prefill(enable_trace=False, **prefill_kwargs)
271
+ coordinator.warmup_decode(enable_trace=False, **decode_kwargs)
272
+
273
+ assert coordinator.coverage_manifest is None
274
+ assert bound_calls == [cache, cache]
275
+ assert all(
276
+ tuple(call) == ("tokens", "page_table", "prompt_lens", "start_pos", "empty_slots", "sampling_params")
277
+ for call in execution.prefill_calls
278
+ )
279
+ assert all(
280
+ tuple(call) == ("tokens", "start_pos", "page_table", "sampling_params", "reset_batch")
281
+ for call in execution.decode_calls
282
+ )
283
+
284
+
285
+ @pytest.mark.parametrize(
286
+ ("method_name", "plugin_kwargs", "unexpected_name"),
287
+ [
288
+ (
289
+ "warmup_prefill",
290
+ {"kv_cache": "cache", "can_sample_on_device": False, "enable_trace": False},
291
+ "greedy_only",
292
+ ),
293
+ (
294
+ "warmup_decode",
295
+ {
296
+ "kv_cache": "cache",
297
+ "max_batch_size": 1,
298
+ "num_blocks": 8,
299
+ "can_sample_on_device": False,
300
+ "enable_trace": False,
301
+ },
302
+ "read_from_device",
303
+ ),
304
+ ],
305
+ )
306
+ def test_warmup_contract_rejects_unregistered_plugin_keywords(
307
+ method_name,
308
+ plugin_kwargs,
309
+ unexpected_name,
310
+ expect_error,
311
+ ):
312
+ coordinator, *_ = make_coordinator(trace_mode="none", sampling=False, lane_capacity=1)
313
+ plugin_kwargs[unexpected_name] = False
314
+
315
+ with expect_error(TypeError, unexpected_name):
316
+ getattr(coordinator, method_name)(**plugin_kwargs)
317
+
318
+
319
+ def test_configured_prefill_lengths_override_model_supported_defaults():
320
+ coordinator, execution, *_ = make_coordinator(
321
+ warmup_config=WarmupConfig(prefill_seq_lens=(1024,), prefill_batch_sizes=(1,)),
322
+ sequence_lengths=(128,),
323
+ sampling=False,
324
+ )
325
+
326
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
327
+
328
+ regular_lengths = [int(call["tokens"].shape[-1]) for call in execution.prefill_calls if call["start_pos"] is None]
329
+ assert regular_lengths == [1024]
330
+
331
+
332
+ @pytest.mark.parametrize(
333
+ ("sequence_lengths", "message"),
334
+ [
335
+ ((), "non-empty tuple"),
336
+ ((True,), "positive integers"),
337
+ ((128, 128), "unique"),
338
+ ],
339
+ )
340
+ def test_model_supported_prefill_lengths_are_validated_once(sequence_lengths, message, expect_error):
341
+ with expect_error(ValueError, message):
342
+ make_coordinator(sequence_lengths=sequence_lengths)
343
+
344
+
345
+ def test_sampler_argmax_capability_is_resolved_once():
346
+ class SamplingConfig:
347
+ reads = 0
348
+ max_top_k = 32
349
+
350
+ @property
351
+ def allow_force_argmax(self):
352
+ self.reads += 1
353
+ return True
354
+
355
+ sampling_config = SamplingConfig()
356
+ coordinator, *_ = make_coordinator(sampling_config=sampling_config)
357
+ resolved_reads = sampling_config.reads
358
+
359
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=True)
360
+ coordinator.warmup_decode(
361
+ kv_cache="cache",
362
+ enable_trace=False,
363
+ max_batch_size=4,
364
+ num_blocks=8,
365
+ can_sample_on_device=True,
366
+ )
367
+
368
+ assert sampling_config.reads == resolved_reads
369
+
370
+
371
+ def test_page_table_layout_can_be_reconfigured_only_before_use(expect_error):
372
+ coordinator, execution, *_ = make_coordinator(
373
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
374
+ sampling=False,
375
+ )
376
+ final_layout = PageTableLayout(
377
+ block_size=32,
378
+ raw_capacity_width=4,
379
+ prefill_width=64,
380
+ decode_width=8,
381
+ )
382
+
383
+ coordinator.configure_page_table_layout(final_layout)
384
+ assert coordinator.config.page_table_layout is final_layout
385
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
386
+
387
+ assert all(call["start_pos"] is None for call in execution.prefill_calls)
388
+ with expect_error(RuntimeError, "configuration is sealed"):
389
+ coordinator.configure_page_table_layout(final_layout)
390
+
391
+
392
+ def test_explicit_configuration_seal_precedes_physical_kv_allocation(expect_error):
393
+ coordinator, *_ = make_coordinator()
394
+ coordinator.seal_configuration()
395
+
396
+ with expect_error(RuntimeError, "configuration is sealed"):
397
+ coordinator.configure_page_table_layout(
398
+ PageTableLayout(
399
+ block_size=32,
400
+ raw_capacity_width=64,
401
+ prefill_width=128,
402
+ decode_width=64,
403
+ )
404
+ )
405
+
406
+
407
+ def test_page_table_layout_reconfiguration_requires_immutable_layout(expect_error):
408
+ coordinator, *_ = make_coordinator()
409
+
410
+ with expect_error(TypeError, "PageTableLayout"):
411
+ coordinator.configure_page_table_layout(SimpleNamespace(block_size=32))
412
+
413
+
414
+ def test_resolved_config_is_frozen_and_owns_both_coverage_plans(expect_error):
415
+ coordinator, *_ = make_coordinator(
416
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
417
+ lane_capacity=2,
418
+ )
419
+
420
+ assert coordinator.config.eager_plan.decode == (coordinator.config.eager_plan.decode[0],)
421
+ assert [case.sampling_path for case in coordinator.config.sampled_plan.decode] == ["logits", "argmax"]
422
+ with expect_error(AttributeError, "cannot assign"):
423
+ coordinator.config.lane_batch_size = 4
424
+
425
+
426
+ def test_direct_config_construction_rejects_inconsistent_derived_plan(expect_error):
427
+ coordinator, *_ = make_coordinator()
428
+
429
+ with expect_error(ValueError, "plans must match"):
430
+ replace(coordinator.config, eager_plan=coordinator.config.sampled_plan)
431
+
432
+
433
+ @pytest.mark.parametrize(
434
+ ("mismatch", "message"),
435
+ [
436
+ ("model", "share one model"),
437
+ ("layout", "share one page-table layout"),
438
+ ("lane", "share one lane capacity"),
439
+ ("sampling", "share device-sampling policy"),
440
+ ("argmax", "share force-argmax capability"),
441
+ ("raw_ceiling", "share one page-table layout ceiling"),
442
+ ("decode_ceiling", "share one page-table layout ceiling"),
443
+ ],
444
+ )
445
+ def test_resolution_rejects_inconsistent_runtime_configs(mismatch, message, expect_error):
446
+ prefill, decode = make_runtime_configs()
447
+ if mismatch == "model":
448
+ _, decode = make_runtime_configs()
449
+ elif mismatch == "layout":
450
+ decode = decode.with_page_table_layout(PageTableLayout(32, 64, 128, 64))
451
+ elif mismatch == "lane":
452
+ decode = DecodeRuntimeConfig.resolve(
453
+ model=prefill.model,
454
+ output_reader=prefill.output_reader,
455
+ lane_capacity=2,
456
+ page_table_layout=prefill.page_table_layout,
457
+ device_sampling_enabled=True,
458
+ )
459
+ elif mismatch == "sampling":
460
+ decode = DecodeRuntimeConfig.resolve(
461
+ model=prefill.model,
462
+ output_reader=prefill.output_reader,
463
+ lane_capacity=prefill.max_batch_size,
464
+ page_table_layout=prefill.page_table_layout,
465
+ device_sampling_enabled=False,
466
+ )
467
+ elif mismatch == "argmax":
468
+ prefill.model.sampling.config.allow_force_argmax = False
469
+ decode = DecodeRuntimeConfig.resolve(
470
+ model=prefill.model,
471
+ output_reader=prefill.output_reader,
472
+ lane_capacity=prefill.max_batch_size,
473
+ page_table_layout=prefill.page_table_layout,
474
+ device_sampling_enabled=True,
475
+ )
476
+ elif mismatch == "raw_ceiling":
477
+ ceiling = prefill.page_table_layout_ceiling
478
+ prefill = replace(
479
+ prefill,
480
+ page_table_layout_ceiling=replace(
481
+ ceiling,
482
+ raw_capacity_width=ceiling.raw_capacity_width + 1,
483
+ decode_width=ceiling.decode_width + 8,
484
+ ),
485
+ )
486
+ else:
487
+ ceiling = prefill.page_table_layout_ceiling
488
+ prefill = replace(
489
+ prefill,
490
+ page_table_layout_ceiling=replace(ceiling, decode_width=ceiling.decode_width + 8),
491
+ )
492
+
493
+ with expect_error(ValueError, message):
494
+ WarmupCoordinatorConfig.resolve(
495
+ warmup=WarmupConfig(),
496
+ trace=TraceConfig("all"),
497
+ prefill=prefill,
498
+ decode=decode,
499
+ prefill_sequence_lengths=(128,),
500
+ )
501
+
502
+
503
+ def test_constructor_rejects_execution_disagreement_with_resolved_config(expect_error):
504
+ coordinator, execution, *_ = make_coordinator()
505
+ config = coordinator.config
506
+ execution.prefill.config = replace(
507
+ execution.prefill.config,
508
+ page_table_layout=PageTableLayout(32, 64, 128, 64),
509
+ )
510
+
511
+ with expect_error(ValueError, "warmup config page-table layout"):
512
+ WarmupCoordinator(
513
+ config=config,
514
+ execution=execution,
515
+ ensure_sampling_buffers=lambda: None,
516
+ validate_bound_cache=lambda _: None,
517
+ )
518
+
519
+
520
+ def test_runtime_does_not_copy_static_config_fields():
521
+ coordinator, *_ = make_coordinator()
522
+
523
+ assert {
524
+ "page_table_layout",
525
+ "prefill_sequence_lengths",
526
+ "lane_batch_size",
527
+ "device_sampling_enabled",
528
+ "allow_force_argmax",
529
+ "prime_q128_tile_ends",
530
+ "prefill_trace_enabled",
531
+ "decode_trace_enabled",
532
+ "eager_plan",
533
+ "sampled_plan",
534
+ }.isdisjoint(vars(coordinator))
535
+
536
+
537
+ def test_layout_replacement_is_immutable_bounded_and_rebuilds_coverage(expect_error):
538
+ coordinator, *_ = make_coordinator(
539
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
540
+ page_table_layout=PageTableLayout(32, 128, 192, 128),
541
+ )
542
+ original = coordinator.config
543
+ replacement = PageTableLayout(32, 4, 64, 8)
544
+
545
+ coordinator.configure_page_table_layout(replacement)
546
+
547
+ assert coordinator.config is not original
548
+ assert coordinator.config.page_table_layout_ceiling is original.page_table_layout
549
+ assert original.page_table_layout.raw_capacity_width == 128
550
+ assert not any(case.cached_tokens for case in coordinator.config.eager_plan.prefill)
551
+ with expect_error(ValueError, "cannot change block_size"):
552
+ original.with_page_table_layout(PageTableLayout(16, 4, 64, 8))
553
+ with expect_error(ValueError, "capacity ceiling"):
554
+ original.with_page_table_layout(PageTableLayout(32, 129, 192, 136))
555
+ with expect_error(ValueError, "canonical geometry"):
556
+ original.with_page_table_layout(PageTableLayout(32, 128, 200, 128))
557
+
558
+
559
+ def test_q128_batches_are_capped_by_lane_and_non128_is_batch_one():
560
+ config = WarmupConfig(prefill_batch_sizes=(1, 2, 4, 8, 16, 32))
561
+ coordinator, execution, *_ = make_coordinator(warmup_config=config, lane_capacity=8)
562
+
563
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
564
+
565
+ regular_q128 = [
566
+ int(call["tokens"].shape[0])
567
+ for call in execution.prefill_calls
568
+ if int(call["tokens"].shape[-1]) == 128 and call["start_pos"] is None
569
+ ]
570
+ regular_q1024 = [
571
+ int(call["tokens"].shape[0])
572
+ for call in execution.prefill_calls
573
+ if int(call["tokens"].shape[-1]) == 1024 and call["start_pos"] is None
574
+ ]
575
+ assert regular_q128 == [1, 2, 4, 8]
576
+ assert regular_q1024 == [1]
577
+
578
+
579
+ def test_sampling_paths_include_forced_prefill_topk_and_opt_in_true_topk_decode():
580
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,), include_decode_top_k=True)
581
+ coordinator, execution, *_ = make_coordinator(warmup_config=config, sequence_lengths=(128,), lane_capacity=2)
582
+
583
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=True)
584
+ coordinator.warmup_decode(
585
+ kv_cache="cache",
586
+ enable_trace=False,
587
+ max_batch_size=2,
588
+ num_blocks=8,
589
+ can_sample_on_device=True,
590
+ )
591
+
592
+ assert execution.prefill_calls[0]["sampling_params"] is None
593
+ assert execution.prefill_calls[1]["sampling_params"].top_k.tolist() == [32]
594
+ assert execution.decode_calls[0]["sampling_params"] is None
595
+ assert execution.decode_calls[1]["sampling_params"].top_k.tolist() == [1, 1]
596
+ assert execution.decode_calls[2]["sampling_params"].top_k.tolist() == [32, 32]
597
+ # Preserve the established true-top-k recipe, not merely a top-k label.
598
+ assert execution.decode_calls[2]["sampling_params"].top_p.tolist() == pytest.approx([0.08, 0.08])
599
+
600
+
601
+ def test_q128_single_topk_primes_all_tile_ends():
602
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
603
+ coordinator, execution, *_ = make_coordinator(
604
+ warmup_config=config,
605
+ sequence_lengths=(128,),
606
+ lane_capacity=32,
607
+ )
608
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=True)
609
+
610
+ topk_calls = [
611
+ call
612
+ for call in execution.prefill_calls
613
+ if call["sampling_params"] is not None
614
+ and float(call["sampling_params"].temperature[0]) == 1.0
615
+ and call["start_pos"] is None
616
+ ]
617
+ assert [int(call["prompt_lens"][0]) for call in topk_calls] == [32, 64, 96, 128]
618
+ assert [int(call["tokens"].shape[-1]) for call in topk_calls] == [32, 64, 96, 128]
619
+ argmax_calls = [
620
+ call
621
+ for call in execution.prefill_calls
622
+ if call["sampling_params"] is not None
623
+ and float(call["sampling_params"].temperature[0]) == 0.0
624
+ and call["start_pos"] is None
625
+ ]
626
+ assert [int(call["prompt_lens"][0]) for call in argmax_calls] == [32, 64, 96, 128]
627
+
628
+
629
+ def test_decode_warmup_uses_topk_as_the_platform_greedy_path_when_argmax_is_disabled():
630
+ coordinator, execution, *_ = make_coordinator(
631
+ sequence_lengths=(128,),
632
+ lane_capacity=2,
633
+ allow_force_argmax=False,
634
+ )
635
+
636
+ coordinator.warmup_decode(
637
+ kv_cache="cache",
638
+ enable_trace=False,
639
+ max_batch_size=2,
640
+ num_blocks=8,
641
+ can_sample_on_device=True,
642
+ )
643
+
644
+ assert execution.decode_calls[0]["sampling_params"] is None
645
+ assert execution.decode_calls[1]["sampling_params"].top_k.tolist() == [32, 32]
646
+
647
+
648
+ def test_eager_and_trace_coverage_are_separately_idempotent():
649
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
650
+ coordinator, execution, trace_compiler, *_ = make_coordinator(
651
+ warmup_config=config, sequence_lengths=(128,), lane_capacity=1
652
+ )
653
+
654
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
655
+ eager_calls = len(execution.prefill_calls)
656
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
657
+ assert len(execution.prefill_calls) == eager_calls
658
+
659
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=True)
660
+ trace_calls = len(execution.prefill_calls)
661
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=True)
662
+ assert len(execution.prefill_calls) == trace_calls
663
+ assert trace_compiler.calls == 0
664
+
665
+
666
+ def test_trace_warmup_routes_cached_prefill_through_traced_execution_target():
667
+ coordinator, eager, *_ = make_coordinator(
668
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
669
+ sequence_lengths=(128,),
670
+ lane_capacity=1,
671
+ sampling=False,
672
+ )
673
+ traced = RecordingExecution()
674
+ coordinator.execution = traced
675
+
676
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=False)
677
+
678
+ assert not eager.prefill_calls
679
+ assert any(call["start_pos"] is None for call in traced.prefill_calls)
680
+ assert any(call["start_pos"] is not None for call in traced.prefill_calls)
681
+
682
+
683
+ def test_coverage_manifest_uses_compiler_registries_and_deduplicates_trace_identities():
684
+ eager_program = CompiledProgram(ProgramKey("0" * 64), "eager", OutputSpec((1,), torch.float32))
685
+ first_traced = CompiledProgram(ProgramKey("1" * 64), "traced-a", OutputSpec((1,), torch.float32))
686
+ second_traced = CompiledProgram(ProgramKey("2" * 64), "traced-b", OutputSpec((1,), torch.float32))
687
+ shared_trace_key = ProgramKey("a" * 64)
688
+ program_compiler = SimpleNamespace(compiled_programs=(eager_program, first_traced, second_traced))
689
+ eager = SimpleNamespace(program_compiler=program_compiler)
690
+ trace_compiler = SimpleNamespace(
691
+ trace_key_for_program=lambda key: None if key == eager_program.key else shared_trace_key,
692
+ get=lambda key: SimpleNamespace(signature="shared-trace") if key == shared_trace_key else None,
693
+ )
694
+
695
+ manifest = _resolve_coverage_manifest(eager, trace_compiler)
696
+
697
+ assert manifest.eager_program_signatures == ("eager",)
698
+ assert manifest.traced_source_program_signatures == ("traced-a", "traced-b")
699
+ assert manifest.trace_signatures == ("shared-trace",)
700
+ assert manifest.aliases == (
701
+ CoverageAlias("traced-a", "shared-trace"),
702
+ CoverageAlias("traced-b", "shared-trace"),
703
+ )
704
+
705
+
706
+ def test_coverage_manifest_rejects_any_missing_required_trace_alias(expect_error):
707
+ first_traced = CompiledProgram(ProgramKey("1" * 64), "traced-a", OutputSpec((1,), torch.float32))
708
+ second_traced = CompiledProgram(ProgramKey("2" * 64), "traced-b", OutputSpec((1,), torch.float32))
709
+ trace_key = ProgramKey("a" * 64)
710
+ eager = SimpleNamespace(program_compiler=SimpleNamespace(compiled_programs=(first_traced, second_traced)))
711
+ trace_compiler = SimpleNamespace(
712
+ trace_key_for_program=lambda key: trace_key if key == first_traced.key else None,
713
+ get=lambda key: SimpleNamespace(signature="trace-a") if key == trace_key else None,
714
+ )
715
+
716
+ with expect_error(RuntimeError, "required trace alias"):
717
+ _resolve_coverage_manifest(
718
+ eager,
719
+ trace_compiler,
720
+ required_program_keys={first_traced.key, second_traced.key},
721
+ required_trace_program_keys={first_traced.key, second_traced.key},
722
+ )
723
+
724
+
725
+ def test_activation_validates_every_program_returned_by_trace_warmup(expect_error):
726
+ coordinator, execution, trace_compiler, *_ = make_coordinator(
727
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
728
+ sequence_lengths=(128,),
729
+ lane_capacity=1,
730
+ sampling=False,
731
+ )
732
+ programs = tuple(
733
+ CompiledProgram(ProgramKey(str(index) * 64), f"program-{index}", OutputSpec((1,), torch.float32))
734
+ for index in range(1, 4)
735
+ )
736
+ execution.program_compiler = SimpleNamespace(compiled_programs=programs)
737
+ prefill_programs = iter(programs[:2])
738
+ execution.compile_prefill = lambda **_kwargs: (next(prefill_programs),)
739
+ execution.compile_decode = lambda **_kwargs: programs[2]
740
+ trace_keys = {
741
+ programs[0].key: ProgramKey("a" * 64),
742
+ programs[2].key: ProgramKey("c" * 64),
743
+ }
744
+ trace_compiler.trace_key_for_program = trace_keys.get
745
+ trace_compiler.get = lambda key: SimpleNamespace(signature=f"trace-{key.digest[0]}")
746
+
747
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=False)
748
+ with expect_error(RuntimeError, programs[1].key.digest):
749
+ coordinator.warmup_decode(
750
+ kv_cache="cache",
751
+ enable_trace=True,
752
+ max_batch_size=1,
753
+ num_blocks=8,
754
+ can_sample_on_device=False,
755
+ )
756
+
757
+ assert trace_compiler.calls == 0
758
+
759
+
760
+ @pytest.mark.parametrize("order", [("prefill", "decode"), ("decode", "prefill")])
761
+ def test_prefill_decode_order_is_independent_and_capture_waits_for_both(order):
762
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
763
+ coordinator, execution, trace_compiler, *_ = make_coordinator(
764
+ warmup_config=config, sequence_lengths=(128,), lane_capacity=1
765
+ )
766
+
767
+ def run(operation):
768
+ if operation == "prefill":
769
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=True)
770
+ else:
771
+ coordinator.warmup_decode(
772
+ kv_cache="cache",
773
+ enable_trace=True,
774
+ max_batch_size=1,
775
+ num_blocks=8,
776
+ can_sample_on_device=True,
777
+ )
778
+
779
+ run(order[0])
780
+ assert trace_compiler.calls == 0
781
+ run(order[1])
782
+ assert trace_compiler.calls == 1
783
+ run(order[0])
784
+ run(order[1])
785
+ assert trace_compiler.calls == 1
786
+
787
+
788
+ @pytest.mark.parametrize("order", [("prefill", "decode"), ("decode", "prefill")])
789
+ def test_capture_uses_phase_specific_sampling_decisions(order):
790
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
791
+ coordinator, execution, trace_compiler, *_ = make_coordinator(
792
+ trace_mode="all",
793
+ warmup_config=config,
794
+ sequence_lengths=(128,),
795
+ lane_capacity=1,
796
+ )
797
+
798
+ def run(operation):
799
+ if operation == "prefill":
800
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=False)
801
+ else:
802
+ coordinator.warmup_decode(
803
+ kv_cache="cache",
804
+ enable_trace=True,
805
+ max_batch_size=1,
806
+ num_blocks=8,
807
+ can_sample_on_device=True,
808
+ )
809
+
810
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
811
+ coordinator.warmup_decode(
812
+ kv_cache="cache",
813
+ enable_trace=False,
814
+ max_batch_size=1,
815
+ num_blocks=8,
816
+ can_sample_on_device=True,
817
+ )
818
+ with coordinator.defer_capture():
819
+ run(order[0])
820
+ assert trace_compiler.calls == 0
821
+ run(order[1])
822
+ assert coordinator.capture_pending
823
+ coordinator.activate_pending_capture()
824
+
825
+ assert trace_compiler.calls == 1
826
+ assert coordinator.already_warmed_up_prefill
827
+ assert all(call["sampling_params"] is None for call in execution.prefill_calls)
828
+ assert any(call["sampling_params"] is not None for call in execution.decode_calls)
829
+ assert not execution.prefill_replays
830
+
831
+
832
+ def test_capture_deferral_stages_complete_registration_until_explicit_activation():
833
+ coordinator, _, trace_compiler, *_ = make_coordinator(
834
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
835
+ sequence_lengths=(128,),
836
+ lane_capacity=1,
837
+ sampling=False,
838
+ )
839
+
840
+ with coordinator.defer_capture():
841
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=False)
842
+ coordinator.warmup_decode(
843
+ kv_cache="cache",
844
+ enable_trace=True,
845
+ max_batch_size=1,
846
+ num_blocks=8,
847
+ can_sample_on_device=False,
848
+ )
849
+ assert coordinator.capture_pending
850
+ assert not coordinator.trace_activated
851
+ assert trace_compiler.calls == 0
852
+ coordinator.activate_pending_capture()
853
+ assert coordinator.trace_activated
854
+ assert trace_compiler.calls == 1
855
+
856
+ assert not coordinator.capture_pending
857
+
858
+
859
+ def test_capture_deferral_exception_discards_pending_activation(expect_error):
860
+ coordinator, _, trace_compiler, *_ = make_coordinator(
861
+ warmup_config=WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,)),
862
+ sequence_lengths=(128,),
863
+ lane_capacity=1,
864
+ sampling=False,
865
+ )
866
+
867
+ with expect_error(RuntimeError, "staging failed"):
868
+ with coordinator.defer_capture():
869
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=False)
870
+ coordinator.warmup_decode(
871
+ kv_cache="cache",
872
+ enable_trace=True,
873
+ max_batch_size=1,
874
+ num_blocks=8,
875
+ can_sample_on_device=False,
876
+ )
877
+ assert coordinator.capture_pending
878
+ raise RuntimeError("staging failed")
879
+
880
+ assert not coordinator.capture_pending
881
+ assert not coordinator.trace_activated
882
+ assert trace_compiler.calls == 0
883
+
884
+
885
+ def test_static_all_can_capture_decode_only_runtime_trace():
886
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
887
+ coordinator, execution, trace_compiler, *_ = make_coordinator(
888
+ trace_mode="all",
889
+ warmup_config=config,
890
+ sequence_lengths=(128,),
891
+ lane_capacity=1,
892
+ )
893
+
894
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=True)
895
+ assert coordinator.already_warmed_up_prefill
896
+ assert trace_compiler.calls == 0
897
+
898
+ coordinator.warmup_decode(
899
+ kv_cache="cache",
900
+ enable_trace=True,
901
+ max_batch_size=1,
902
+ num_blocks=8,
903
+ can_sample_on_device=True,
904
+ )
905
+
906
+ assert trace_compiler.calls == 1
907
+ assert not execution.prefill_replays
908
+
909
+
910
+ def test_two_phase_static_all_waits_for_phase_two_decode_before_capture():
911
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
912
+ coordinator, _, trace_compiler, *_ = make_coordinator(
913
+ trace_mode="all",
914
+ warmup_config=config,
915
+ sequence_lengths=(128,),
916
+ lane_capacity=1,
917
+ )
918
+
919
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=True)
920
+ coordinator.warmup_decode(
921
+ kv_cache="cache",
922
+ enable_trace=False,
923
+ max_batch_size=1,
924
+ num_blocks=8,
925
+ can_sample_on_device=True,
926
+ )
927
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=True)
928
+ assert trace_compiler.calls == 0
929
+
930
+ coordinator.warmup_decode(
931
+ kv_cache="cache",
932
+ enable_trace=True,
933
+ max_batch_size=1,
934
+ num_blocks=8,
935
+ can_sample_on_device=True,
936
+ )
937
+
938
+ assert trace_compiler.calls == 1
939
+
940
+
941
+ def test_sampling_buffers_are_materialized_before_first_compile_and_capture():
942
+ events = []
943
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
944
+ coordinator, _, _, _, _, events = make_coordinator(
945
+ warmup_config=config,
946
+ sequence_lengths=(128,),
947
+ lane_capacity=1,
948
+ events=events,
949
+ )
950
+
951
+ coordinator.warmup_decode(
952
+ kv_cache="cache",
953
+ enable_trace=True,
954
+ max_batch_size=1,
955
+ num_blocks=8,
956
+ can_sample_on_device=True,
957
+ )
958
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=True)
959
+
960
+ assert events.index("sampling") < events.index("compile_decode")
961
+ assert events.index("sampling") < events.index("compile_prefill")
962
+ assert max(index for index, event in enumerate(events) if event.startswith("compile_")) < events.index("capture")
963
+
964
+
965
+ def test_failed_case_is_not_marked_complete_and_retry_skips_completed_case(expect_error):
966
+ execution = RecordingExecution()
967
+ execution.fail_decode_call = 2
968
+ config = WarmupConfig(prefill_seq_lens=(128,), prefill_batch_sizes=(1,))
969
+ coordinator, execution, *_ = make_coordinator(
970
+ trace_mode="none",
971
+ sampling=True,
972
+ warmup_config=config,
973
+ sequence_lengths=(128,),
974
+ lane_capacity=1,
975
+ execution=execution,
976
+ )
977
+
978
+ with expect_error(RuntimeError, "decode compile failed"):
979
+ coordinator.warmup_decode(
980
+ kv_cache="cache",
981
+ enable_trace=False,
982
+ max_batch_size=1,
983
+ num_blocks=8,
984
+ can_sample_on_device=True,
985
+ )
986
+ assert len(execution.decode_calls) == 1
987
+
988
+ coordinator.warmup_decode(
989
+ kv_cache="cache",
990
+ enable_trace=False,
991
+ max_batch_size=1,
992
+ num_blocks=8,
993
+ can_sample_on_device=True,
994
+ )
995
+ assert len(execution.decode_calls) == 2
996
+
997
+
998
+ def test_dynamic_hints_cannot_expand_static_trace_or_sampling_ceilings(expect_error):
999
+ coordinator, *_ = make_coordinator(trace_mode="decode_only", sampling=False)
1000
+
1001
+ with expect_error(ValueError, "prefill trace warmup exceeds"):
1002
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=True, can_sample_on_device=False)
1003
+ with expect_error(ValueError, "statically disabled"):
1004
+ coordinator.warmup_decode(
1005
+ kv_cache="cache",
1006
+ enable_trace=False,
1007
+ max_batch_size=4,
1008
+ num_blocks=8,
1009
+ can_sample_on_device=True,
1010
+ )
1011
+
1012
+ coordinator.warmup_prefill(kv_cache="cache", enable_trace=False, can_sample_on_device=False)
code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_demo_contract.py ADDED
@@ -0,0 +1,452 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import ast
5
+ import os
6
+ from pathlib import Path
7
+ from types import SimpleNamespace
8
+
9
+ import pytest
10
+ import torch
11
+
12
+ from models.common.llm_runtime.config import TraceConfig
13
+ from models.common.llm_runtime.prefill.plan import _plan_prefill_requests
14
+
15
+ _DEMO_PATH = "models/common/tests/demos/deepseek_r1_distill_qwen_14b/demo.py"
16
+ _DEMO_TREE = ast.parse(Path(_DEMO_PATH).read_text(encoding="utf-8"), filename=_DEMO_PATH)
17
+
18
+
19
+ def _demo_function(name, namespace=None):
20
+ function = next(node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == name)
21
+ namespace = {} if namespace is None else namespace
22
+ exec(compile(ast.Module(body=[function], type_ignores=[]), _DEMO_PATH, "exec"), namespace)
23
+ return namespace[name]
24
+
25
+
26
+ def _called_names(function_name):
27
+ function = next(
28
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == function_name
29
+ )
30
+ return [
31
+ node.func.id for node in ast.walk(function) if isinstance(node, ast.Call) and isinstance(node.func, ast.Name)
32
+ ]
33
+
34
+
35
+ def test_demo_case_manifest_is_preserved():
36
+ test_function = next(
37
+ node
38
+ for node in _DEMO_TREE.body
39
+ if isinstance(node, ast.FunctionDef) and node.name == "test_deepseek_r1_qwen_14b"
40
+ )
41
+ decorators = [node for node in test_function.decorator_list if isinstance(node, ast.Call)]
42
+ test_config = next(node for node in decorators if ast.literal_eval(node.args[0]) == "test_config")
43
+ optimizations = next(node for node in decorators if ast.literal_eval(node.args[0]) == "optimizations")
44
+ case_ids = [ast.literal_eval(element.keywords[0].value) for element in test_config.args[1].elts]
45
+ assert case_ids == [
46
+ "token-accuracy",
47
+ "batch-1",
48
+ "batch-32",
49
+ "batch-32-ci",
50
+ "eval-32",
51
+ "ci-b1-DP-2",
52
+ "ci-b1-DP-4",
53
+ "ci-b1-DP-8",
54
+ "ci-b1-DP-16",
55
+ "ci-b1-DP-32",
56
+ ]
57
+ assert ast.literal_eval(optimizations.args[1]) == ["performance", "accuracy"]
58
+
59
+
60
+ def test_demo_reserves_trace_space_by_mesh(monkeypatch):
61
+ for mesh_name, mesh_shape, trace_region_size in (
62
+ ("N300", (1, 2), 50_000_000),
63
+ ("T3K", (1, 8), 100_000_000),
64
+ ):
65
+ monkeypatch.setenv("MESH_DEVICE", mesh_name)
66
+ device_params = _demo_function(
67
+ "_ttnn_mesh_device_param_from_env",
68
+ {
69
+ "os": os,
70
+ "pytest": pytest,
71
+ "_MESH_DEVICE_TO_SHAPE": {mesh_name: mesh_shape},
72
+ "ttnn": SimpleNamespace(FabricConfig=SimpleNamespace(FABRIC_1D=object())),
73
+ },
74
+ )()
75
+
76
+ assert device_params["mesh_shape"] == mesh_shape
77
+ assert device_params["trace_region_size"] == trace_region_size
78
+
79
+
80
+ def test_demo_warmup_compiles_eager_programs_before_trace_capture():
81
+ calls = []
82
+ config = SimpleNamespace(trace=TraceConfig("all"), device_sampling_enabled=True)
83
+ executor = SimpleNamespace(
84
+ config=config,
85
+ model=SimpleNamespace(config=SimpleNamespace(max_batch_size=4)),
86
+ warmup_model_prefill=lambda **kwargs: calls.append(("prefill", kwargs)),
87
+ warmup_model_decode=lambda **kwargs: calls.append(("decode", kwargs)),
88
+ )
89
+ warmup = _demo_function("_warmup_demo_executor")
90
+ kv_cache = object()
91
+ warmup(executor, kv_cache=kv_cache, page_table=SimpleNamespace(shape=(4, 8)))
92
+
93
+ assert [(kind, kwargs["enable_trace"]) for kind, kwargs in calls] == [
94
+ ("decode", False),
95
+ ("prefill", False),
96
+ ("prefill", True),
97
+ ("decode", True),
98
+ ]
99
+ assert all(kwargs["kv_cache"] is kv_cache for _, kwargs in calls)
100
+
101
+
102
+ def test_demo_warmup_registers_concrete_prefill_before_trace_capture():
103
+ calls = []
104
+ config = SimpleNamespace(trace=TraceConfig("all"), device_sampling_enabled=False)
105
+ eager_execution = object()
106
+ executor = SimpleNamespace(
107
+ config=config,
108
+ eager_execution=eager_execution,
109
+ model=SimpleNamespace(config=SimpleNamespace(max_batch_size=32)),
110
+ warmup_model_prefill=lambda **kwargs: calls.append(("prefill", kwargs)),
111
+ warmup_model_decode=lambda **kwargs: calls.append(("decode", kwargs)),
112
+ compile_prefill=lambda **kwargs: calls.append(("compile_prefill", kwargs)),
113
+ )
114
+ warmup = _demo_function("_warmup_demo_executor")
115
+ tokens = torch.zeros((32, 700), dtype=torch.long)
116
+ prompt_lens = torch.tensor([64] * 30 + [400, 700])
117
+ page_table = torch.zeros((32, 64), dtype=torch.int32)
118
+ kv_cache = object()
119
+
120
+ warmup(
121
+ executor,
122
+ kv_cache=kv_cache,
123
+ page_table=page_table,
124
+ prefill_compile_case=(tokens, prompt_lens),
125
+ )
126
+
127
+ assert [(kind, kwargs.get("enable_trace")) for kind, kwargs in calls] == [
128
+ ("decode", False),
129
+ ("prefill", False),
130
+ ("compile_prefill", None),
131
+ ("prefill", True),
132
+ ("decode", True),
133
+ ]
134
+ compile_kwargs = calls[2][1]
135
+ assert compile_kwargs["tokens"] is tokens
136
+ assert compile_kwargs["prompt_lens"] is prompt_lens
137
+ assert compile_kwargs["page_table"] is page_table
138
+ assert compile_kwargs["kv_cache"] is kv_cache
139
+ assert compile_kwargs["empty_slots"] == list(range(32))
140
+ assert compile_kwargs["execution"] is eager_execution
141
+
142
+
143
+ def test_demo_warmup_uses_lane_group_capacity_and_lane_trace_policy():
144
+ calls = []
145
+ lane_config = SimpleNamespace(trace=TraceConfig("all"), device_sampling_enabled=True)
146
+ group = SimpleNamespace(
147
+ lanes=[SimpleNamespace(config=lane_config) for _ in range(4)],
148
+ max_batch_size=4,
149
+ warmup_model_prefill=lambda **kwargs: calls.append(("prefill", kwargs)),
150
+ warmup_model_decode=lambda **kwargs: calls.append(("decode", kwargs)),
151
+ )
152
+ warmup = _demo_function("_warmup_demo_executor")
153
+ kv_cache = [object() for _ in range(4)]
154
+ warmup(group, kv_cache=kv_cache, page_table=SimpleNamespace(shape=(4, 128)))
155
+
156
+ decode_calls = [kwargs for kind, kwargs in calls if kind == "decode"]
157
+ assert len(decode_calls) == 2
158
+ assert all(kwargs["max_batch_size"] == 4 for kwargs in decode_calls)
159
+ assert all(kwargs["num_blocks"] == 128 for kwargs in decode_calls)
160
+ assert all(kwargs["kv_cache"] is kv_cache for _, kwargs in calls)
161
+
162
+
163
+ def test_eval_prefill_signature_multiset_is_rotation_invariant_and_not_static_warmup_shaped():
164
+ tokens = torch.zeros((32, 700), dtype=torch.long)
165
+ prompt_lens = torch.tensor([64] * 30 + [400, 700])
166
+ page_table = torch.zeros((32, 64), dtype=torch.int32)
167
+
168
+ def planned_shapes(offset):
169
+ rotated_tokens = torch.roll(tokens, shifts=-offset, dims=0)
170
+ rotated_lens = torch.roll(prompt_lens, shifts=-offset, dims=0)
171
+ requests = _plan_prefill_requests(
172
+ tokens=rotated_tokens,
173
+ page_table=page_table,
174
+ prompt_lens=rotated_lens,
175
+ empty_slots=list(range(32)),
176
+ start_pos=None,
177
+ block_size=32,
178
+ max_batch_size=32,
179
+ max_prefill_chunk_size=1024,
180
+ supports_batched_prefill=True,
181
+ max_prefill_batch_size=8,
182
+ max_actual_page_table_width=32,
183
+ canonical_page_table_width=64,
184
+ )
185
+ return sorted(
186
+ (request.padded_sequence_length, request.padded_batch_size, len(request.source_rows))
187
+ for request in requests
188
+ )
189
+
190
+ expected = [(128, 8, 6), (128, 8, 8), (128, 8, 8), (128, 8, 8), (1024, 2, 2)]
191
+ assert planned_shapes(0) == expected
192
+ assert planned_shapes(1) == expected
193
+ assert planned_shapes(2) == expected
194
+
195
+
196
+ @pytest.mark.parametrize("function_name", ["_run_perf_benchmark", "_run_eval_repeat_batch32"])
197
+ def test_traced_demo_paths_warm_up_fresh_executor(function_name):
198
+ assert "_warmup_demo_executor" in _called_names(function_name)
199
+
200
+
201
+ def test_create_executor_uses_model_owned_executor_and_resolved_cache():
202
+ captured = {}
203
+
204
+ def executor_config(**kwargs):
205
+ captured.update(kwargs)
206
+ return SimpleNamespace(**kwargs)
207
+
208
+ namespace = {
209
+ "DeepSeekR1Qwen14B": object,
210
+ "DeepSeekR1Qwen14BExecutor": lambda model, runtime_config, config: config,
211
+ "DeepSeekR1Qwen14BExecutorConfig": executor_config,
212
+ "PagedKVCacheConfig": lambda **kwargs: SimpleNamespace(**kwargs),
213
+ "TraceConfig": TraceConfig,
214
+ "WarmupConfig": lambda: object(),
215
+ }
216
+ create_executor = _demo_function("create_executor", namespace)
217
+ model = SimpleNamespace(
218
+ model_args=object(),
219
+ config=SimpleNamespace(
220
+ max_seq_len=2048,
221
+ max_batch_size=32,
222
+ block_configs=[SimpleNamespace(attention_config=SimpleNamespace(kv_cache_dtype=object()))],
223
+ ),
224
+ )
225
+
226
+ result = create_executor(model, traced=True, device_sampling_enabled=True)
227
+
228
+ assert result.trace.mode == "all"
229
+ assert result.device_sampling_enabled is True
230
+ assert captured["paged_kv_cache"].num_blocks == 2048
231
+
232
+
233
+ def test_eval_uses_decode_only_trace_while_ordinary_traced_executor_uses_all():
234
+ create_executor = next(
235
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "create_executor"
236
+ )
237
+ trace_config = next(
238
+ node
239
+ for node in ast.walk(create_executor)
240
+ if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "TraceConfig"
241
+ )
242
+ assert isinstance(trace_config.keywords[0].value, ast.Name)
243
+ assert trace_config.keywords[0].value.id == "trace_mode"
244
+ derived_mode = next(
245
+ node
246
+ for node in ast.walk(create_executor)
247
+ if isinstance(node, ast.Assign)
248
+ and any(isinstance(target, ast.Name) and target.id == "trace_mode" for target in node.targets)
249
+ )
250
+ assert ast.unparse(derived_mode.value) == "'all' if traced else 'none'"
251
+
252
+ eval_function = next(
253
+ node
254
+ for node in _DEMO_TREE.body
255
+ if isinstance(node, ast.FunctionDef) and node.name == "_run_eval_repeat_batch32"
256
+ )
257
+ eval_create = next(
258
+ node
259
+ for node in ast.walk(eval_function)
260
+ if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "create_executor"
261
+ )
262
+ keywords = {keyword.arg: keyword.value for keyword in eval_create.keywords}
263
+ assert ast.literal_eval(keywords["traced"]) is True
264
+ assert ast.literal_eval(keywords["trace_mode"]) == "decode_only"
265
+
266
+
267
+ def test_deepseek_stop_guard_truncates_eos_but_not_ordinary_reasoning_tokens(expect_error, monkeypatch):
268
+ shared_calls = []
269
+
270
+ def shared_guard(generated_token_ids, tokenizer, **kwargs):
271
+ shared_calls.append((generated_token_ids, kwargs))
272
+ if kwargs["is_ci_env"] is None and os.environ.get("TT_DEMO_STRICT_SPECIAL_TOKENS") == "1":
273
+ outputs_before_eos = [
274
+ output[: output.index(tokenizer.eos_token_id)] if tokenizer.eos_token_id in output else output
275
+ for output in generated_token_ids
276
+ ]
277
+ if any(99 in output for output in outputs_before_eos):
278
+ raise AssertionError("model produced special tokens")
279
+
280
+ guard = _demo_function("assert_no_special_tokens", {"assert_no_special_tokens_shared": shared_guard})
281
+ tokenizer = SimpleNamespace(
282
+ all_special_ids=[10, 99],
283
+ eos_token_id=10,
284
+ )
285
+
286
+ monkeypatch.setenv("TT_DEMO_STRICT_SPECIAL_TOKENS", "1")
287
+ guard([[1, 10, 99], [2, 3, 4]], tokenizer)
288
+ assert shared_calls[-1][0] == [[1], [2, 3, 4]]
289
+ with expect_error(AssertionError, "model produced special tokens"):
290
+ guard([[1, 99]], tokenizer)
291
+
292
+
293
+ def test_dp_smoke_uses_model_owned_lane_group_execution():
294
+ calls = _called_names("_run_dp_smoke")
295
+ assert "_dp_lane_tp_or_skip" in calls
296
+ assert "_create_dp_submeshes" in calls
297
+ assert "create_executor" in calls
298
+ assert "LaneGroupExecutor" in calls
299
+ assert "run_perf_benchmark" in calls
300
+ assert "cleanup_dp_model_case" in calls
301
+ assert "_skip_below_min_tp_devices" not in calls
302
+
303
+
304
+ def test_runnable_dp_lane_build_errors_are_not_converted_to_topology_skips():
305
+ function = next(
306
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_dp_smoke"
307
+ )
308
+ pytest_skip_calls = [
309
+ node
310
+ for node in ast.walk(function)
311
+ if isinstance(node, ast.Call)
312
+ and isinstance(node.func, ast.Attribute)
313
+ and isinstance(node.func.value, ast.Name)
314
+ and node.func.value.id == "pytest"
315
+ and node.func.attr == "skip"
316
+ ]
317
+ assert pytest_skip_calls == []
318
+
319
+
320
+ def test_deepseek_dp_topology_accepts_t3k_dp2_tp4_and_dp4_tp2(expect_error):
321
+ topology = _demo_function(
322
+ "_dp_lane_tp_or_skip",
323
+ {"ttnn": SimpleNamespace(MeshDevice=object), "pytest": pytest, "_MIN_TP_DEVICES": 2},
324
+ )
325
+ t3k = SimpleNamespace(get_num_devices=lambda: 8)
326
+
327
+ assert topology(t3k, 2) == 4
328
+ assert topology(t3k, 4) == 2
329
+ with expect_error(pytest.skip.Exception, "DP-8 on 8 devices creates TP1 lanes"):
330
+ topology(t3k, 8)
331
+ with expect_error(pytest.skip.Exception, "DP-16 cannot partition 8 devices"):
332
+ topology(t3k, 16)
333
+
334
+
335
+ def test_deepseek_dp4_partitions_four_tp2_submeshes():
336
+ calls = []
337
+ submeshes = [object() for _ in range(4)]
338
+ parent = SimpleNamespace(
339
+ create_submeshes=lambda shape: calls.append(shape) or submeshes,
340
+ )
341
+ fake_ttnn = SimpleNamespace(MeshDevice=object, MeshShape=lambda rows, columns: (rows, columns))
342
+ create_submeshes = _demo_function("_create_dp_submeshes", {"ttnn": fake_ttnn})
343
+
344
+ assert create_submeshes(parent, 4, 2) == submeshes
345
+ assert calls == [(1, 2)]
346
+
347
+
348
+ def test_deepseek_dp_lane_cache_reuses_lane_topology(tmp_path):
349
+ cache_dir = tmp_path / "DeepSeek-R1-Distill-Qwen-14B" / "T3K"
350
+ cache_dir.mkdir(parents=True)
351
+ lane_cache_dir = _demo_function("_dp_lane_cache_dir", {"Path": Path})(cache_dir, 2)
352
+
353
+ assert lane_cache_dir == cache_dir.parent / "N300"
354
+ assert lane_cache_dir.is_dir()
355
+ assert _demo_function("_dp_lane_cache_dir", {"Path": Path})(cache_dir, 4) == cache_dir.parent / "N150x4"
356
+
357
+
358
+ def test_deepseek_dp_lane_contract_checks_heads_capacity_and_cache(expect_error):
359
+ validate = _demo_function(
360
+ "_validate_dp_lane",
361
+ {
362
+ "DeepSeekR1Qwen14B": object,
363
+ "DeepSeekR1Qwen14BExecutor": object,
364
+ "math": __import__("math"),
365
+ },
366
+ )
367
+ attention = SimpleNamespace(n_heads=40, n_kv_heads=8)
368
+ model = SimpleNamespace(
369
+ config=SimpleNamespace(
370
+ num_devices=2,
371
+ max_batch_size=1,
372
+ block_configs=[SimpleNamespace(attention_config=attention)],
373
+ )
374
+ )
375
+ cache = SimpleNamespace(max_num_blocks=128, num_blocks=128)
376
+ lane = SimpleNamespace(config=SimpleNamespace(paged_kv_cache=cache))
377
+
378
+ validate(model, lane, 2, 4096)
379
+ model.config.num_devices = 4
380
+ with expect_error(ValueError, "expected TP2, model uses TP4"):
381
+ validate(model, lane, 2, 4096)
382
+ model.config.num_devices = 2
383
+ model.config.max_batch_size = 2
384
+ with expect_error(ValueError, "capacity 1"):
385
+ validate(model, lane, 2, 4096)
386
+ model.config.max_batch_size = 1
387
+ cache.num_blocks = None
388
+ with expect_error(ValueError, "cache must contain 128 blocks"):
389
+ validate(model, lane, 2, 4096)
390
+
391
+
392
+ def test_token_accuracy_cleans_up_executor_in_finally():
393
+ function = next(
394
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_token_accuracy"
395
+ )
396
+ cleanup_calls = [
397
+ statement
398
+ for node in ast.walk(function)
399
+ if isinstance(node, ast.Try)
400
+ for statement in node.finalbody
401
+ if isinstance(statement, ast.Expr)
402
+ and isinstance(statement.value, ast.Call)
403
+ and isinstance(statement.value.func, ast.Attribute)
404
+ and statement.value.func.attr == "cleanup"
405
+ ]
406
+ assert len(cleanup_calls) == 1
407
+
408
+
409
+ def test_main_demo_does_not_synchronize_parent_mesh_after_prebuild_skip():
410
+ function = next(
411
+ node
412
+ for node in _DEMO_TREE.body
413
+ if isinstance(node, ast.FunctionDef) and node.name == "test_deepseek_r1_qwen_14b"
414
+ )
415
+ try_node = next(node for node in function.body if isinstance(node, ast.Try))
416
+
417
+ assert len(try_node.finalbody) == 1
418
+ guard = try_node.finalbody[0]
419
+ assert isinstance(guard, ast.If)
420
+ assert ast.unparse(guard.test) == "model is not None"
421
+ assert any(
422
+ isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "cleanup_model_case"
423
+ for node in ast.walk(guard)
424
+ )
425
+
426
+
427
+ @pytest.mark.parametrize("function_name", ["_run_token_accuracy", "_run_perf_benchmark", "_run_eval_repeat_batch32"])
428
+ def test_demo_reads_model_geometry_from_model_config(function_name):
429
+ function = next(
430
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == function_name
431
+ )
432
+ model_args_aliases = [
433
+ node
434
+ for node in ast.walk(function)
435
+ if isinstance(node, ast.Assign)
436
+ and isinstance(node.value, ast.Attribute)
437
+ and isinstance(node.value.value, ast.Name)
438
+ and node.value.value.id == "model"
439
+ and node.value.attr == "model_args"
440
+ ]
441
+ config_fields = {
442
+ node.attr
443
+ for node in ast.walk(function)
444
+ if isinstance(node, ast.Attribute)
445
+ and isinstance(node.value, ast.Attribute)
446
+ and isinstance(node.value.value, ast.Name)
447
+ and node.value.value.id == "model"
448
+ and node.value.attr == "config"
449
+ }
450
+
451
+ assert model_args_aliases == []
452
+ assert {"max_batch_size", "max_seq_len"} <= config_fields
code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_hf_adaptor.py ADDED
@@ -0,0 +1,290 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import inspect
5
+ from types import SimpleNamespace
6
+
7
+ import torch
8
+ from transformers import Qwen2Config, Qwen2ForCausalLM
9
+ from transformers.models.qwen2.modeling_qwen2 import Qwen2RotaryEmbedding
10
+
11
+ from models.common.models.deepseek_r1_distill_qwen_14b import generator, hf_adaptor
12
+ from models.common.models.deepseek_r1_distill_qwen_14b import model as qwen_model
13
+ from models.common.models.deepseek_r1_distill_qwen_14b import weight_utils
14
+ from models.common.models.deepseek_r1_distill_qwen_14b.hf_adaptor import DeepSeekR1Qwen14BForCausalLM as DeepSeekProduct
15
+ from models.common.models.deepseek_r1_distill_qwen_14b.hf_adaptor import (
16
+ DeepSeekR1Qwen14BRuntimeConfig,
17
+ _trace_seq_lens,
18
+ convert_hf_model_weights,
19
+ )
20
+
21
+
22
+ def test_runtime_config_preserves_tp2_trace_and_batched_prefill_policy():
23
+ runtime = DeepSeekR1Qwen14BRuntimeConfig(
24
+ model_name="DeepSeek-R1-Distill-Qwen-14B",
25
+ model_cache_path=None,
26
+ max_prefill_chunk_size=2048,
27
+ max_context_len=32768,
28
+ max_seq_len=4096,
29
+ trace_prefill_supported_seq_lens=(128, 1024),
30
+ )
31
+ assert runtime.can_enable_trace(128, num_cached_tokens=32)
32
+ assert runtime.can_enable_trace(1024)
33
+ assert not runtime.can_enable_trace(2048)
34
+ assert runtime.supports_batched_prefill
35
+ assert runtime.max_prefill_batch_size == 32
36
+ assert runtime.batched_prefill_batched_extract
37
+ assert _trace_seq_lens(2, 2048, 4096) == (128, 1024)
38
+ assert _trace_seq_lens(4, 2048, 4096) == (128,)
39
+ assert _trace_seq_lens(8, 2048, 4096) == (128, 1024)
40
+
41
+
42
+ def test_pinned_revision_is_the_provider_and_generator_default():
43
+ expected = "1df8507178afcc1bef68cd8c393f61a886323761"
44
+ assert hf_adaptor.DEFAULT_HF_REVISION == expected
45
+ assert generator.DeepSeekR1Qwen14BGeneratorConfig.__dataclass_fields__["hf_revision"].default == expected
46
+
47
+
48
+ def test_generator_keeps_deepseek_chat_template_enabled():
49
+ source = inspect.getsource(generator.build_deepseek_r1_distill_qwen_14b_generator)
50
+ assert "instruct=True" in source
51
+
52
+
53
+ def test_provider_rejects_below_capacity_before_loading_hf(expect_error):
54
+ mesh = SimpleNamespace(get_num_devices=lambda: 1)
55
+ with expect_error(ValueError, "supports logical TP2/TP4/TP8"):
56
+ hf_adaptor.from_pretrained(mesh)
57
+
58
+
59
+ def test_product_binds_runtime_config_and_stop_tokens():
60
+ model = SimpleNamespace(config=SimpleNamespace(max_seq_len=4096), model_args=None)
61
+ tokenizer = SimpleNamespace(stop_tokens=[151643, 151644])
62
+ runtime = DeepSeekR1Qwen14BRuntimeConfig(
63
+ model_name="model",
64
+ model_cache_path=None,
65
+ max_prefill_chunk_size=2048,
66
+ max_context_len=32768,
67
+ max_seq_len=4096,
68
+ trace_prefill_supported_seq_lens=(128, 1024),
69
+ )
70
+ product = DeepSeekProduct(model=model, tokenizer=tokenizer, runtime_config=runtime)
71
+ assert model.model_args is runtime
72
+ assert product.generation_config.stop_token_ids == (151643, 151644)
73
+ assert product.max_seq_len == 4096
74
+ assert product.max_context_len == 32768
75
+
76
+
77
+ def test_tokenizer_adds_eos_and_threads_revision(monkeypatch):
78
+ tokenizer = SimpleNamespace(
79
+ eos_token_id=151643,
80
+ convert_tokens_to_ids=lambda token: -1,
81
+ )
82
+ seen = {}
83
+
84
+ def fake_from_pretrained(model, **kwargs):
85
+ seen.update(model=model, **kwargs)
86
+ return tokenizer
87
+
88
+ monkeypatch.setattr(hf_adaptor.AutoTokenizer, "from_pretrained", fake_from_pretrained)
89
+ assert hf_adaptor.load_tokenizer("deepseek-ai/DeepSeek-R1-Distill-Qwen-14B", "revision") is tokenizer
90
+ assert tokenizer.stop_tokens == [151643]
91
+ assert seen["revision"] == "revision"
92
+
93
+
94
+ def test_qkv_weights_and_bias_use_reverse_permutation_and_device_major_packing():
95
+ hidden_size = 16
96
+ n_heads = 4
97
+ n_kv_heads = 2
98
+ head_dim = 4
99
+ num_devices = 2
100
+ kv_width = n_kv_heads * head_dim
101
+ q = torch.arange(hidden_size * hidden_size, dtype=torch.float32).reshape(hidden_size, hidden_size)
102
+ k = torch.arange(kv_width * hidden_size, dtype=torch.float32).reshape(kv_width, hidden_size) + 10_000
103
+ v = k + 10_000
104
+ o = q + 30_000
105
+ bq = torch.arange(hidden_size, dtype=torch.float32)
106
+ bk = torch.arange(kv_width, dtype=torch.float32) + 100
107
+ bv = torch.arange(kv_width, dtype=torch.float32) + 200
108
+ attention = SimpleNamespace(
109
+ config=SimpleNamespace(
110
+ hidden_size=hidden_size,
111
+ num_attention_heads=n_heads,
112
+ num_key_value_heads=n_kv_heads,
113
+ ),
114
+ q_proj=SimpleNamespace(weight=q, bias=bq),
115
+ k_proj=SimpleNamespace(weight=k, bias=bk),
116
+ v_proj=SimpleNamespace(weight=v, bias=bv),
117
+ o_proj=SimpleNamespace(weight=o),
118
+ )
119
+
120
+ wqkv, wo, q_norm, k_norm, bias = weight_utils.attention_wqkv_wo_from_hf_layer(attention, num_devices)
121
+ q_meta = weight_utils.reverse_permute(q, n_heads, hidden_size, hidden_size).T
122
+ k_meta = weight_utils.reverse_permute(k, n_kv_heads, kv_width, hidden_size).T
123
+ bq_meta = weight_utils.reverse_permute_1d(bq.view(n_heads, head_dim)).view(-1)
124
+ bk_meta = weight_utils.reverse_permute_1d(bk.view(n_kv_heads, head_dim)).view(-1)
125
+ expected_weights = (
126
+ torch.cat(
127
+ [
128
+ torch.cat(parts, dim=-1)
129
+ for parts in zip(
130
+ torch.chunk(q_meta, num_devices, dim=1),
131
+ torch.chunk(k_meta, num_devices, dim=1),
132
+ torch.chunk(v.T, num_devices, dim=1),
133
+ )
134
+ ],
135
+ dim=-1,
136
+ )
137
+ .unsqueeze(0)
138
+ .unsqueeze(0)
139
+ )
140
+ expected_bias = torch.cat(
141
+ [
142
+ torch.cat(parts, dim=-1)
143
+ for parts in zip(
144
+ torch.chunk(bq_meta, num_devices),
145
+ torch.chunk(bk_meta, num_devices),
146
+ torch.chunk(bv, num_devices),
147
+ )
148
+ ]
149
+ )
150
+
151
+ torch.testing.assert_close(wqkv, expected_weights)
152
+ torch.testing.assert_close(wo, o.T.unsqueeze(0).unsqueeze(0))
153
+ torch.testing.assert_close(bias, expected_bias)
154
+ assert q_norm is None and k_norm is None
155
+
156
+
157
+ def test_hf_rope_tables_preserve_plain_theta_one_million():
158
+ head_dim = 16
159
+ table_len = 128
160
+ config = Qwen2Config(
161
+ hidden_size=64,
162
+ intermediate_size=128,
163
+ num_hidden_layers=1,
164
+ num_attention_heads=4,
165
+ num_key_value_heads=2,
166
+ max_position_embeddings=32768,
167
+ rope_parameters={"rope_type": "default", "rope_theta": 1_000_000.0},
168
+ )
169
+ rotary = Qwen2RotaryEmbedding(config)
170
+ cos, sin = weight_utils.build_rope_cos_sin_torch(rotary, table_len, head_dim, torch.bfloat16)
171
+ x = torch.zeros(1, 1, table_len, head_dim, dtype=torch.bfloat16)
172
+ positions = torch.arange(table_len).unsqueeze(0)
173
+ with torch.no_grad():
174
+ hf_cos, hf_sin = rotary(x, positions)
175
+ expected_cos, expected_sin = weight_utils.permute_hf_rope_to_meta_tables(hf_cos.float(), hf_sin.float())
176
+ assert config.rope_parameters["rope_theta"] == 1_000_000.0
177
+ torch.testing.assert_close(cos, expected_cos.to(torch.bfloat16))
178
+ torch.testing.assert_close(sin, expected_sin.to(torch.bfloat16))
179
+
180
+
181
+ def test_conversion_covers_qkv_bias_and_untied_lm_head():
182
+ config = Qwen2Config(
183
+ hidden_size=64,
184
+ intermediate_size=128,
185
+ num_hidden_layers=1,
186
+ num_attention_heads=4,
187
+ num_key_value_heads=2,
188
+ vocab_size=128,
189
+ max_position_embeddings=32768,
190
+ rope_parameters={"rope_type": "default", "rope_theta": 1_000_000.0},
191
+ tie_word_embeddings=False,
192
+ )
193
+ hf = Qwen2ForCausalLM(config).eval()
194
+ weights = convert_hf_model_weights(
195
+ hf,
196
+ config,
197
+ n_layers=1,
198
+ num_devices=2,
199
+ rope_table_len=128,
200
+ head_dim=16,
201
+ )
202
+ layer = weights.layers[0]
203
+ assert layer.wqkv.shape == (1, 1, 64, 128)
204
+ assert layer.wqkv_bias.shape == (128,)
205
+ assert layer.wo.shape == (1, 1, 64, 64)
206
+ assert layer.w1.shape == layer.w3.shape == (64, 2048)
207
+ assert layer.w2.shape == (2048, 64)
208
+ assert torch.count_nonzero(layer.w1[:, 128:]) == 0
209
+ assert torch.count_nonzero(layer.w3[:, 128:]) == 0
210
+ assert torch.count_nonzero(layer.w2[128:, :]) == 0
211
+ torch.testing.assert_close(weights.lm_head, hf.lm_head.weight.detach().to(torch.bfloat16))
212
+ assert weights.lm_head.data_ptr() != weights.embedding.data_ptr()
213
+
214
+
215
+ def test_config_builder_is_owned_by_model_module():
216
+ assert (
217
+ hf_adaptor.build_deepseek_r1_distill_qwen_14b_transformer_config
218
+ is qwen_model.build_deepseek_r1_distill_qwen_14b_transformer_config
219
+ )
220
+ assert qwen_model.build_deepseek_r1_distill_qwen_14b_transformer_config.__module__ == qwen_model.__name__
221
+
222
+
223
+ def test_post_attention_norm_program_and_memory_use_same_mlp_grid(monkeypatch):
224
+ grid = SimpleNamespace(num_cores=28)
225
+ program = object()
226
+ memory = object()
227
+ captured = {}
228
+
229
+ monkeypatch.setattr(qwen_model, "get_padded_hidden_dim", lambda *_: 18944)
230
+ monkeypatch.setattr(qwen_model, "_dram_shard_core_grid_k_n", lambda *_: grid)
231
+ monkeypatch.setattr(
232
+ qwen_model,
233
+ "_create_sharded_norm_program_config",
234
+ lambda dim, selected_grid, rows, tile: captured.update(program=(dim, selected_grid, rows, tile)) or program,
235
+ )
236
+ monkeypatch.setattr(
237
+ qwen_model.ttnn,
238
+ "create_sharded_memory_config",
239
+ lambda shape, selected_grid, *args, **kwargs: captured.update(memory=(shape, selected_grid)) or memory,
240
+ )
241
+
242
+ assert qwen_model._post_attn_norm_decode_configs(
243
+ dim=3584,
244
+ hidden_dim=18944,
245
+ num_devices=2,
246
+ max_batch_size=32,
247
+ ) == (program, memory)
248
+ assert captured["program"] == (3584, grid, 32, 32)
249
+ assert captured["memory"] == ((32, 128), grid)
250
+
251
+
252
+ def test_decoder_layer_prefill_calls_chunk_capable_attention_entrypoint(monkeypatch):
253
+ captured = {}
254
+ attention_output = object()
255
+ final_output = object()
256
+ attention = SimpleNamespace(
257
+ prefill_forward=lambda x, rot_mats, **kwargs: captured.update(attention=(x, rot_mats, kwargs))
258
+ or attention_output
259
+ )
260
+ layer = qwen_model.DeepSeekR1Qwen14BDecoderLayer(
261
+ input_layernorm=SimpleNamespace(prefill_forward=lambda x: x),
262
+ self_attn=attention,
263
+ post_attention_layernorm=SimpleNamespace(prefill_forward=lambda x: x),
264
+ mlp=SimpleNamespace(prefill_forward=lambda x: x),
265
+ )
266
+ monkeypatch.setattr(qwen_model, "_all_gather_rmsnorm_tensor", lambda _norm, x: x)
267
+ monkeypatch.setattr(
268
+ qwen_model.ttnn,
269
+ "add",
270
+ lambda *_args, **_kwargs: final_output,
271
+ )
272
+
273
+ chunk_start_idx_tensor = object()
274
+ rot_mats = (object(), object())
275
+ assert (
276
+ layer.prefill_forward(
277
+ object(),
278
+ rot_mats,
279
+ user_id=[0, 1],
280
+ page_table=object(),
281
+ chunk_page_table=object(),
282
+ chunk_start_idx=128,
283
+ batch_size=2,
284
+ chunk_start_idx_tensor=chunk_start_idx_tensor,
285
+ )
286
+ is final_output
287
+ )
288
+ assert captured["attention"][1] is rot_mats
289
+ assert captured["attention"][2]["chunk_start_idx_tensor"] is chunk_start_idx_tensor
290
+ assert captured["attention"][2]["batch_size"] == 2
code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_prefill_last_token_contract.py ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+
6
+ from models.common.models.deepseek_r1_distill_qwen_14b import model as qwen_model
7
+
8
+
9
+ def test_prefill_runtime_slice_and_index_override_full_hidden_state_return(monkeypatch):
10
+ calls = []
11
+ hidden = SimpleNamespace(shape=(1, 1, 128, 640), dtype=qwen_model.ttnn.bfloat16)
12
+ sliced = SimpleNamespace(dtype=qwen_model.ttnn.bfloat16)
13
+ selected = object()
14
+ selected_4d = object()
15
+ logits = object()
16
+ slice_start = object()
17
+ slice_end = object()
18
+ last_token_index = object()
19
+ model = SimpleNamespace(
20
+ layers=[],
21
+ _last_tile_logits=lambda value: calls.append(("last_tile_logits", value)) or logits,
22
+ )
23
+
24
+ monkeypatch.setattr(
25
+ qwen_model.ttnn,
26
+ "slice",
27
+ lambda value, start, end, **kwargs: calls.append(("slice", value, start, end, kwargs)) or sliced,
28
+ )
29
+ monkeypatch.setattr(
30
+ qwen_model.ttnn,
31
+ "embedding",
32
+ lambda index, value, **kwargs: calls.append(("embedding", index, value, kwargs)) or selected,
33
+ )
34
+ monkeypatch.setattr(
35
+ qwen_model.ttnn,
36
+ "unsqueeze_to_4D",
37
+ lambda value: calls.append(("unsqueeze_to_4D", value)) or selected_4d,
38
+ )
39
+ monkeypatch.setattr(qwen_model.ttnn, "deallocate", lambda value: calls.append(("deallocate", value)))
40
+
41
+ result = qwen_model.DeepSeekR1Qwen14B.prefill_forward(
42
+ model,
43
+ hidden,
44
+ rot_mats=(object(), object()),
45
+ get_last_token=-1,
46
+ last_token_slice=(slice_start, slice_end),
47
+ last_token_index=last_token_index,
48
+ )
49
+
50
+ assert result is logits
51
+ assert calls == [
52
+ ("slice", hidden, slice_start, slice_end, {"slice_dim": 2, "num_devices": 4}),
53
+ ("deallocate", hidden),
54
+ ("embedding", last_token_index, sliced, {"layout": qwen_model.ttnn.TILE_LAYOUT}),
55
+ ("unsqueeze_to_4D", selected),
56
+ ("deallocate", sliced),
57
+ ("last_tile_logits", selected_4d),
58
+ ]
59
+
60
+
61
+ def test_prefill_runtime_index_requires_runtime_slice(expect_error):
62
+ model = SimpleNamespace(layers=[])
63
+
64
+ with expect_error(ValueError, "last_token_index is required with a runtime last_token_slice"):
65
+ qwen_model.DeepSeekR1Qwen14B.prefill_forward(
66
+ model,
67
+ object(),
68
+ rot_mats=(object(), object()),
69
+ get_last_token=-1,
70
+ last_token_index=object(),
71
+ )
code/models/common/tests/models/llama32_1b/test_batched_prefill_postprocess.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+
6
+ import torch
7
+
8
+ import models.common.models.llama32_1b.model as model_module
9
+ from models.common.models.llama32_1b.model import Llama32_1BTransformer1D
10
+
11
+
12
+ def test_batched_prefill_postprocess_gathers_each_slots_last_token_before_norm(monkeypatch):
13
+ calls = []
14
+ hidden = object()
15
+ selector_tt = object()
16
+ gathered = object()
17
+ normalized = object()
18
+ all_gathered = object()
19
+ logits = object()
20
+ output = object()
21
+ mesh = SimpleNamespace(arch=lambda: "wormhole")
22
+
23
+ class FakeTTNN:
24
+ bfloat16 = "bfloat16"
25
+ TILE_LAYOUT = "tile"
26
+ DRAM_MEMORY_CONFIG = "dram"
27
+ MathFidelity = SimpleNamespace(HiFi4="hifi4")
28
+
29
+ @staticmethod
30
+ def ReplicateTensorToMesh(device):
31
+ assert device is mesh
32
+ return "replicate"
33
+
34
+ @staticmethod
35
+ def from_torch(selector, **kwargs):
36
+ calls.append(("from_torch", selector.clone(), kwargs))
37
+ return selector_tt
38
+
39
+ @staticmethod
40
+ def init_device_compute_kernel_config(arch, **kwargs):
41
+ assert arch == "wormhole"
42
+ return kwargs
43
+
44
+ @staticmethod
45
+ def matmul(lhs, rhs, **kwargs):
46
+ calls.append(("matmul", lhs, rhs, kwargs))
47
+ return gathered
48
+
49
+ @staticmethod
50
+ def deallocate(tensor):
51
+ calls.append(("deallocate", tensor))
52
+
53
+ @staticmethod
54
+ def to_memory_config(tensor, memory_config):
55
+ calls.append(("to_memory_config", tensor, memory_config))
56
+ return output
57
+
58
+ fake_norm = SimpleNamespace(prefill_forward=lambda tensor: calls.append(("norm", tensor)) or normalized)
59
+ fake_lm_head = SimpleNamespace(
60
+ config=SimpleNamespace(input_memcfg=None),
61
+ forward=lambda tensor: calls.append(("lm_head", tensor)) or logits,
62
+ )
63
+ model = SimpleNamespace(mesh_device=mesh, norm=fake_norm, lm_head=fake_lm_head)
64
+ monkeypatch.setattr(model_module, "ttnn", FakeTTNN)
65
+ monkeypatch.setattr(
66
+ model_module,
67
+ "_all_gather_rmsnorm_tensor",
68
+ lambda norm, tensor: calls.append(("all_gather", norm, tensor)) or all_gathered,
69
+ )
70
+
71
+ result = Llama32_1BTransformer1D.post_process_batched_prefill_output(
72
+ model,
73
+ hidden,
74
+ last_token_idx_list=[3, 7, 11, 0],
75
+ padded_batch=4,
76
+ prefill_seq_len=32,
77
+ )
78
+
79
+ assert result is output
80
+ selector = calls[0][1]
81
+ assert selector.shape == (1, 1, 32, 128)
82
+ assert selector.dtype == torch.bfloat16
83
+ assert torch.count_nonzero(selector).item() == 4
84
+ assert [selector[0, 0, row].nonzero().item() for row in range(4)] == [3, 39, 75, 96]
85
+ assert calls[0][2] == {
86
+ "device": mesh,
87
+ "dtype": "bfloat16",
88
+ "layout": "tile",
89
+ "mesh_mapper": "replicate",
90
+ }
91
+ assert [call[0] for call in calls] == [
92
+ "from_torch",
93
+ "matmul",
94
+ "deallocate",
95
+ "norm",
96
+ "all_gather",
97
+ "lm_head",
98
+ "to_memory_config",
99
+ ]
100
+ assert calls[1][1:3] == (selector_tt, hidden)
101
+ assert calls[2] == ("deallocate", selector_tt)
102
+ assert calls[3] == ("norm", gathered)
103
+ assert calls[4] == ("all_gather", fake_norm, normalized)
104
+ assert calls[5] == ("lm_head", all_gathered)
code/models/common/tests/models/llama32_1b/test_demo_warmup.py ADDED
@@ -0,0 +1,134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import ast
5
+ from pathlib import Path
6
+ from types import SimpleNamespace
7
+
8
+ import pytest
9
+
10
+ from models.common.llm_runtime.config import TraceConfig
11
+
12
+ _DEMO_PATH = "models/common/tests/demos/llama32_1b/demo.py"
13
+ _DEMO_TREE = ast.parse(Path(_DEMO_PATH).read_text(encoding="utf-8"), filename=_DEMO_PATH)
14
+
15
+
16
+ def _demo_function(name, namespace=None):
17
+ function = next(node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == name)
18
+ namespace = {} if namespace is None else namespace
19
+ exec(compile(ast.Module(body=[function], type_ignores=[]), _DEMO_PATH, "exec"), namespace)
20
+ return namespace[name]
21
+
22
+
23
+ _warmup_demo_executor = _demo_function("_warmup_demo_executor")
24
+
25
+
26
+ @pytest.mark.parametrize("lane_group", [False, True])
27
+ def test_demo_warmup_compiles_eager_programs_before_trace_capture(lane_group):
28
+ calls = []
29
+ config = SimpleNamespace(trace=TraceConfig("all"), device_sampling_enabled=True)
30
+
31
+ def warmup_prefill(**kwargs):
32
+ calls.append(("prefill", kwargs))
33
+
34
+ def warmup_decode(**kwargs):
35
+ calls.append(("decode", kwargs))
36
+
37
+ executor = SimpleNamespace(
38
+ warmup_model_prefill=warmup_prefill,
39
+ warmup_model_decode=warmup_decode,
40
+ max_batch_size=4,
41
+ )
42
+ if lane_group:
43
+ executor.lanes = [SimpleNamespace(config=config)]
44
+ else:
45
+ executor.config = config
46
+ executor.model = SimpleNamespace(config=SimpleNamespace(max_batch_size=4))
47
+
48
+ kv_cache = object()
49
+ page_table = SimpleNamespace(shape=(4, 8))
50
+ _warmup_demo_executor(executor, kv_cache=kv_cache, page_table=page_table)
51
+
52
+ assert [(kind, kwargs["enable_trace"]) for kind, kwargs in calls] == [
53
+ ("decode", False),
54
+ ("prefill", False),
55
+ ("prefill", True),
56
+ ("decode", True),
57
+ ]
58
+ for _, kwargs in calls:
59
+ assert kwargs["kv_cache"] is kv_cache
60
+ assert kwargs["can_sample_on_device"] is True
61
+ for kind, kwargs in calls:
62
+ if kind == "decode":
63
+ assert kwargs["max_batch_size"] == 4
64
+ assert kwargs["num_blocks"] == 8
65
+
66
+
67
+ def _called_names(function_name):
68
+ function = next(
69
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == function_name
70
+ )
71
+ return [
72
+ node.func.id for node in ast.walk(function) if isinstance(node, ast.Call) and isinstance(node.func, ast.Name)
73
+ ]
74
+
75
+
76
+ @pytest.mark.parametrize(("data_parallel", "expected_tp_devices"), [(4, 2), (8, 1)])
77
+ def test_t3k_dp_topology_preserves_supported_tp_lanes(data_parallel, expected_tp_devices):
78
+ helper = _demo_function("_dp_tp_devices_or_skip", {"pytest": pytest, "ttnn": SimpleNamespace(MeshDevice=object)})
79
+ mesh = SimpleNamespace(get_num_devices=lambda: 8)
80
+
81
+ assert helper(mesh, data_parallel) == expected_tp_devices
82
+
83
+
84
+ def test_t3k_dp2_skips_unsupported_tp4_lanes(expect_error):
85
+ helper = _demo_function("_dp_tp_devices_or_skip", {"pytest": pytest, "ttnn": SimpleNamespace(MeshDevice=object)})
86
+ mesh = SimpleNamespace(get_num_devices=lambda: 8)
87
+
88
+ with expect_error(pytest.skip.Exception, "creates TP4 lanes"):
89
+ helper(mesh, 2)
90
+
91
+
92
+ def test_dp_build_validates_and_resolves_cache_from_each_lane_submesh():
93
+ function = next(
94
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_dp_smoke"
95
+ )
96
+ lane_loop = next(
97
+ node
98
+ for node in ast.walk(function)
99
+ if isinstance(node, ast.For) and isinstance(node.target, ast.Name) and node.target.id == "sm"
100
+ )
101
+ calls = [node for node in ast.walk(lane_loop) if isinstance(node, ast.Call) and isinstance(node.func, ast.Name)]
102
+ call_names = [node.func.id for node in calls]
103
+ assert "_skip_unless_heads_divide_mesh" in call_names
104
+ assert "lazy_weight_cache_dir_for_demo" in call_names
105
+
106
+ from_pretrained_call = next(node for node in calls if node.func.id == "from_pretrained")
107
+ cache_dir = next(keyword.value for keyword in from_pretrained_call.keywords if keyword.arg == "cache_dir")
108
+ assert isinstance(cache_dir, ast.Name)
109
+ assert cache_dir.id == "lane_cache_dir"
110
+
111
+
112
+ @pytest.mark.parametrize("function_name", ["_run_perf_benchmark", "_run_dp_smoke"])
113
+ def test_traced_demo_paths_warm_up_before_benchmark(function_name):
114
+ calls = _called_names(function_name)
115
+ assert calls.index("_warmup_demo_executor") < calls.index("run_perf_benchmark")
116
+
117
+
118
+ def test_eval_repeat_warms_each_fresh_executor():
119
+ calls = _called_names("_run_eval_repeat_batch32")
120
+ assert "_warmup_demo_executor" in calls
121
+
122
+
123
+ def test_perf_path_enables_pipeline_readback_by_default():
124
+ function = next(
125
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_perf_benchmark"
126
+ )
127
+ benchmark_call = next(
128
+ node
129
+ for node in ast.walk(function)
130
+ if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "run_perf_benchmark"
131
+ )
132
+ keywords = {keyword.arg: keyword.value for keyword in benchmark_call.keywords}
133
+ assert isinstance(keywords["pipeline_readback"], ast.Name)
134
+ assert keywords["pipeline_readback"].id == "pipeline_readback"
code/models/common/tests/models/llama32_1b/test_hf_adaptor.py ADDED
@@ -0,0 +1,264 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+
6
+ import torch
7
+ from transformers import LlamaConfig, LlamaForCausalLM
8
+ from transformers.models.llama.modeling_llama import LlamaRotaryEmbedding
9
+
10
+ from models.common.models.llama32_1b import hf_adaptor
11
+ from models.common.models.llama32_1b import model as llama_model
12
+ from models.common.models.llama32_1b import weight_utils
13
+ from models.common.models.llama32_1b.hf_adaptor import (
14
+ Llama32_1BForCausalLM,
15
+ Llama32_1BRuntimeConfig,
16
+ _trace_seq_lens,
17
+ convert_hf_model_weights,
18
+ )
19
+
20
+ LLAMA32_ROPE_PARAMETERS = {
21
+ "rope_type": "llama3",
22
+ "factor": 32.0,
23
+ "low_freq_factor": 1.0,
24
+ "high_freq_factor": 4.0,
25
+ "original_max_position_embeddings": 8192,
26
+ "rope_theta": 500000.0,
27
+ }
28
+
29
+
30
+ def test_runtime_config_preserves_trace_and_batched_prefill_policy():
31
+ runtime = Llama32_1BRuntimeConfig(
32
+ model_name="Llama-3.2-1B-Instruct",
33
+ model_cache_path=None,
34
+ max_prefill_chunk_size=2048,
35
+ max_context_len=131072,
36
+ max_seq_len=4096,
37
+ trace_prefill_supported_seq_lens=(128, 1024),
38
+ )
39
+ assert runtime.can_enable_trace(128, num_cached_tokens=32)
40
+ assert runtime.can_enable_trace(1024)
41
+ assert not runtime.can_enable_trace(2048)
42
+ assert runtime.supports_batched_prefill
43
+ assert runtime.max_prefill_batch_size == 32
44
+ assert runtime.batched_prefill_batched_extract
45
+
46
+
47
+ def test_trace_matrix_is_device_specific_and_bounded():
48
+ assert _trace_seq_lens(1, 2048, 4096) == (128,)
49
+ assert _trace_seq_lens(2, 2048, 4096) == (128, 1024)
50
+ assert _trace_seq_lens(8, 2048, 4096) == (128, 1024)
51
+
52
+
53
+ def test_product_binds_runtime_config_unconditionally():
54
+ model = SimpleNamespace(config=SimpleNamespace(max_seq_len=4096), model_args=None)
55
+ tokenizer = SimpleNamespace(stop_tokens=[128001])
56
+ runtime = Llama32_1BRuntimeConfig(
57
+ model_name="model",
58
+ model_cache_path=None,
59
+ max_prefill_chunk_size=2048,
60
+ max_context_len=131072,
61
+ max_seq_len=4096,
62
+ trace_prefill_supported_seq_lens=(128,),
63
+ )
64
+ product = Llama32_1BForCausalLM(model=model, tokenizer=tokenizer, runtime_config=runtime)
65
+ assert model.model_args is runtime
66
+ assert product.generation_config.stop_token_ids == (128001,)
67
+ assert product.model_name == "model"
68
+ assert product.model_cache_path is None
69
+ assert product.max_seq_len == 4096
70
+ assert product.max_context_len == 131072
71
+
72
+
73
+ def test_hf_attention_and_mlp_weights_match_reference_layouts():
74
+ hidden_size = 128
75
+ num_attention_heads = 32
76
+ num_key_value_heads = 8
77
+ num_devices = 8
78
+ head_dim = hidden_size // num_attention_heads
79
+ kv_width = num_key_value_heads * head_dim
80
+ config = SimpleNamespace(
81
+ num_attention_heads=num_attention_heads,
82
+ num_key_value_heads=num_key_value_heads,
83
+ hidden_size=hidden_size,
84
+ )
85
+ q = torch.arange(hidden_size * hidden_size, dtype=torch.float32).reshape(hidden_size, hidden_size)
86
+ k = torch.arange(kv_width * hidden_size, dtype=torch.float32).reshape(kv_width, hidden_size) + 100_000
87
+ v = k + 100_000
88
+ o = q + 300_000
89
+ attention = SimpleNamespace(
90
+ config=config,
91
+ q_proj=SimpleNamespace(weight=q),
92
+ k_proj=SimpleNamespace(weight=k),
93
+ v_proj=SimpleNamespace(weight=v),
94
+ o_proj=SimpleNamespace(weight=o),
95
+ )
96
+
97
+ wqkv, wo = weight_utils.attention_wqkv_wo_from_hf_layer(attention, num_devices=num_devices)
98
+ q_meta = q.view(num_attention_heads, 2, head_dim // 2, hidden_size).transpose(1, 2).reshape(q.shape).T
99
+ k_meta = k.view(num_key_value_heads, 2, head_dim // 2, hidden_size).transpose(1, 2).reshape(k.shape).T
100
+ expected_qkv = (
101
+ torch.cat(
102
+ [
103
+ torch.cat(parts, dim=-1)
104
+ for parts in zip(
105
+ torch.chunk(q_meta, num_devices, dim=1),
106
+ torch.chunk(k_meta, num_devices, dim=1),
107
+ torch.chunk(v.T, num_devices, dim=1),
108
+ )
109
+ ],
110
+ dim=-1,
111
+ )
112
+ .unsqueeze(0)
113
+ .unsqueeze(0)
114
+ )
115
+ assert wqkv.shape == (1, 1, hidden_size, hidden_size + 2 * kv_width)
116
+ torch.testing.assert_close(wqkv, expected_qkv)
117
+ torch.testing.assert_close(wo, o.T.unsqueeze(0).unsqueeze(0))
118
+
119
+ gate = torch.arange(48, dtype=torch.float32).reshape(6, 8)
120
+ down = torch.arange(48, dtype=torch.float32).reshape(8, 6)
121
+ up = gate + 100
122
+ mlp = SimpleNamespace(
123
+ gate_proj=SimpleNamespace(weight=gate),
124
+ down_proj=SimpleNamespace(weight=down),
125
+ up_proj=SimpleNamespace(weight=up),
126
+ )
127
+ w1, w2, w3 = weight_utils.mlp_weights_from_hf_layer(mlp)
128
+ torch.testing.assert_close(w1, gate.T)
129
+ torch.testing.assert_close(w2, down.T)
130
+ torch.testing.assert_close(w3, up.T)
131
+
132
+
133
+ def test_hf_rope_tables_match_real_llama32_scaled_rotary_reference():
134
+ head_dim = 64
135
+ table_len = LLAMA32_ROPE_PARAMETERS["original_max_position_embeddings"] + 128
136
+ config = LlamaConfig(
137
+ hidden_size=128,
138
+ intermediate_size=256,
139
+ num_hidden_layers=1,
140
+ num_attention_heads=2,
141
+ num_key_value_heads=2,
142
+ head_dim=head_dim,
143
+ max_position_embeddings=131072,
144
+ rope_parameters=LLAMA32_ROPE_PARAMETERS,
145
+ )
146
+ rotary = LlamaRotaryEmbedding(config)
147
+
148
+ cos, sin = weight_utils.build_rope_cos_sin_torch(
149
+ rotary, table_len=table_len, head_dim=head_dim, dtype=torch.bfloat16
150
+ )
151
+ x = torch.zeros(1, 1, table_len, head_dim, dtype=torch.bfloat16)
152
+ position_ids = torch.arange(table_len, dtype=torch.long).unsqueeze(0)
153
+ with torch.no_grad():
154
+ hf_cos, hf_sin = rotary(x, position_ids)
155
+ expected_cos = hf_cos.float().squeeze(0)[:, : head_dim // 2].repeat_interleave(2, dim=-1).unsqueeze(0).unsqueeze(0)
156
+ expected_sin = hf_sin.float().squeeze(0)[:, : head_dim // 2].repeat_interleave(2, dim=-1).unsqueeze(0).unsqueeze(0)
157
+
158
+ assert config.rope_parameters == LLAMA32_ROPE_PARAMETERS
159
+ assert cos.shape == sin.shape == (1, 1, table_len, head_dim)
160
+ assert cos.dtype == sin.dtype == torch.bfloat16
161
+ torch.testing.assert_close(cos, expected_cos.to(torch.bfloat16))
162
+ torch.testing.assert_close(sin, expected_sin.to(torch.bfloat16))
163
+
164
+
165
+ def test_convert_hf_model_weights_covers_real_nonempty_llama_layer():
166
+ config = LlamaConfig(
167
+ hidden_size=128,
168
+ intermediate_size=256,
169
+ num_hidden_layers=1,
170
+ num_attention_heads=32,
171
+ num_key_value_heads=8,
172
+ head_dim=4,
173
+ vocab_size=128,
174
+ max_position_embeddings=131072,
175
+ rope_parameters=LLAMA32_ROPE_PARAMETERS,
176
+ tie_word_embeddings=True,
177
+ )
178
+ hf = LlamaForCausalLM(config).eval()
179
+ weights = convert_hf_model_weights(
180
+ hf,
181
+ config,
182
+ n_layers=1,
183
+ num_devices=8,
184
+ rope_table_len=128,
185
+ head_dim=4,
186
+ )
187
+
188
+ assert len(weights.layers) == 1
189
+ layer_weights = weights.layers[0]
190
+ assert layer_weights.wqkv.shape == (1, 1, 128, 192)
191
+ assert layer_weights.wo.shape == (1, 1, 128, 128)
192
+ assert layer_weights.w1.shape == (128, 256)
193
+ assert layer_weights.w2.shape == (256, 128)
194
+ assert layer_weights.w3.shape == (128, 256)
195
+ assert layer_weights.attention_norm.shape == layer_weights.ff_norm.shape == (128,)
196
+ assert weights.embedding.shape == (1, 1, 128, 128)
197
+ assert weights.rope_cos.shape == weights.rope_sin.shape == (1, 1, 128, 4)
198
+ assert weights.final_norm.shape == (128,)
199
+ torch.testing.assert_close(weights.lm_head, hf.model.embed_tokens.weight.detach().to(torch.bfloat16))
200
+
201
+
202
+ def test_tied_embedding_is_explicit_lm_head_construction_source():
203
+ class Rotary:
204
+ def __call__(self, x, position_ids):
205
+ return torch.ones(1, position_ids.shape[-1], x.shape[-1]), torch.zeros(
206
+ 1, position_ids.shape[-1], x.shape[-1]
207
+ )
208
+
209
+ tied_weight = torch.arange(24, dtype=torch.float32).reshape(6, 4)
210
+ decoy_lm_head = torch.full((6, 4), -99.0)
211
+ base = SimpleNamespace(
212
+ embed_tokens=SimpleNamespace(weight=tied_weight),
213
+ rotary_emb=Rotary(),
214
+ layers=[],
215
+ norm=SimpleNamespace(weight=torch.ones(4)),
216
+ )
217
+ hf = SimpleNamespace(model=base, lm_head=SimpleNamespace(weight=decoy_lm_head))
218
+ config = SimpleNamespace(tie_word_embeddings=True)
219
+ weights = convert_hf_model_weights(
220
+ hf,
221
+ config,
222
+ n_layers=0,
223
+ num_devices=1,
224
+ rope_table_len=8,
225
+ head_dim=4,
226
+ )
227
+
228
+ torch.testing.assert_close(weights.lm_head, tied_weight.to(torch.bfloat16))
229
+ assert not torch.equal(weights.lm_head, decoy_lm_head.to(torch.bfloat16))
230
+ assert weights.embedding.shape == (1, 1, 6, 4)
231
+
232
+
233
+ def test_config_builder_is_owned_by_model_module():
234
+ assert hf_adaptor.build_llama32_1b_transformer_1d_config is llama_model.build_llama32_1b_transformer_1d_config
235
+ assert llama_model.build_llama32_1b_transformer_1d_config.__module__ == llama_model.__name__
236
+
237
+
238
+ def test_post_attention_norm_program_and_memory_use_same_mlp_grid(monkeypatch):
239
+ grid = SimpleNamespace(num_cores=64)
240
+ program = object()
241
+ memory = object()
242
+ captured = {}
243
+
244
+ monkeypatch.setattr(llama_model, "get_padded_hidden_dim", lambda *_: 8192)
245
+ monkeypatch.setattr(llama_model, "_dram_shard_core_grid_k_n", lambda *_: grid)
246
+ monkeypatch.setattr(
247
+ llama_model,
248
+ "_create_sharded_norm_program_config",
249
+ lambda dim, selected_grid, rows, tile: captured.update(program=(dim, selected_grid, rows, tile)) or program,
250
+ )
251
+ monkeypatch.setattr(
252
+ llama_model.ttnn,
253
+ "create_sharded_memory_config",
254
+ lambda shape, selected_grid, *args, **kwargs: captured.update(memory=(shape, selected_grid)) or memory,
255
+ )
256
+
257
+ assert llama_model._post_attn_norm_decode_configs(
258
+ dim=2048,
259
+ hidden_dim=8192,
260
+ num_devices=1,
261
+ max_batch_size=1,
262
+ ) == (program, memory)
263
+ assert captured["program"] == (2048, grid, 32, 32)
264
+ assert captured["memory"] == ((32, 32), grid)
code/models/common/tests/models/llama32_3b/test_batched_prefill_postprocess.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ from types import SimpleNamespace
5
+
6
+ import torch
7
+
8
+ import models.common.models.llama32_3b.model as model_module
9
+ from models.common.models.llama32_3b.model import Llama32_3BTransformer1D
10
+
11
+
12
+ def test_batched_prefill_postprocess_gathers_each_slots_last_token_before_norm(monkeypatch):
13
+ calls = []
14
+ hidden = object()
15
+ selector_tt = object()
16
+ gathered = object()
17
+ normalized = object()
18
+ all_gathered = object()
19
+ logits = object()
20
+ output = object()
21
+ mesh = SimpleNamespace(arch=lambda: "wormhole")
22
+
23
+ class FakeTTNN:
24
+ bfloat16 = "bfloat16"
25
+ TILE_LAYOUT = "tile"
26
+ DRAM_MEMORY_CONFIG = "dram"
27
+ MathFidelity = SimpleNamespace(HiFi4="hifi4")
28
+
29
+ @staticmethod
30
+ def ReplicateTensorToMesh(device):
31
+ assert device is mesh
32
+ return "replicate"
33
+
34
+ @staticmethod
35
+ def from_torch(selector, **kwargs):
36
+ calls.append(("from_torch", selector.clone(), kwargs))
37
+ return selector_tt
38
+
39
+ @staticmethod
40
+ def init_device_compute_kernel_config(arch, **kwargs):
41
+ assert arch == "wormhole"
42
+ return kwargs
43
+
44
+ @staticmethod
45
+ def matmul(lhs, rhs, **kwargs):
46
+ calls.append(("matmul", lhs, rhs, kwargs))
47
+ return gathered
48
+
49
+ @staticmethod
50
+ def deallocate(tensor):
51
+ calls.append(("deallocate", tensor))
52
+
53
+ @staticmethod
54
+ def to_memory_config(tensor, memory_config):
55
+ calls.append(("to_memory_config", tensor, memory_config))
56
+ return output
57
+
58
+ fake_norm = SimpleNamespace(prefill_forward=lambda tensor: calls.append(("norm", tensor)) or normalized)
59
+ fake_lm_head = SimpleNamespace(
60
+ config=SimpleNamespace(input_memcfg=None),
61
+ forward=lambda tensor: calls.append(("lm_head", tensor)) or logits,
62
+ )
63
+ model = SimpleNamespace(mesh_device=mesh, norm=fake_norm, lm_head=fake_lm_head)
64
+ monkeypatch.setattr(model_module, "ttnn", FakeTTNN)
65
+ monkeypatch.setattr(
66
+ model_module,
67
+ "_all_gather_rmsnorm_tensor",
68
+ lambda norm, tensor: calls.append(("all_gather", norm, tensor)) or all_gathered,
69
+ )
70
+
71
+ result = Llama32_3BTransformer1D.post_process_batched_prefill_output(
72
+ model,
73
+ hidden,
74
+ last_token_idx_list=[3, 7, 11, 0],
75
+ padded_batch=4,
76
+ prefill_seq_len=32,
77
+ )
78
+
79
+ assert result is output
80
+ selector = calls[0][1]
81
+ assert selector.shape == (1, 1, 32, 128)
82
+ assert selector.dtype == torch.bfloat16
83
+ assert torch.count_nonzero(selector).item() == 4
84
+ assert [selector[0, 0, row].nonzero().item() for row in range(4)] == [3, 39, 75, 96]
85
+ assert calls[0][2] == {
86
+ "device": mesh,
87
+ "dtype": "bfloat16",
88
+ "layout": "tile",
89
+ "mesh_mapper": "replicate",
90
+ }
91
+ assert [call[0] for call in calls] == [
92
+ "from_torch",
93
+ "matmul",
94
+ "deallocate",
95
+ "norm",
96
+ "all_gather",
97
+ "lm_head",
98
+ "to_memory_config",
99
+ ]
100
+ assert calls[1][1:3] == (selector_tt, hidden)
101
+ assert calls[2] == ("deallocate", selector_tt)
102
+ assert calls[3] == ("norm", gathered)
103
+ assert calls[4] == ("all_gather", fake_norm, normalized)
104
+ assert calls[5] == ("lm_head", all_gathered)
code/models/common/tests/models/llama32_3b/test_demo_warmup.py ADDED
@@ -0,0 +1,218 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
2
+ # SPDX-License-Identifier: Apache-2.0
3
+
4
+ import ast
5
+ import os
6
+ from pathlib import Path
7
+ from types import SimpleNamespace
8
+
9
+ import pytest
10
+
11
+ from models.common.llm_runtime.config import TraceConfig
12
+
13
+ _DEMO_PATH = "models/common/tests/demos/llama32_3b/demo.py"
14
+ _DEMO_TREE = ast.parse(Path(_DEMO_PATH).read_text(encoding="utf-8"), filename=_DEMO_PATH)
15
+
16
+
17
+ def _demo_function(name, namespace=None):
18
+ function = next(node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == name)
19
+ namespace = {} if namespace is None else namespace
20
+ exec(compile(ast.Module(body=[function], type_ignores=[]), _DEMO_PATH, "exec"), namespace)
21
+ return namespace[name]
22
+
23
+
24
+ _warmup_demo_executor = _demo_function("_warmup_demo_executor")
25
+
26
+
27
+ @pytest.mark.parametrize("lane_group", [False, True])
28
+ @pytest.mark.parametrize(
29
+ ("trace_mode", "expected_trace_calls"),
30
+ [
31
+ ("all", [("prefill", True), ("decode", True)]),
32
+ ("decode_only", [("decode", True)]),
33
+ ],
34
+ )
35
+ def test_demo_warmup_compiles_eager_programs_before_enabled_trace_capture(lane_group, trace_mode, expected_trace_calls):
36
+ calls = []
37
+ config = SimpleNamespace(trace=TraceConfig(trace_mode), device_sampling_enabled=True)
38
+
39
+ def warmup_prefill(**kwargs):
40
+ calls.append(("prefill", kwargs))
41
+
42
+ def warmup_decode(**kwargs):
43
+ calls.append(("decode", kwargs))
44
+
45
+ executor = SimpleNamespace(
46
+ warmup_model_prefill=warmup_prefill,
47
+ warmup_model_decode=warmup_decode,
48
+ max_batch_size=4,
49
+ )
50
+ if lane_group:
51
+ executor.lanes = [SimpleNamespace(config=config)]
52
+ else:
53
+ executor.config = config
54
+ executor.model = SimpleNamespace(config=SimpleNamespace(max_batch_size=4))
55
+
56
+ kv_cache = object()
57
+ page_table = SimpleNamespace(shape=(4, 8))
58
+ _warmup_demo_executor(executor, kv_cache=kv_cache, page_table=page_table)
59
+
60
+ eager_calls = [("decode", False), ("prefill", False)]
61
+ assert [(kind, kwargs["enable_trace"]) for kind, kwargs in calls] == eager_calls + expected_trace_calls
62
+ for _, kwargs in calls:
63
+ assert kwargs["kv_cache"] is kv_cache
64
+ assert kwargs["can_sample_on_device"] is True
65
+ for kind, kwargs in calls:
66
+ if kind == "decode":
67
+ assert kwargs["max_batch_size"] == 4
68
+ assert kwargs["num_blocks"] == 8
69
+
70
+
71
+ @pytest.mark.parametrize(
72
+ ("num_devices", "traced", "expected_mode"),
73
+ [(1, True, "decode_only"), (2, True, "all"), (8, True, "all"), (1, False, "none")],
74
+ )
75
+ def test_create_executor_preserves_3b_trace_device_matrix(num_devices, traced, expected_mode):
76
+ captured = {}
77
+
78
+ def executor_config(**kwargs):
79
+ captured.update(kwargs)
80
+ return SimpleNamespace(**kwargs)
81
+
82
+ namespace = {
83
+ "Llama32_3BTransformer1D": object,
84
+ "Llama32_3BExecutor": lambda model, model_args, config: config,
85
+ "Llama32_3BExecutorConfig": executor_config,
86
+ "PagedKVCacheConfig": lambda **kwargs: SimpleNamespace(**kwargs),
87
+ "TraceConfig": TraceConfig,
88
+ "WarmupConfig": lambda: object(),
89
+ }
90
+ create_executor = _demo_function("create_executor", namespace)
91
+ model = SimpleNamespace(
92
+ model_args=object(),
93
+ config=SimpleNamespace(
94
+ max_seq_len=4096,
95
+ max_batch_size=32,
96
+ num_devices=num_devices,
97
+ block_configs=[SimpleNamespace(attention_config=SimpleNamespace(kv_cache_dtype=object()))],
98
+ ),
99
+ )
100
+
101
+ result = create_executor(model, traced=traced, device_sampling_enabled=True)
102
+
103
+ assert result.trace.mode == expected_mode
104
+ assert captured["device_sampling_enabled"] is True
105
+ assert captured["paged_kv_cache"].num_blocks == 4096
106
+
107
+
108
+ def _called_names(function_name):
109
+ function = next(
110
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == function_name
111
+ )
112
+ return [
113
+ node.func.id for node in ast.walk(function) if isinstance(node, ast.Call) and isinstance(node.func, ast.Name)
114
+ ]
115
+
116
+
117
+ @pytest.mark.parametrize(("data_parallel", "expected_tp_devices"), [(4, 2), (8, 1)])
118
+ def test_t3k_dp_topology_preserves_supported_tp_lanes(data_parallel, expected_tp_devices):
119
+ helper = _demo_function("_dp_tp_devices_or_skip", {"pytest": pytest, "ttnn": SimpleNamespace(MeshDevice=object)})
120
+ mesh = SimpleNamespace(get_num_devices=lambda: 8)
121
+
122
+ assert helper(mesh, data_parallel) == expected_tp_devices
123
+
124
+
125
+ def test_t3k_dp2_skips_unsupported_tp4_lanes(expect_error):
126
+ helper = _demo_function("_dp_tp_devices_or_skip", {"pytest": pytest, "ttnn": SimpleNamespace(MeshDevice=object)})
127
+ mesh = SimpleNamespace(get_num_devices=lambda: 8)
128
+
129
+ with expect_error(pytest.skip.Exception, "creates TP4 lanes"):
130
+ helper(mesh, 2)
131
+
132
+
133
+ def test_dp_build_validates_and_resolves_cache_from_each_lane_submesh():
134
+ function = next(
135
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_dp_smoke"
136
+ )
137
+ lane_loop = next(
138
+ node
139
+ for node in ast.walk(function)
140
+ if isinstance(node, ast.For) and isinstance(node.target, ast.Name) and node.target.id == "sm"
141
+ )
142
+ calls = [node for node in ast.walk(lane_loop) if isinstance(node, ast.Call) and isinstance(node.func, ast.Name)]
143
+ call_names = [node.func.id for node in calls]
144
+ assert "_skip_unless_heads_divide_mesh" in call_names
145
+ assert "lazy_weight_cache_dir_for_demo" in call_names
146
+
147
+ from_pretrained_call = next(node for node in calls if node.func.id == "from_pretrained")
148
+ cache_dir = next(keyword.value for keyword in from_pretrained_call.keywords if keyword.arg == "cache_dir")
149
+ assert isinstance(cache_dir, ast.Name)
150
+ assert cache_dir.id == "lane_cache_dir"
151
+
152
+
153
+ @pytest.mark.parametrize("function_name", ["_run_perf_benchmark", "_run_dp_smoke"])
154
+ def test_traced_demo_paths_warm_up_before_benchmark(function_name):
155
+ calls = _called_names(function_name)
156
+ assert calls.index("_warmup_demo_executor") < calls.index("run_perf_benchmark")
157
+
158
+
159
+ def test_eval_repeat_warms_each_fresh_executor():
160
+ calls = _called_names("_run_eval_repeat_batch32")
161
+ assert "_warmup_demo_executor" in calls
162
+
163
+
164
+ def test_perf_path_enables_pipeline_readback_by_default():
165
+ function = next(
166
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_perf_benchmark"
167
+ )
168
+ benchmark_call = next(
169
+ node
170
+ for node in ast.walk(function)
171
+ if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "run_perf_benchmark"
172
+ )
173
+ keywords = {keyword.arg: keyword.value for keyword in benchmark_call.keywords}
174
+ assert isinstance(keywords["pipeline_readback"], ast.Name)
175
+ assert keywords["pipeline_readback"].id == "pipeline_readback"
176
+
177
+
178
+ def test_create_model_preserves_reduced_layer_diagnostic_override(monkeypatch):
179
+ captured = {}
180
+ model = SimpleNamespace()
181
+
182
+ def from_pretrained(*args, **kwargs):
183
+ captured.update(kwargs)
184
+ return SimpleNamespace(model=model, tokenizer=object())
185
+
186
+ namespace = {
187
+ "Path": Path,
188
+ "Llama32_3BTransformer1D": object,
189
+ "LLAMA32_3B_ACCURACY": object(),
190
+ "LLAMA32_3B_PERFORMANCE": object(),
191
+ "_skip_unless_heads_divide_mesh": lambda *_: None,
192
+ "from_pretrained": from_pretrained,
193
+ "os": os,
194
+ "pytest": pytest,
195
+ "ttnn": SimpleNamespace(MeshDevice=object),
196
+ }
197
+ create_model = _demo_function("create_model", namespace)
198
+ monkeypatch.setenv("LLAMA32_3B_DEMO_NUM_LAYERS", "3")
199
+
200
+ assert create_model(object(), "performance", Path("cache")) is model
201
+ assert captured["n_layers"] == 3
202
+
203
+
204
+ def test_token_accuracy_cleans_up_executor_in_finally():
205
+ function = next(
206
+ node for node in _DEMO_TREE.body if isinstance(node, ast.FunctionDef) and node.name == "_run_token_accuracy"
207
+ )
208
+ cleanup_finally = [
209
+ statement
210
+ for node in ast.walk(function)
211
+ if isinstance(node, ast.Try)
212
+ for statement in node.finalbody
213
+ if isinstance(statement, ast.Expr)
214
+ and isinstance(statement.value, ast.Call)
215
+ and isinstance(statement.value.func, ast.Attribute)
216
+ and statement.value.func.attr == "cleanup"
217
+ ]
218
+ assert len(cleanup_finally) == 1