Add files using upload-large-folder tool
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- README.md +113 -0
- code/models/common/modules/moe/configs/deepseek_ocr.yaml +28 -0
- code/models/common/modules/moe/configs/deepseek_v3.yaml +28 -0
- code/models/common/modules/moe/configs/deepseek_v3_single_glx.yaml +28 -0
- code/models/common/modules/moe/configs/deepseek_v4_flash.yaml +32 -0
- code/models/common/modules/moe/configs/deepseek_v4_pro.yaml +32 -0
- code/models/common/modules/moe/configs/gemma_4_26b.yaml +35 -0
- code/models/common/modules/moe/configs/glm5.yaml +30 -0
- code/models/common/modules/moe/configs/glm_47.yaml +31 -0
- code/models/common/modules/moe/configs/gpt_oss.yaml +33 -0
- code/models/common/modules/moe/configs/kimi_k25.yaml +30 -0
- code/models/common/modules/moe/configs/ling_1t.yaml +29 -0
- code/models/common/modules/moe/configs/models_table.md +18 -0
- code/models/common/modules/moe/configs/qwen35_35b.yaml +26 -0
- code/models/common/modules/moe/configs/qwen35_397b.yaml +31 -0
- code/models/common/modules/moe/configs/qwen3_235b.yaml +27 -0
- code/models/common/modules/moe/configs/qwen3_omni_talker.yaml +29 -0
- code/models/common/modules/moe/configs/qwen3_omni_thinker.yaml +27 -0
- code/models/common/tests/demos/cleanup_utils.py +113 -0
- code/models/common/tests/demos/run_helpers.py +1071 -0
- code/models/common/tests/demos/test_cleanup_utils.py +154 -0
- code/models/common/tests/demos/test_run_helpers.py +687 -0
- code/models/common/tests/host/test_metrics_pytorch_only.py +61 -0
- code/models/common/tests/host/test_utility_functions_imports.py +31 -0
- code/models/common/tests/llm_runtime/test_config.py +147 -0
- code/models/common/tests/llm_runtime/test_decode_runtime.py +1337 -0
- code/models/common/tests/llm_runtime/test_execution.py +831 -0
- code/models/common/tests/llm_runtime/test_executor_integration.py +1833 -0
- code/models/common/tests/llm_runtime/test_lane_group.py +1373 -0
- code/models/common/tests/llm_runtime/test_llama3_8b_integration.py +1065 -0
- code/models/common/tests/llm_runtime/test_llama3_8b_model_contract.py +349 -0
- code/models/common/tests/llm_runtime/test_model_contract.py +972 -0
- code/models/common/tests/llm_runtime/test_model_executor.py +291 -0
- code/models/common/tests/llm_runtime/test_output_reader.py +221 -0
- code/models/common/tests/llm_runtime/test_paged_kv_cache.py +415 -0
- code/models/common/tests/llm_runtime/test_prefill_inputs.py +173 -0
- code/models/common/tests/llm_runtime/test_prefill_runtime.py +0 -0
- code/models/common/tests/llm_runtime/test_program_compiler.py +268 -0
- code/models/common/tests/llm_runtime/test_tensor_resources.py +62 -0
- code/models/common/tests/llm_runtime/test_trace_compiler.py +576 -0
- code/models/common/tests/llm_runtime/test_vllm_adapter.py +1025 -0
- code/models/common/tests/llm_runtime/test_warmup.py +1012 -0
- code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_demo_contract.py +452 -0
- code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_hf_adaptor.py +290 -0
- code/models/common/tests/models/deepseek_r1_distill_qwen_14b/test_prefill_last_token_contract.py +71 -0
- code/models/common/tests/models/llama32_1b/test_batched_prefill_postprocess.py +104 -0
- code/models/common/tests/models/llama32_1b/test_demo_warmup.py +134 -0
- code/models/common/tests/models/llama32_1b/test_hf_adaptor.py +264 -0
- code/models/common/tests/models/llama32_3b/test_batched_prefill_postprocess.py +104 -0
- 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
|