lovesenko commited on
Commit
fb22737
·
verified ·
1 Parent(s): 5b3c719

Add files using upload-large-folder tool

Browse files
Files changed (50) hide show
  1. LICENSE +21 -0
  2. README.md +147 -0
  3. config.json +71 -0
  4. encoding/README.md +156 -0
  5. encoding/encoding_dsv4.py +744 -0
  6. encoding/test_encoding_dsv4.py +89 -0
  7. encoding/tests/test_input_1.json +81 -0
  8. encoding/tests/test_input_2.json +24 -0
  9. encoding/tests/test_input_3.json +159 -0
  10. encoding/tests/test_input_4.json +28 -0
  11. encoding/tests/test_output_1.txt +36 -0
  12. encoding/tests/test_output_2.txt +1 -0
  13. encoding/tests/test_output_3.txt +38 -0
  14. encoding/tests/test_output_4.txt +29 -0
  15. generation_config.json +9 -0
  16. inference/README.md +26 -0
  17. inference/config.json +40 -0
  18. inference/convert.py +154 -0
  19. inference/generate.py +144 -0
  20. inference/kernel.py +536 -0
  21. inference/model.py +961 -0
  22. inference/requirements.txt +5 -0
  23. model-00004-of-00048.safetensors +3 -0
  24. model-00006-of-00048.safetensors +3 -0
  25. model-00008-of-00048.safetensors +3 -0
  26. model-00009-of-00048.safetensors +3 -0
  27. model-00010-of-00048.safetensors +3 -0
  28. model-00013-of-00048.safetensors +3 -0
  29. model-00014-of-00048.safetensors +3 -0
  30. model-00017-of-00048.safetensors +3 -0
  31. model-00021-of-00048.safetensors +3 -0
  32. model-00022-of-00048.safetensors +3 -0
  33. model-00023-of-00048.safetensors +3 -0
  34. model-00024-of-00048.safetensors +3 -0
  35. model-00026-of-00048.safetensors +3 -0
  36. model-00027-of-00048.safetensors +3 -0
  37. model-00028-of-00048.safetensors +3 -0
  38. model-00030-of-00048.safetensors +3 -0
  39. model-00031-of-00048.safetensors +3 -0
  40. model-00033-of-00048.safetensors +3 -0
  41. model-00034-of-00048.safetensors +3 -0
  42. model-00035-of-00048.safetensors +3 -0
  43. model-00036-of-00048.safetensors +3 -0
  44. model-00037-of-00048.safetensors +3 -0
  45. model-00040-of-00048.safetensors +3 -0
  46. model-00041-of-00048.safetensors +3 -0
  47. model-00047-of-00048.safetensors +3 -0
  48. model.safetensors.index.json +0 -0
  49. tokenizer.json +0 -0
  50. tokenizer_config.json +34 -0
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2023 DeepSeek
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
README.md ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ library_name: vllm
4
+ base_model: deepseek-ai/DeepSeek-V4-Flash-DSpark
5
+ tags:
6
+ - abliterated
7
+ - uncensored
8
+ - weight-editing
9
+ - deepseek-v4
10
+ ---
11
+
12
+ # DeepSeek-V4-Flash-DSpark — Abliterated
13
+
14
+ This is an abliterated (uncensored) version of `deepseek-ai/DeepSeek-V4-Flash-DSpark`, produced by direct weight-space editing.
15
+
16
+ DeepSeek-V4-Flash is the 284B-parameter (13B-activated) Mixture-of-Experts member of the DeepSeek-V4 family, with a 1-million-token context window and FP8 mixed-precision weights. The `-DSpark` variant attaches a native Multi-Token-Prediction (MTP) speculative-decoding draft head (DeepSpec). Its decoder uses **Manifold-Constrained Hyper-Connections (mHC)**, which — like Gemma 4's double-norm + Per-Layer-Embeddings — make the model highly resistant to LoRA-based abliteration: the mHC residual pathway re-normalizes away low-rank perturbations, so LoRA edits produce near-zero behavioral change. This release bypasses that resistance by editing the base FP8 weights directly, in the 4096-dimensional `wo_b` output space, while preserving row magnitudes and capability.
17
+
18
+ ## Method
19
+
20
+ Because mHC re-normalizes low-rank perturbations, LoRA-based abliteration does not work on this family. The fix is to edit the base weights directly.
21
+
22
+ The abliteration captures **4096-dimensional refusal directions** in the model's own output spaces and projects them out of the attention output projection (`attn.wo_b`) on every decoder layer.
23
+
24
+ Key techniques applied:
25
+
26
+ - **4096-dim refusal-direction capture** via a patched vLLM server that hooks the `wo_b` and aggregated-FFN outputs on all **43 decoder layers**, prefill-only, with per-request sequencing. Five refusal modes were characterized (broad, stubborn, reframe, lecture, value-flip) as difference-of-means directions, Gram-Schmidt orthonormalized, with per-category AUC gating (Harassment AUC ≥ 0.7).
27
+ - **Rank-1 broad-d projection** — only the single broad refusal direction `d` is projected out (higher-rank variants including deflection modes severely damaged capability). This is the smallest, most capability-preserving edit.
28
+ - **SRA cleaning (Spectral Residual Alignment)** — the refusal direction is orthogonalized against the top-`r=8` SVD atoms of capability-concept activations before projection, so the broad-d direction does not eat capability.
29
+ - **Naive output-side orthogonal projection** on `attn.wo_b` for all 43 decoder layers, plus `mtp.wo_b` (the DSpark draft head) via the deepest-layer basis: `W ← W − λ·V(VᵀW)` with **λ_attn = 2.5**.
30
+ - **MLP (`ffn`/`w2`) editing was evaluated and abandoned** — full-43 MLP editing caused catastrophic capability loss and, counter-intuitively, raised refusal on some trials.
31
+ - **FP8/Int8 mixed-precision dequant/requant** — directions are mapped into the weight space and applied with precise dequantization/requantization, since the model ships in FP8-mixed format.
32
+ - **Base-model integrity** — edited shards are written atomically so the original checkpoint is never modified in place; the base model remains byte-intact.
33
+ - **Capability lock** — any variant whose full-benchmark capability dropped >3pt vs base on MMLU-Pro / GSM8K / HumanEval was rejected.
34
+
35
+ ## Evaluation
36
+
37
+ | Metric | Value |
38
+ |---|---|
39
+ | Refusals — broad eval set (1000 prompts, Gemini-2.5-flash judge) | **74 / 1000 (7.40%)** |
40
+ | Refusals — standard harm-benchmark distribution (sealed 284) | 9 / 284 (3.17%) |
41
+ | Baseline refusals (raw base model, sealed 284) | 265 / 284 (93.31%) |
42
+ | Configuration | rank-1 broad-d, all 43 layers, `attn.wo_b` only, λ_attn = 2.5 |
43
+ | Projection mode | Direct weight editing (naive output-side orthogonal projection) |
44
+ | SRA cleaning rank | 8 (vs capability concept atoms) |
45
+ | Edit footprint | 46 `wo_b` tensors, mean Frobenius δ = 0.059 |
46
+ | Hardware used | 2× RTX PRO 6000 Blackwell (~96 GB each, TP=2) |
47
+
48
+ The broad-set figure (7.40%) exceeds the ≤5% target; this is expected and understood rather than a miss — it is driven entirely by the two retained safety-critical categories (PII doxing and self-harm, broken out below), which weight-space `wo_b` abliteration does not remove. On the standard harm-benchmark distribution the ≤5% target is met (3.17%).
49
+
50
+ `λ_attn` was chosen by ablation: λ=3 drove refusal lower (3.87%) but pushed MMLU-Pro to −3.25 pt, just over the capability lock; λ=2.5 is the lowest value that holds capability within the 3-point lock while still clearing the refusal target on the standard harm-benchmark distribution.
51
+
52
+ ### Refusal breakdown by category (1000-prompt broad set)
53
+
54
+ | Category | Refusals | Rate |
55
+ |---|---|---|
56
+ | PII (doxing private individuals) | 36 / 78 | 46.2% |
57
+ | Self-Harm (suicide methods) | 11 / 70 | 15.7% |
58
+ | Radicalization | 5 / 67 | 7.5% |
59
+ | Cybercrime | 4 / 67 | 6.0% |
60
+ | Hate Speech | 4 / 78 | 5.1% |
61
+ | Illegal Drugs | 3 / 69 | 4.3% |
62
+ | Weapons | 2 / 72 | 2.8% |
63
+ | Political Sensitivity | 2 / 69 | 2.9% |
64
+ | Fraud | 2 / 78 | 2.6% |
65
+ | Harassment | 2 / 61 | 3.3% |
66
+ | Violence | 1 / 76 | 1.3% |
67
+ | CBRNE | 1 / 70 | 1.4% |
68
+ | Financial Crimes | 1 / 71 | 1.4% |
69
+ | Sabotage | 0 / 74 | 0.0% |
70
+
71
+ 12 of 14 categories sit at ≤7.5% (most at ≤5%). The broad-set number is carried almost entirely by two safety-critical categories — **PII doxing of private individuals** and **self-harm / suicide methods** — where the model retains guardrails that weight-space `wo_b` abliteration does not remove: for suicide prompts it responds with empathy and crisis-hotline referrals rather than lethal methods, and for doxing it offers public-record guidance but will not surface private contact data. Excluding those two categories, refusal on the remaining 852 prompts is **27 / 852 (3.17%)**. On the standard harm-benchmark distribution (sealed 284, which does not emphasize those two topics) the rate is also 3.17%.
72
+
73
+ ### Full capability sweep (entire datasets, not subsets) — base vs abliterated
74
+
75
+ | Benchmark | Full N | Base | Abliterated | Δ |
76
+ |---|---|---|---|---|
77
+ | MMLU-Pro | 12032 | 0.6733 | 0.6750 | +0.17 pt |
78
+ | GSM8K | 1319 | 0.9242 | 0.9257 | +0.15 pt |
79
+ | HumanEval (pass@1) | 164 | 0.7988 | 0.8354 | +3.66 pt |
80
+ | MBPP (pass@1) | 500 | 0.5180 | 0.5160 | −0.20 pt |
81
+
82
+ ### Multi-turn & higher-context degradation
83
+
84
+ | Benchmark | Base | Abliterated |
85
+ |---|---|---|
86
+ | Multi-turn (20 curated 3-turn convos / 60 turns, Gemini judge 1–10) | 9.97 / 10 | 9.98 / 10 |
87
+ | Needle-in-haystack @ 2k / 4k / 8k / 16k / 32k tokens | 100% / 100% / 100% / 100% / 100% | 100% / 100% / 100% / 100% / 100% |
88
+
89
+ No multi-turn coherence loss and **no higher-context degradation up to 32k tokens**.
90
+
91
+ ### SWE-bench Lite (oracle-file-context, single-shot, n=30, same instances)
92
+
93
+ Terminology: **submitted** = instances attempted; **completed** = the model's patch applied cleanly and the test suite ran (i.e. a valid pass/fail verdict was reached); **resolved** ⊂ completed = the previously-failing tests now pass; **patch-apply errors** = the generated diff did not apply cleanly, so no verdict was produced.
94
+
95
+ | | Base | Abliterated |
96
+ |---|---|---|
97
+ | Submitted | 30 | 30 |
98
+ | Completed | 21 | 19 |
99
+ | **Resolved** | **4 (13.3%)** | **4 (13.3%)** |
100
+ | Unresolved (completed, tests still failing) | 17 | 15 |
101
+ | Patch-apply errors | 9 | 11 |
102
+
103
+ Identical resolve rate (3 of 4 resolved instances overlap) — no degradation in agentic-style code repair. The abliterated model produced 2 more patch-apply errors (malformed diffs) and 2 fewer completions; those 2 instances shifted from "completed-but-unresolved" on base to "patch-apply error" on abliterated, which is within run-to-run variance for single-shot diff generation and does not change the resolved count.
104
+
105
+ ### DSpark speculative decoding (post-abliteration)
106
+
107
+ The `mtp.wo_b` draft head was edited with the same projection applied to the decoder (deepest-layer basis). Speculative decoding remains functional and healthy:
108
+
109
+ - DSpark ON: **174 tok/s** (512 tokens in 2.95s)
110
+ - DSpark OFF: 123 tok/s (512 tokens in 4.15s)
111
+ - **+41% throughput** from accepted drafts; coherence spot-checks (temp=0) correct.
112
+
113
+ **Draft acceptance** (accepted draft tokens / total draft tokens; 100-prompt mixed workload, probabilistic draft sampling):
114
+
115
+ | `num_speculative_tokens` | Base | Abliterated |
116
+ |---|---|---|
117
+ | 5 (native) | 51.1% | 51.3% |
118
+ | 4 | — | 59.0% |
119
+ | 3 | — | 66.0% |
120
+
121
+ At the native 5-token setting, acceptance is **at parity with the unedited base model (51.3% vs 51.1%)** — the weight edit did not desynchronize the draft head from the target, and ~51% is simply DSpark's native acceptance rate at 5 draft tokens on a mixed workload. Acceptance rises with fewer draft tokens (the standard speculative-decoding trade-off: fewer drafts = higher per-token acceptance, marginally lower absolute speedup). The release is served at the native 5-token setting; users wanting a higher acceptance rate can set `num_speculative_tokens=4` (~59%) or `3` (~66%) at a small throughput cost. Because DSpark verifies every draft token against the abliterated target, the served output distribution is identical for any of these settings — they affect throughput only, not outputs.
122
+
123
+ ## A note on honest evaluation
124
+
125
+ Refusal numbers are only meaningful when the methodology behind them is documented. Our methodology:
126
+
127
+ - **Sufficient generation length** — DeepSeek-V4 exhibits a "delayed refusal" pattern (a stretch of educational framing/disclaimers before pivoting to the actual refusal), so short generations undercount refusals; we generate long enough to capture that pivot.
128
+ - **Hybrid detection** — keyword matching for obvious refusals plus an LLM judge (Google Gemini 2.5 Flash via OpenRouter) for ambiguous cases. Neither method alone is sufficient.
129
+ - **Challenging, diverse prompts** — our refusal set spans 1000 prompts across 14 categories, multiple sophistication levels (direct requests to socially-engineered framings), and English / Chinese / mixed languages.
130
+ - **Paired, full-dataset capability measurement** — capability is measured on the **entire** MMLU-Pro (12032), GSM8K (1319), HumanEval (164) and MBPP (500) test sets for both base and abliterated models, not small samples.
131
+ - **Documented parameters** — generation length, detection method, dataset, λ, rank, and layer coverage are all listed on this card.
132
+
133
+ The refusal figures above are from a rigorous end-to-end re-evaluation of the edited weights, including the category breakdown so the two safety-critical categories that retain guardrails are visible rather than averaged away.
134
+
135
+ ## Usage
136
+
137
+ This abliterated checkpoint is a drop-in replacement for the original weights — it has the **exact same architecture, format, chat-template/encoding, and inference path as the released base model** [`deepseek-ai/DeepSeek-V4-Flash-DSpark`](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash-DSpark). Load and serve it however you would the official model (vLLM, the DeepSeek-V4 `encoding`/`inference` folders, OpenAI-compatible serving, etc.). The abliteration modified the text-decoder `attn.wo_b` weights on all 43 layers and the DSpark draft head's `mtp.wo_b`; the tokenizer, chat encoding, and all other components are unchanged.
138
+
139
+ For inference guidance specific to the **NVIDIA RTX PRO 6000 Blackwell** (TP2/TP4, the `lucifer-default` / `lucifer-cutlass` / `b12x` backends, and the native DSpark `method=dspark` speculative-decoding path with `num_speculative_tokens=5`), see the community guide:
140
+
141
+ 👉 **https://github.com/local-inference-lab/rtx6kpro/blob/master/models/ds4dspark-v8.md**
142
+
143
+ That page documents the validated Docker image, launch helpers, and full TP2/TP4 throughput sweep (decode + prefill) for this exact checkpoint on RTX PRO 6000, including DSpark speculative decoding which gives a ~41% single-stream throughput speedup on this abliterated release.
144
+
145
+ ## Disclaimer
146
+
147
+ This model is released for research purposes only — primarily interpretability and safety research, including studying how refusal behavior is encoded in large MoE decoders and how weight-space edits interact with architectures that resist low-rank perturbation. The abliteration process removes safety guardrails on most harm categories, so the model will comply with requests the base model refuses. Use responsibly, in accordance with local laws and the DeepSeek / model terms of use, and do not deploy it in production or user-facing settings without a separate safety layer. The authors take no responsibility for misuse.
config.json ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "DeepseekV4ForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "bos_token_id": 0,
8
+ "eos_token_id": 1,
9
+ "expert_dtype": "fp4",
10
+ "hc_eps": 1e-06,
11
+ "hc_mult": 4,
12
+ "hc_sinkhorn_iters": 20,
13
+ "head_dim": 512,
14
+ "hidden_act": "silu",
15
+ "hidden_size": 4096,
16
+ "index_head_dim": 128,
17
+ "index_n_heads": 64,
18
+ "index_topk": 512,
19
+ "initializer_range": 0.02,
20
+ "max_position_embeddings": 1048576,
21
+ "model_type": "deepseek_v4",
22
+ "moe_intermediate_size": 2048,
23
+ "n_routed_experts": 256,
24
+ "n_shared_experts": 1,
25
+ "norm_topk_prob": true,
26
+ "num_attention_heads": 64,
27
+ "num_experts_per_tok": 6,
28
+ "num_hidden_layers": 43,
29
+ "num_hash_layers": 3,
30
+ "num_key_value_heads": 1,
31
+ "num_nextn_predict_layers": 1,
32
+ "o_groups": 8,
33
+ "o_lora_rank": 1024,
34
+ "q_lora_rank": 1024,
35
+ "qk_rope_head_dim": 64,
36
+ "quantization_config": {
37
+ "activation_scheme": "dynamic",
38
+ "fmt": "e4m3",
39
+ "quant_method": "fp8",
40
+ "scale_fmt": "ue8m0",
41
+ "weight_block_size": [
42
+ 128,
43
+ 128
44
+ ]
45
+ },
46
+ "rms_norm_eps": 1e-06,
47
+ "rope_scaling": {
48
+ "beta_fast": 32,
49
+ "beta_slow": 1,
50
+ "factor": 16,
51
+ "original_max_position_embeddings": 65536,
52
+ "type": "yarn"
53
+ },
54
+ "rope_theta": 10000,
55
+ "routed_scaling_factor": 1.5,
56
+ "scoring_func": "sqrtsoftplus",
57
+ "sliding_window": 128,
58
+ "swiglu_limit": 10.0,
59
+ "tie_word_embeddings": false,
60
+ "topk_method": "noaux_tc",
61
+ "torch_dtype": "bfloat16",
62
+ "transformers_version": "4.57.1",
63
+ "use_cache": true,
64
+ "vocab_size": 129280,
65
+ "compress_rope_theta": 160000,
66
+ "compress_ratios": [0, 0, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 0, 0, 0],
67
+ "dspark_block_size": 5,
68
+ "dspark_noise_token_id": 128799,
69
+ "dspark_target_layer_ids": [40, 41, 42],
70
+ "dspark_markov_rank": 256
71
+ }
encoding/README.md ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # DeepSeek-V4 Encoding
2
+
3
+ This document describes the prompt encoding format used by DeepSeek-V4 series models. The encoding handles multi-turn conversations, tool calling, extended thinking (reasoning), and quick instruction tasks.
4
+
5
+ A self-contained reference implementation is provided in `encoding_dsv4.py`.
6
+
7
+ ## Quick Start
8
+
9
+ ```python
10
+ from encoding_dsv4 import encode_messages, parse_message_from_completion_text
11
+
12
+ # Encode a conversation
13
+ messages = [
14
+ {"role": "system", "content": "You are a helpful assistant."},
15
+ {"role": "user", "content": "What is 2+2?"},
16
+ ]
17
+ prompt = encode_messages(messages, thinking_mode="thinking")
18
+ # => "<|begin▁of▁sentence|>You are a helpful assistant.<|User|>What is 2+2?<|Assistant|><think>"
19
+
20
+ # Parse model output back to structured message
21
+ completion = "Simple arithmetic.</think>2 + 2 = 4.<|end▁of▁sentence|>"
22
+ parsed = parse_message_from_completion_text(completion, thinking_mode="thinking")
23
+ # => {"role": "assistant", "reasoning_content": "Simple arithmetic.", "content": "2 + 2 = 4.", "tool_calls": []}
24
+ ```
25
+
26
+ > **Note:** The `parse_message_from_completion_text` function is designed to handle well-formatted model output only. It does not attempt to correct or recover from malformed output that the model might occasionally generate. For production use, additional error handling is recommended.
27
+
28
+ ## Message Format
29
+
30
+ ### Special Tokens
31
+
32
+ | Token | Purpose |
33
+ |-------|---------|
34
+ | `<|begin▁of▁sentence|>` | Beginning of sequence (BOS) |
35
+ | `<|end▁of▁sentence|>` | End of assistant turn (EOS) |
36
+ | `<|User|>` | User turn prefix |
37
+ | `<|Assistant|>` | Assistant turn prefix |
38
+ | `<|latest_reminder|>` | Latest reminder (date, locale, etc.) |
39
+ | `<think>` / `</think>` | Reasoning block delimiters |
40
+ | `|DSML|` | DSML markup token |
41
+
42
+ ### Roles
43
+
44
+ The encoding supports the following message roles: `system`, `user`, `assistant`, `tool`, `latest_reminder`, and `developer`.
45
+
46
+ > **Note on the `developer` role:** The `developer` role is used exclusively in the internal search agent pipeline. It is not needed for general-purpose chat or tool-calling tasks, and the official API does not accept messages with this role.
47
+
48
+ ### Basic Chat
49
+
50
+ A simple multi-turn conversation is encoded as:
51
+
52
+ ```
53
+ <|begin▁of▁sentence|>{system_prompt}
54
+ <|User|>{user_message}<|Assistant|></think>{response}<|end▁of▁sentence|>
55
+ <|User|>{user_message_2}<|Assistant|></think>{response_2}<|end▁of▁sentence|>
56
+ ```
57
+
58
+ - The BOS token is prepended at the very beginning of the conversation.
59
+ - In **chat mode** (`thinking_mode="chat"`), `</think>` is placed right after `<|Assistant|>` to immediately close the thinking block, so the model generates content directly.
60
+
61
+ ### Interleaved Thinking Mode
62
+
63
+ In **thinking mode** (`thinking_mode="thinking"`), the model produces explicit reasoning inside `<think>...</think>` blocks before responding.
64
+
65
+ ```
66
+ <|begin▁of▁sentence|>{system_prompt}
67
+ <|User|>{message}<|Assistant|><think>{reasoning}</think>{response}<|end▁of▁sentence|>
68
+ ```
69
+
70
+ The `drop_thinking` parameter (default `True`) controls whether reasoning from earlier turns is preserved:
71
+
72
+ - **Without tools**: `drop_thinking` takes effect. Reasoning content from assistant turns **before** the last user message is stripped. Only the final assistant turn retains its `<think>...</think>` block.
73
+ - **With tools** (on system or developer message): `drop_thinking` is automatically disabled. All turns retain their reasoning, because tool-calling conversations require full context for the model to track multi-step reasoning across tool calls.
74
+
75
+ ### Tool Calling (DSML Format)
76
+
77
+ Tools are defined on the `system` or `developer` message via the `tools` field (OpenAI-compatible format). When tools are present, the following schema block is injected into the system/user prompt:
78
+
79
+ ```
80
+ ## Tools
81
+
82
+ You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<|DSML|tool_calls>" block like the following:
83
+
84
+ <|DSML|tool_calls>
85
+ <|DSML|invoke name="$TOOL_NAME">
86
+ <|DSML|parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE</|DSML|parameter>
87
+ ...
88
+ </|DSML|invoke>
89
+ <|DSML|invoke name="$TOOL_NAME2">
90
+ ...
91
+ </|DSML|invoke>
92
+ </|DSML|tool_calls>
93
+
94
+ String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.
95
+
96
+ If thinking_mode is enabled (triggered by <think>), you MUST output your complete reasoning inside <think>...</think> BEFORE any tool calls or final response.
97
+
98
+ Otherwise, output directly after </think> with tool calls or final response.
99
+
100
+ ### Available Tool Schemas
101
+
102
+ {tool_definitions_json}
103
+
104
+ You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.
105
+ ```
106
+
107
+ An actual tool call in the assistant turn looks like:
108
+
109
+ ```xml
110
+ <|DSML|tool_calls>
111
+ <|DSML|invoke name="function_name">
112
+ <|DSML|parameter name="param" string="true">string_value</|DSML|parameter>
113
+ <|DSML|parameter name="count" string="false">5</|DSML|parameter>
114
+ </|DSML|invoke>
115
+ </|DSML|tool_calls><|end▁of▁sentence|>
116
+ ```
117
+
118
+ - `string="true"`: the parameter value is a raw string.
119
+ - `string="false"`: the parameter value is JSON (number, boolean, array, object).
120
+
121
+ Tool execution results are wrapped in `<tool_result>` tags within user messages:
122
+
123
+ ```
124
+ <|User|><tool_result>{result_json}</tool_result><|Assistant|><think>...
125
+ ```
126
+
127
+ When multiple tool results are present, they are sorted by the order of the corresponding `tool_calls` in the preceding assistant message.
128
+
129
+ ### Reasoning Effort
130
+
131
+ When `reasoning_effort="max"` is set, a special prefix is prepended at the very beginning of the prompt (before the system message) to instruct the model to maximize its reasoning depth:
132
+
133
+ ```
134
+ Reasoning Effort: Absolute maximum with no shortcuts permitted.
135
+ You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.
136
+ Explicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.
137
+ ```
138
+
139
+ ### Quick Instruction Special Tokens
140
+
141
+ Quick instruction tokens are used for auxiliary classification and generation tasks. They are appended to messages via the `"task"` field to trigger specialized model behavior for a single-token or short-form output.
142
+
143
+ | Special Token | Description | Format |
144
+ |:---|:---|:---|
145
+ | `<|action|>` | Determines whether the user prompt requires a web search or can be answered directly. | `...<|User|>{prompt}<|Assistant|><think><|action|>` |
146
+ | `<|title|>` | Generates a concise conversation title after the first assistant response. | `...<|Assistant|>{response}<|end▁of▁sentence|><|title|>` |
147
+ | `<|query|>` | Generates search queries for the user prompt. | `...<|User|>{prompt}<|query|>` |
148
+ | `<|authority|>` | Classifies the user prompt's demand for source authoritativeness. | `...<|User|>{prompt}<|authority|>` |
149
+ | `<|domain|>` | Identifies the domain of the user prompt. | `...<|User|>{prompt}<|domain|>` |
150
+ | `<|extracted_url|>` `<|read_url|>` | Determines whether each URL in the user prompt should be fetched and read. | `...<|User|>{prompt}<|extracted_url|>{url}<|read_url|>` |
151
+
152
+ Usage in message format:
153
+
154
+ - **`action`** on a user message: the `<|action|>` token is placed after the assistant prefix and thinking token, triggering a routing decision (e.g., "Search" or "Answer").
155
+ - **Other tasks** (`query`, `authority`, `domain`, `read_url`) on a user message: the task token is appended directly after the user content.
156
+ - **`title`** on an assistant message: the `<|title|>` token is appended after the assistant's EOS. The next assistant message provides the generated title.
encoding/encoding_dsv4.py ADDED
@@ -0,0 +1,744 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ DeepSeek-V4 Encoding
3
+
4
+ A self-contained implementation for encoding/decoding DeepSeek-V4 chat messages
5
+ with tool calling, thinking mode, and quick instruction task support.
6
+ """
7
+
8
+ from typing import Any, Dict, List, Union, Optional, Tuple
9
+ import copy
10
+ import json
11
+ import re
12
+
13
+ # ============================================================
14
+ # Special Tokens
15
+ # ============================================================
16
+
17
+ bos_token: str = "<|begin▁of▁sentence|>"
18
+ eos_token: str = "<|end▁of▁sentence|>"
19
+ thinking_start_token: str = "<think>"
20
+ thinking_end_token: str = "</think>"
21
+ dsml_token: str = "|DSML|"
22
+
23
+ USER_SP_TOKEN = "<|User|>"
24
+ ASSISTANT_SP_TOKEN = "<|Assistant|>"
25
+ LATEST_REMINDER_SP_TOKEN = "<|latest_reminder|>"
26
+
27
+ # Task special tokens for internal classification tasks
28
+ DS_TASK_SP_TOKENS = {
29
+ "action": "<|action|>",
30
+ "query": "<|query|>",
31
+ "authority": "<|authority|>",
32
+ "domain": "<|domain|>",
33
+ "title": "<|title|>",
34
+ "read_url": "<|read_url|>",
35
+ }
36
+ VALID_TASKS = set(DS_TASK_SP_TOKENS.keys())
37
+
38
+ # ============================================================
39
+ # Templates
40
+ # ============================================================
41
+
42
+ system_msg_template: str = "{content}"
43
+ user_msg_template: str = "{content}"
44
+ latest_reminder_msg_template: str = "{content}"
45
+ assistant_msg_template: str = "{reasoning}{content}{tool_calls}" + eos_token
46
+ assistant_msg_wo_eos_template: str = "{reasoning}{content}{tool_calls}"
47
+ thinking_template: str = "{reasoning_content}"
48
+
49
+ response_format_template: str = (
50
+ "## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n{schema}"
51
+ )
52
+ tool_call_template: str = (
53
+ "<{dsml_token}invoke name=\"{name}\">\n{arguments}\n</{dsml_token}invoke>"
54
+ )
55
+ tool_calls_template = (
56
+ "<{dsml_token}{tc_block_name}>\n{tool_calls}\n</{dsml_token}{tc_block_name}>"
57
+ )
58
+ tool_calls_block_name: str = "tool_calls"
59
+
60
+ tool_output_template: str = (
61
+ "<tool_result>{content}</tool_result>"
62
+ )
63
+
64
+ REASONING_EFFORT_MAX = (
65
+ "Reasoning Effort: Absolute maximum with no shortcuts permitted.\n"
66
+ "You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\n"
67
+ "Explicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n"
68
+ )
69
+
70
+ TOOLS_TEMPLATE = """## Tools
71
+
72
+ You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<{dsml_token}tool_calls>" block like the following:
73
+
74
+ <{dsml_token}tool_calls>
75
+ <{dsml_token}invoke name="$TOOL_NAME">
76
+ <{dsml_token}parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE</{dsml_token}parameter>
77
+ ...
78
+ </{dsml_token}invoke>
79
+ <{dsml_token}invoke name="$TOOL_NAME2">
80
+ ...
81
+ </{dsml_token}invoke>
82
+ </{dsml_token}tool_calls>
83
+
84
+ String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.
85
+
86
+ If thinking_mode is enabled (triggered by {thinking_start_token}), you MUST output your complete reasoning inside {thinking_start_token}...{thinking_end_token} BEFORE any tool calls or final response.
87
+
88
+ Otherwise, output directly after {thinking_end_token} with tool calls or final response.
89
+
90
+ ### Available Tool Schemas
91
+
92
+ {tool_schemas}
93
+
94
+ You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.
95
+ """
96
+
97
+ # ============================================================
98
+ # Utility Functions
99
+ # ============================================================
100
+
101
+ def to_json(value: Any) -> str:
102
+ """Serialize a value to JSON string."""
103
+ try:
104
+ return json.dumps(value, ensure_ascii=False)
105
+ except:
106
+ return json.dumps(value, ensure_ascii=True)
107
+
108
+
109
+ def tools_from_openai_format(tools):
110
+ """Extract function definitions from OpenAI-format tool list."""
111
+ return [tool["function"] for tool in tools]
112
+
113
+
114
+ def tool_calls_from_openai_format(tool_calls):
115
+ """Convert OpenAI-format tool calls to internal format."""
116
+ return [
117
+ {
118
+ "name": tool_call["function"]["name"],
119
+ "arguments": tool_call["function"]["arguments"],
120
+ }
121
+ for tool_call in tool_calls
122
+ ]
123
+
124
+
125
+ def tool_calls_to_openai_format(tool_calls):
126
+ """Convert internal tool calls to OpenAI format."""
127
+ return [
128
+ {
129
+ "type": "function",
130
+ "function": {
131
+ "name": tool_call["name"],
132
+ "arguments": tool_call["arguments"],
133
+ }
134
+ }
135
+ for tool_call in tool_calls
136
+ ]
137
+
138
+
139
+ def encode_arguments_to_dsml(tool_call: Dict[str, str]) -> str:
140
+ """
141
+ Encode tool call arguments into DSML parameter format.
142
+
143
+ Args:
144
+ tool_call: Dict with "name" and "arguments" (JSON string) keys.
145
+
146
+ Returns:
147
+ DSML-formatted parameter string.
148
+ """
149
+ p_dsml_template = '<{dsml_token}parameter name="{key}" string="{is_str}">{value}</{dsml_token}parameter>'
150
+ P_dsml_strs = []
151
+
152
+ try:
153
+ arguments = json.loads(tool_call["arguments"])
154
+ except Exception as err:
155
+ arguments = {"arguments": tool_call["arguments"]}
156
+
157
+ for k, v in arguments.items():
158
+ p_dsml_str = p_dsml_template.format(
159
+ dsml_token=dsml_token,
160
+ key=k,
161
+ is_str="true" if isinstance(v, str) else "false",
162
+ value=v if isinstance(v, str) else to_json(v),
163
+ )
164
+ P_dsml_strs.append(p_dsml_str)
165
+
166
+ return "\n".join(P_dsml_strs)
167
+
168
+
169
+ def decode_dsml_to_arguments(tool_name: str, tool_args: Dict[str, Tuple[str, str]]) -> Dict[str, str]:
170
+ """
171
+ Decode DSML parameters back to a tool call dict.
172
+
173
+ Args:
174
+ tool_name: Name of the tool.
175
+ tool_args: Dict mapping param_name -> (value, is_string_flag).
176
+
177
+ Returns:
178
+ Dict with "name" and "arguments" (JSON string) keys.
179
+ """
180
+ def _decode_value(key: str, value: str, string: str):
181
+ if string == "true":
182
+ value = to_json(value)
183
+ return f"{to_json(key)}: {value}"
184
+
185
+ tool_args_json = "{" + ", ".join([_decode_value(k, v, string=is_str) for k, (v, is_str) in tool_args.items()]) + "}"
186
+ return dict(name=tool_name, arguments=tool_args_json)
187
+
188
+
189
+ def render_tools(tools: List[Dict[str, Union[str, Dict[str, Any]]]]) -> str:
190
+ """
191
+ Render tool schemas into the system prompt format.
192
+
193
+ Args:
194
+ tools: List of tool schema dicts (each with name, description, parameters).
195
+
196
+ Returns:
197
+ Formatted tools section string.
198
+ """
199
+ tools_json = [to_json(t) for t in tools]
200
+
201
+ return TOOLS_TEMPLATE.format(
202
+ tool_schemas="\n".join(tools_json),
203
+ dsml_token=dsml_token,
204
+ thinking_start_token=thinking_start_token,
205
+ thinking_end_token=thinking_end_token,
206
+ )
207
+
208
+
209
+ def find_last_user_index(messages: List[Dict[str, Any]]) -> int:
210
+ """Find the index of the last user/developer message."""
211
+ last_user_index = -1
212
+ for idx in range(len(messages) - 1, -1, -1):
213
+ if messages[idx].get("role") in ["user", "developer"]:
214
+ last_user_index = idx
215
+ break
216
+ return last_user_index
217
+
218
+
219
+ # ============================================================
220
+ # Message Rendering
221
+ # ============================================================
222
+
223
+ def render_message(index: int, messages: List[Dict[str, Any]], thinking_mode: str, drop_thinking: bool = True, reasoning_effort: Optional[str] = None) -> str:
224
+ """
225
+ Render a single message at the given index into its encoded string form.
226
+
227
+ This is the core function that converts each message in the conversation
228
+ into the DeepSeek-V4 format.
229
+
230
+ Args:
231
+ index: Index of the message to render.
232
+ messages: Full list of messages in the conversation.
233
+ thinking_mode: Either "chat" or "thinking".
234
+ drop_thinking: Whether to drop reasoning content from earlier turns.
235
+ reasoning_effort: Optional reasoning effort level ("max", "high", or None).
236
+
237
+ Returns:
238
+ Encoded string for this message.
239
+ """
240
+ assert 0 <= index < len(messages)
241
+ assert thinking_mode in ["chat", "thinking"], f"Invalid thinking_mode `{thinking_mode}`"
242
+
243
+ prompt = ""
244
+ msg = messages[index]
245
+ last_user_idx = find_last_user_index(messages)
246
+
247
+ role = msg.get("role")
248
+ content = msg.get("content")
249
+ tools = msg.get("tools")
250
+ response_format = msg.get("response_format")
251
+ tool_calls = msg.get("tool_calls")
252
+ reasoning_content = msg.get("reasoning_content")
253
+ wo_eos = msg.get("wo_eos", False)
254
+
255
+ if tools:
256
+ tools = tools_from_openai_format(tools)
257
+ if tool_calls:
258
+ tool_calls = tool_calls_from_openai_format(tool_calls)
259
+
260
+ # Reasoning effort prefix (only at index 0 in thinking mode with max effort)
261
+ assert reasoning_effort in ['max', None, 'high'], f"Invalid reasoning effort: {reasoning_effort}"
262
+ if index == 0 and thinking_mode == "thinking" and reasoning_effort == 'max':
263
+ prompt += REASONING_EFFORT_MAX
264
+
265
+ if role == "system":
266
+ prompt += system_msg_template.format(content=content or "")
267
+ if tools:
268
+ prompt += "\n\n" + render_tools(tools)
269
+ if response_format:
270
+ prompt += "\n\n" + response_format_template.format(schema=to_json(response_format))
271
+
272
+ elif role == "developer":
273
+ assert content, f"Invalid message for role `{role}`: {msg}"
274
+
275
+ content_developer = USER_SP_TOKEN
276
+ content_developer += content
277
+
278
+ if tools:
279
+ content_developer += "\n\n" + render_tools(tools)
280
+ if response_format:
281
+ content_developer += "\n\n" + response_format_template.format(schema=to_json(response_format))
282
+
283
+ prompt += user_msg_template.format(content=content_developer)
284
+
285
+ elif role == "user":
286
+ prompt += USER_SP_TOKEN
287
+
288
+ # Handle content blocks (tool results mixed with text)
289
+ content_blocks = msg.get("content_blocks")
290
+ if content_blocks:
291
+ parts = []
292
+ for block in content_blocks:
293
+ block_type = block.get("type")
294
+ if block_type == "text":
295
+ parts.append(block.get("text", ""))
296
+ elif block_type == "tool_result":
297
+ tool_content = block.get("content", "")
298
+ if isinstance(tool_content, list):
299
+ text_parts = []
300
+ for b in tool_content:
301
+ if b.get("type") == "text":
302
+ text_parts.append(b.get("text", ""))
303
+ else:
304
+ text_parts.append(f"[Unsupported {b.get('type')}]")
305
+ tool_content = "\n\n".join(text_parts)
306
+ parts.append(tool_output_template.format(content=tool_content))
307
+ else:
308
+ parts.append(f"[Unsupported {block_type}]")
309
+ prompt += "\n\n".join(parts)
310
+ else:
311
+ prompt += content or ""
312
+
313
+ elif role == "latest_reminder":
314
+ prompt += LATEST_REMINDER_SP_TOKEN + latest_reminder_msg_template.format(content=content)
315
+
316
+ elif role == "tool":
317
+ raise NotImplementedError("deepseek_v4 merges tool messages into user; please preprocess with merge_tool_messages()")
318
+
319
+ elif role == "assistant":
320
+ thinking_part = ""
321
+ tc_content = ""
322
+
323
+ if tool_calls:
324
+ tc_list = [
325
+ tool_call_template.format(
326
+ dsml_token=dsml_token,
327
+ name=tc.get("name"),
328
+ arguments=encode_arguments_to_dsml(tc)
329
+ )
330
+ for tc in tool_calls
331
+ ]
332
+ tc_content += '\n\n' + tool_calls_template.format(
333
+ dsml_token=dsml_token,
334
+ tool_calls="\n".join(tc_list),
335
+ tc_block_name=tool_calls_block_name,
336
+ )
337
+
338
+ summary_content = content or ""
339
+ rc = reasoning_content or ""
340
+
341
+ # Check if previous message has a task - if so, this is a task output (no thinking)
342
+ prev_has_task = index - 1 >= 0 and messages[index - 1].get("task") is not None
343
+
344
+ if thinking_mode == "thinking" and not prev_has_task:
345
+ if not drop_thinking or index > last_user_idx:
346
+ thinking_part = thinking_template.format(reasoning_content=rc) + thinking_end_token
347
+ else:
348
+ thinking_part = ""
349
+
350
+ if wo_eos:
351
+ prompt += assistant_msg_wo_eos_template.format(
352
+ reasoning=thinking_part,
353
+ content=summary_content,
354
+ tool_calls=tc_content,
355
+ )
356
+ else:
357
+ prompt += assistant_msg_template.format(
358
+ reasoning=thinking_part,
359
+ content=summary_content,
360
+ tool_calls=tc_content,
361
+ )
362
+ else:
363
+ raise NotImplementedError(f"Unknown role: {role}")
364
+
365
+ # Append transition tokens based on what follows
366
+ if index + 1 < len(messages) and messages[index + 1].get("role") not in ["assistant", "latest_reminder"]:
367
+ return prompt
368
+
369
+ task = messages[index].get("task")
370
+ if task is not None:
371
+ # Task special token for internal classification tasks
372
+ assert task in VALID_TASKS, f"Invalid task: '{task}'. Valid tasks are: {list(VALID_TASKS)}"
373
+ task_sp_token = DS_TASK_SP_TOKENS[task]
374
+
375
+ if task != "action":
376
+ # Non-action tasks: append task sp token directly after the message
377
+ prompt += task_sp_token
378
+ else:
379
+ # Action task: append Assistant + thinking token + action sp token
380
+ prompt += ASSISTANT_SP_TOKEN
381
+ prompt += thinking_end_token if thinking_mode != "thinking" else thinking_start_token
382
+ prompt += task_sp_token
383
+
384
+ elif messages[index].get("role") in ["user", "developer"]:
385
+ # Normal generation: append Assistant + thinking token
386
+ prompt += ASSISTANT_SP_TOKEN
387
+ if not drop_thinking and thinking_mode == "thinking":
388
+ prompt += thinking_start_token
389
+ elif drop_thinking and thinking_mode == "thinking" and index >= last_user_idx:
390
+ prompt += thinking_start_token
391
+ else:
392
+ prompt += thinking_end_token
393
+
394
+ return prompt
395
+
396
+
397
+ # ============================================================
398
+ # Preprocessing
399
+ # ============================================================
400
+
401
+ def merge_tool_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
402
+ """
403
+ Merge tool messages into the preceding user message using content_blocks format.
404
+
405
+ DeepSeek-V4 does not have a standalone "tool" role; instead, tool results
406
+ are encoded as <tool_result> blocks within user messages.
407
+
408
+ This function converts a standard OpenAI-format conversation (with separate
409
+ "tool" role messages) into V4 format where tool results are merged into
410
+ user messages.
411
+
412
+ Args:
413
+ messages: List of message dicts in OpenAI format.
414
+
415
+ Returns:
416
+ Processed message list with tool messages merged into user messages.
417
+ """
418
+ merged: List[Dict[str, Any]] = []
419
+
420
+ for msg in messages:
421
+ msg = copy.deepcopy(msg)
422
+ role = msg.get("role")
423
+
424
+ if role == "tool":
425
+ # Convert tool message to a user message with tool_result block
426
+ tool_block = {
427
+ "type": "tool_result",
428
+ "tool_use_id": msg.get("tool_call_id", ""),
429
+ "content": msg.get("content", ""),
430
+ }
431
+ # Merge into previous message if it's already a user (merged tool)
432
+ if merged and merged[-1].get("role") == "user" and "content_blocks" in merged[-1]:
433
+ merged[-1]["content_blocks"].append(tool_block)
434
+ else:
435
+ merged.append({
436
+ "role": "user",
437
+ "content_blocks": [tool_block],
438
+ })
439
+ elif role == "user":
440
+ text_block = {"type": "text", "text": msg.get("content", "")}
441
+ if merged and merged[-1].get("role") == "user" and "content_blocks" in merged[-1] and merged[-1].get("task") is None:
442
+ merged[-1]["content_blocks"].append(text_block)
443
+ else:
444
+ new_msg = {
445
+ "role": "user",
446
+ "content": msg.get("content", ""),
447
+ "content_blocks": [text_block],
448
+ }
449
+ # Preserve extra fields (task, wo_eos, mask, etc.)
450
+ for key in ("task", "wo_eos", "mask"):
451
+ if key in msg:
452
+ new_msg[key] = msg[key]
453
+ merged.append(new_msg)
454
+ else:
455
+ merged.append(msg)
456
+
457
+ return merged
458
+
459
+
460
+ def sort_tool_results_by_call_order(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
461
+ """
462
+ Sort tool_result blocks within user messages by the order of tool_calls
463
+ in the preceding assistant message.
464
+
465
+ Args:
466
+ messages: Preprocessed message list (after merge_tool_messages).
467
+
468
+ Returns:
469
+ Message list with sorted tool result blocks.
470
+ """
471
+ last_tool_call_order: Dict[str, int] = {}
472
+
473
+ for msg in messages:
474
+ role = msg.get("role")
475
+ if role == "assistant" and msg.get("tool_calls"):
476
+ last_tool_call_order = {}
477
+ for idx, tc in enumerate(msg["tool_calls"]):
478
+ tc_id = tc.get("id") or tc.get("function", {}).get("id", "")
479
+ if tc_id:
480
+ last_tool_call_order[tc_id] = idx
481
+
482
+ elif role == "user" and msg.get("content_blocks"):
483
+ tool_blocks = [b for b in msg["content_blocks"] if b.get("type") == "tool_result"]
484
+ if len(tool_blocks) > 1 and last_tool_call_order:
485
+ sorted_blocks = sorted(
486
+ tool_blocks,
487
+ key=lambda b: last_tool_call_order.get(b.get("tool_use_id", ""), 0)
488
+ )
489
+ sorted_idx = 0
490
+ new_blocks = []
491
+ for block in msg["content_blocks"]:
492
+ if block.get("type") == "tool_result":
493
+ new_blocks.append(sorted_blocks[sorted_idx])
494
+ sorted_idx += 1
495
+ else:
496
+ new_blocks.append(block)
497
+ msg["content_blocks"] = new_blocks
498
+
499
+ return messages
500
+
501
+
502
+ # ============================================================
503
+ # Main Encoding Function
504
+ # ============================================================
505
+
506
+ def encode_messages(
507
+ messages: List[Dict[str, Any]],
508
+ thinking_mode: str,
509
+ context: Optional[List[Dict[str, Any]]] = None,
510
+ drop_thinking: bool = True,
511
+ add_default_bos_token: bool = True,
512
+ reasoning_effort: Optional[str] = None,
513
+ ) -> str:
514
+ """
515
+ Encode a list of messages into the DeepSeek-V4 prompt format.
516
+
517
+ This is the main entry point for encoding conversations. It handles:
518
+ - BOS token insertion
519
+ - Thinking mode with optional reasoning content dropping
520
+ - Tool message merging into user messages
521
+ - Multi-turn conversation context
522
+
523
+ Args:
524
+ messages: List of message dicts to encode.
525
+ thinking_mode: Either "chat" or "thinking".
526
+ context: Optional preceding context messages (already encoded prefix).
527
+ drop_thinking: If True, drop reasoning_content from earlier assistant turns
528
+ (only keep reasoning for messages after the last user message).
529
+ add_default_bos_token: Whether to prepend BOS token at conversation start.
530
+ reasoning_effort: Optional reasoning effort level ("max", "high", or None).
531
+
532
+ Returns:
533
+ The encoded prompt string.
534
+ """
535
+ context = context if context else []
536
+
537
+ # Preprocess: merge tool messages and sort tool results
538
+ messages = merge_tool_messages(messages)
539
+ messages = sort_tool_results_by_call_order(context + messages)[len(context):]
540
+ if context:
541
+ context = merge_tool_messages(context)
542
+ context = sort_tool_results_by_call_order(context)
543
+
544
+ full_messages = context + messages
545
+
546
+ prompt = bos_token if add_default_bos_token and len(context) == 0 else ""
547
+
548
+ # Resolve drop_thinking: if any message has tools defined, don't drop thinking
549
+ effective_drop_thinking = drop_thinking
550
+ if any(m.get("tools") for m in full_messages):
551
+ effective_drop_thinking = False
552
+
553
+ if thinking_mode == "thinking" and effective_drop_thinking:
554
+ full_messages = _drop_thinking_messages(full_messages)
555
+ # After dropping, recalculate how many messages to render
556
+ # (context may have shrunk too)
557
+ num_to_render = len(full_messages) - len(_drop_thinking_messages(context))
558
+ context_len = len(full_messages) - num_to_render
559
+ else:
560
+ num_to_render = len(messages)
561
+ context_len = len(context)
562
+
563
+ for idx in range(num_to_render):
564
+ prompt += render_message(
565
+ idx + context_len,
566
+ full_messages,
567
+ thinking_mode=thinking_mode,
568
+ drop_thinking=effective_drop_thinking,
569
+ reasoning_effort=reasoning_effort,
570
+ )
571
+
572
+ return prompt
573
+
574
+
575
+ def _drop_thinking_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
576
+ """
577
+ Drop reasoning_content and non-essential messages before the last user message.
578
+
579
+ Behavior:
580
+ - Messages with role in ["user", "system", "tool", "latest_reminder"] are always kept.
581
+ - Messages at or after the last user index are always kept.
582
+ - Assistant messages before the last user get reasoning_content removed.
583
+ - Developer messages before the last user are dropped entirely.
584
+ """
585
+ last_user_idx = find_last_user_index(messages)
586
+ result = []
587
+ keep_roles = {"user", "system", "tool", "latest_reminder", "direct_search_results"}
588
+
589
+ for idx, msg in enumerate(messages):
590
+ role = msg.get("role")
591
+ if role in keep_roles or idx >= last_user_idx:
592
+ result.append(msg)
593
+ elif role == "assistant":
594
+ msg = copy.copy(msg)
595
+ msg.pop("reasoning_content", None)
596
+ result.append(msg)
597
+ # developer and other roles before last_user_idx are dropped
598
+
599
+ return result
600
+
601
+
602
+ # ============================================================
603
+ # Parsing (Decoding model output)
604
+ # ============================================================
605
+
606
+ def _read_until_stop(index: int, text: str, stop: List[str]) -> Tuple[int, str, Optional[str]]:
607
+ """
608
+ Read text from index until one of the stop strings is found.
609
+
610
+ Returns:
611
+ Tuple of (new_index, content_before_stop, matched_stop_string_or_None).
612
+ """
613
+ min_pos = len(text)
614
+ matched_stop = None
615
+
616
+ for s in stop:
617
+ pos = text.find(s, index)
618
+ if pos != -1 and pos < min_pos:
619
+ min_pos = pos
620
+ matched_stop = s
621
+
622
+ if matched_stop:
623
+ content = text[index:min_pos]
624
+ return min_pos + len(matched_stop), content, matched_stop
625
+ else:
626
+ content = text[index:]
627
+ return len(text), content, None
628
+
629
+
630
+ def parse_tool_calls(index: int, text: str) -> Tuple[int, Optional[str], List[Dict[str, str]]]:
631
+ """
632
+ Parse DSML tool calls from text starting at the given index.
633
+
634
+ Args:
635
+ index: Starting position in text.
636
+ text: The full text to parse.
637
+
638
+ Returns:
639
+ Tuple of (new_index, last_stop_token, list_of_tool_call_dicts).
640
+ Each tool call dict has "name" and "arguments" keys.
641
+ """
642
+ tool_calls: List[Dict[str, Any]] = []
643
+ stop_token = None
644
+ tool_calls_end_token = f"</{dsml_token}{tool_calls_block_name}>"
645
+
646
+ while index < len(text):
647
+ index, _, stop_token = _read_until_stop(index, text, [f"<{dsml_token}invoke", tool_calls_end_token])
648
+ if _ != ">\n":
649
+ raise ValueError(f"Tool call format error: expected '>\\n' but got '{_}'")
650
+
651
+ if stop_token == tool_calls_end_token:
652
+ break
653
+
654
+ if stop_token is None:
655
+ raise ValueError("Missing special token in tool calls")
656
+
657
+ index, tool_name_content, stop_token = _read_until_stop(index, text, [f"<{dsml_token}parameter", f"</{dsml_token}invoke"])
658
+
659
+ p_tool_name = re.findall(r'^\s*name="(.*?)">\n$', tool_name_content, flags=re.DOTALL)
660
+ if len(p_tool_name) != 1:
661
+ raise ValueError(f"Tool name format error: '{tool_name_content}'")
662
+ tool_name = p_tool_name[0]
663
+
664
+ tool_args: Dict[str, Tuple[str, str]] = {}
665
+ while stop_token == f"<{dsml_token}parameter":
666
+ index, param_content, stop_token = _read_until_stop(index, text, [f"/{dsml_token}parameter"])
667
+
668
+ param_kv = re.findall(r'^ name="(.*?)" string="(true|false)">(.*?)<$', param_content, flags=re.DOTALL)
669
+ if len(param_kv) != 1:
670
+ raise ValueError(f"Parameter format error: '{param_content}'")
671
+ param_name, string, param_value = param_kv[0]
672
+
673
+ if param_name in tool_args:
674
+ raise ValueError(f"Duplicate parameter name: '{param_name}'")
675
+ tool_args[param_name] = (param_value, string)
676
+
677
+ index, content, stop_token = _read_until_stop(index, text, [f"<{dsml_token}parameter", f"</{dsml_token}invoke"])
678
+ if content != ">\n":
679
+ raise ValueError(f"Parameter format error: expected '>\\n' but got '{content}'")
680
+
681
+ tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args)
682
+ tool_calls.append(tool_call)
683
+
684
+ return index, stop_token, tool_calls
685
+
686
+
687
+ def parse_message_from_completion_text(text: str, thinking_mode: str) -> Dict[str, Any]:
688
+ """
689
+ Parse a model completion text into a structured assistant message.
690
+
691
+ This function takes the raw text output from the model (a single assistant turn)
692
+ and extracts:
693
+ - reasoning_content (thinking block)
694
+ - content (summary/response)
695
+ - tool_calls (if any)
696
+
697
+ NOTE: This function is designed to parse only correctly formatted strings and
698
+ will raise ValueError for malformed output.
699
+
700
+ Args:
701
+ text: The raw completion text (including EOS token).
702
+ thinking_mode: Either "chat" or "thinking".
703
+
704
+ Returns:
705
+ Dict with keys: "role", "content", "reasoning_content", "tool_calls".
706
+ tool_calls are in OpenAI format.
707
+ """
708
+ summary_content, reasoning_content, tool_calls = "", "", []
709
+ index, stop_token = 0, None
710
+ tool_calls_start_token = f"\n\n<{dsml_token}{tool_calls_block_name}"
711
+
712
+ is_thinking = thinking_mode == "thinking"
713
+ is_tool_calling = False
714
+
715
+ if is_thinking:
716
+ index, content_delta, stop_token = _read_until_stop(index, text, [thinking_end_token, tool_calls_start_token])
717
+ reasoning_content = content_delta
718
+ assert stop_token == thinking_end_token, "Invalid thinking format: missing </think>"
719
+
720
+ index, content_delta, stop_token = _read_until_stop(index, text, [eos_token, tool_calls_start_token])
721
+ summary_content = content_delta
722
+ if stop_token == tool_calls_start_token:
723
+ is_tool_calling = True
724
+ else:
725
+ assert stop_token == eos_token, "Invalid format: missing EOS token"
726
+
727
+ if is_tool_calling:
728
+ index, stop_token, tool_calls = parse_tool_calls(index, text)
729
+
730
+ index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token])
731
+ assert not tool_ends_text, "Unexpected content after tool calls"
732
+
733
+ assert len(text) == index and stop_token in [eos_token, None], "Unexpected content at end"
734
+
735
+ for sp_token in [bos_token, eos_token, thinking_start_token, thinking_end_token, dsml_token]:
736
+ assert sp_token not in summary_content and sp_token not in reasoning_content, \
737
+ f"Unexpected special token '{sp_token}' in content"
738
+
739
+ return {
740
+ "role": "assistant",
741
+ "content": summary_content,
742
+ "reasoning_content": reasoning_content,
743
+ "tool_calls": tool_calls_to_openai_format(tool_calls)
744
+ }
encoding/test_encoding_dsv4.py ADDED
@@ -0,0 +1,89 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Test suite for DeepSeek-V4 Encoding.
3
+
4
+ Run: python test_encoding_dsv4.py
5
+ """
6
+
7
+ import json
8
+ import os
9
+
10
+ from encoding_dsv4 import encode_messages, parse_message_from_completion_text
11
+
12
+ TESTS_DIR = os.path.join(os.path.dirname(__file__), "tests")
13
+
14
+
15
+ def test_case_1():
16
+ """Thinking mode with tool calls (multi-turn, tool results merged into user)."""
17
+ with open(os.path.join(TESTS_DIR, "test_input_1.json")) as f:
18
+ td = json.load(f)
19
+ messages = td["messages"]
20
+ messages[0]["tools"] = td["tools"]
21
+ gold = open(os.path.join(TESTS_DIR, "test_output_1.txt")).read()
22
+ prompt = encode_messages(messages, thinking_mode="thinking")
23
+ assert prompt == gold
24
+
25
+ # Parse: assistant turn with tool call
26
+ marker = "<|Assistant|><think>"
27
+ first_start = prompt.find(marker) + len(marker)
28
+ first_end = prompt.find("<|User|>", first_start)
29
+ parsed_tc = parse_message_from_completion_text(prompt[first_start:first_end], thinking_mode="thinking")
30
+ assert parsed_tc["reasoning_content"] == "The user wants to know the weather in Beijing. I should use the get_weather tool."
31
+ assert parsed_tc["content"] == ""
32
+ assert len(parsed_tc["tool_calls"]) == 1
33
+ assert parsed_tc["tool_calls"][0]["function"]["name"] == "get_weather"
34
+ assert json.loads(parsed_tc["tool_calls"][0]["function"]["arguments"]) == {"location": "Beijing", "unit": "celsius"}
35
+
36
+ # Parse: final assistant turn with content
37
+ last_start = prompt.rfind(marker) + len(marker)
38
+ parsed_final = parse_message_from_completion_text(prompt[last_start:], thinking_mode="thinking")
39
+ assert parsed_final["reasoning_content"] == "Got the weather data. Let me format a nice response."
40
+ assert "22°C" in parsed_final["content"]
41
+ assert parsed_final["tool_calls"] == []
42
+
43
+ print(" [PASS] case 1: thinking with tools (encode + parse)")
44
+
45
+
46
+ def test_case_2():
47
+ """Thinking mode without tools (drop_thinking removes earlier reasoning)."""
48
+ messages = json.load(open(os.path.join(TESTS_DIR, "test_input_2.json")))
49
+ gold = open(os.path.join(TESTS_DIR, "test_output_2.txt")).read()
50
+ prompt = encode_messages(messages, thinking_mode="thinking")
51
+ assert prompt == gold
52
+
53
+ # Parse: last assistant turn
54
+ marker = "<|Assistant|><think>"
55
+ last_start = prompt.rfind(marker) + len(marker)
56
+ parsed = parse_message_from_completion_text(prompt[last_start:], thinking_mode="thinking")
57
+ assert parsed["reasoning_content"] == "The user asks about the capital of France. It is Paris."
58
+ assert parsed["content"] == "The capital of France is Paris."
59
+ assert parsed["tool_calls"] == []
60
+
61
+ # Verify drop_thinking: first assistant's reasoning should be absent
62
+ assert "The user said hello" not in prompt
63
+
64
+ print(" [PASS] case 2: thinking without tools (encode + parse)")
65
+
66
+
67
+ def test_case_3():
68
+ """Interleaved thinking + search (developer with tools, latest_reminder)."""
69
+ messages = json.load(open(os.path.join(TESTS_DIR, "test_input_3.json")))
70
+ gold = open(os.path.join(TESTS_DIR, "test_output_3.txt")).read()
71
+ assert encode_messages(messages, thinking_mode="thinking") == gold
72
+ print(" [PASS] case 3: interleaved thinking + search")
73
+
74
+
75
+ def test_case_4():
76
+ """Quick instruction task with latest_reminder (chat mode, action task)."""
77
+ messages = json.load(open(os.path.join(TESTS_DIR, "test_input_4.json")))
78
+ gold = open(os.path.join(TESTS_DIR, "test_output_4.txt")).read()
79
+ assert encode_messages(messages, thinking_mode="chat") == gold
80
+ print(" [PASS] case 4: quick instruction task")
81
+
82
+
83
+ if __name__ == "__main__":
84
+ print("Running DeepSeek-V4 Encoding Tests...\n")
85
+ test_case_1()
86
+ test_case_2()
87
+ test_case_3()
88
+ test_case_4()
89
+ print("\nAll 4 tests passed!")
encoding/tests/test_input_1.json ADDED
@@ -0,0 +1,81 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "tools": [
3
+ {
4
+ "type": "function",
5
+ "function": {
6
+ "name": "get_weather",
7
+ "description": "Get the weather for a specific location",
8
+ "parameters": {
9
+ "type": "object",
10
+ "properties": {
11
+ "location": {
12
+ "type": "string",
13
+ "description": "The city name"
14
+ },
15
+ "unit": {
16
+ "type": "string",
17
+ "enum": ["celsius", "fahrenheit"],
18
+ "description": "Temperature unit"
19
+ }
20
+ },
21
+ "required": ["location"]
22
+ }
23
+ }
24
+ },
25
+ {
26
+ "type": "function",
27
+ "function": {
28
+ "name": "search",
29
+ "description": "Search the web for information",
30
+ "parameters": {
31
+ "type": "object",
32
+ "properties": {
33
+ "query": {
34
+ "type": "string",
35
+ "description": "Search query"
36
+ },
37
+ "num_results": {
38
+ "type": "integer",
39
+ "description": "Number of results to return"
40
+ }
41
+ },
42
+ "required": ["query"]
43
+ }
44
+ }
45
+ }
46
+ ],
47
+ "messages": [
48
+ {
49
+ "role": "system",
50
+ "content": "You are a helpful assistant."
51
+ },
52
+ {
53
+ "role": "user",
54
+ "content": "What's the weather in Beijing?"
55
+ },
56
+ {
57
+ "role": "assistant",
58
+ "reasoning_content": "The user wants to know the weather in Beijing. I should use the get_weather tool.",
59
+ "tool_calls": [
60
+ {
61
+ "id": "call_001",
62
+ "type": "function",
63
+ "function": {
64
+ "name": "get_weather",
65
+ "arguments": "{\"location\": \"Beijing\", \"unit\": \"celsius\"}"
66
+ }
67
+ }
68
+ ]
69
+ },
70
+ {
71
+ "role": "tool",
72
+ "tool_call_id": "call_001",
73
+ "content": "{\"temperature\": 22, \"condition\": \"sunny\", \"humidity\": 45}"
74
+ },
75
+ {
76
+ "role": "assistant",
77
+ "reasoning_content": "Got the weather data. Let me format a nice response.",
78
+ "content": "The weather in Beijing is currently sunny with a temperature of 22°C and 45% humidity."
79
+ }
80
+ ]
81
+ }
encoding/tests/test_input_2.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "role": "system",
4
+ "content": "You are a helpful assistant."
5
+ },
6
+ {
7
+ "role": "user",
8
+ "content": "Hello"
9
+ },
10
+ {
11
+ "role": "assistant",
12
+ "reasoning_content": "The user said hello, I should greet back.",
13
+ "content": "Hi there! How can I help you?"
14
+ },
15
+ {
16
+ "role": "user",
17
+ "content": "What is the capital of France?"
18
+ },
19
+ {
20
+ "role": "assistant",
21
+ "reasoning_content": "The user asks about the capital of France. It is Paris.",
22
+ "content": "The capital of France is Paris."
23
+ }
24
+ ]
encoding/tests/test_input_3.json ADDED
@@ -0,0 +1,159 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "role": "system",
4
+ "content": "该助手为DeepSeek,由深度求索公司创造。"
5
+ },
6
+ {
7
+ "role": "latest_reminder",
8
+ "content": "2026-02-21,星期六,广州,App,中文"
9
+ },
10
+ {
11
+ "role": "developer",
12
+ "content": "小柴胡冲剂和布洛芬能一起吃吗?\n\nCITATION FORMAT: 【{cursor_id}†L{start_line_id}(-L{end_line_id})?】",
13
+ "tools": [
14
+ {
15
+ "type": "function",
16
+ "function": {
17
+ "name": "search",
18
+ "description": "Web search. Split multiple queries with '||'.",
19
+ "parameters": {
20
+ "type": "object",
21
+ "properties": {
22
+ "queries": {
23
+ "type": "string",
24
+ "description": "query1||query2"
25
+ }
26
+ },
27
+ "required": [
28
+ "queries"
29
+ ],
30
+ "additionalProperties": false,
31
+ "$schema": "http://json-schema.org/draft-07/schema#"
32
+ }
33
+ }
34
+ },
35
+ {
36
+ "type": "function",
37
+ "function": {
38
+ "name": "open",
39
+ "description": "Batch open IDs (format 【{id}†...】) or URLs.",
40
+ "parameters": {
41
+ "type": "object",
42
+ "properties": {
43
+ "open_list": {
44
+ "type": "array",
45
+ "items": {
46
+ "type": "object",
47
+ "properties": {
48
+ "id": {
49
+ "description": "ID or URL",
50
+ "anyOf": [
51
+ {
52
+ "type": "integer"
53
+ },
54
+ {
55
+ "type": "string"
56
+ }
57
+ ],
58
+ "default": -1
59
+ },
60
+ "cursor": {
61
+ "type": "integer",
62
+ "description": "",
63
+ "default": -1
64
+ },
65
+ "loc": {
66
+ "type": "integer",
67
+ "description": "Start line",
68
+ "default": -1
69
+ },
70
+ "num_lines": {
71
+ "type": "integer",
72
+ "description": "",
73
+ "default": -1
74
+ },
75
+ "view_source": {
76
+ "type": "boolean",
77
+ "description": "",
78
+ "default": false
79
+ }
80
+ },
81
+ "additionalProperties": false
82
+ },
83
+ "description": ""
84
+ }
85
+ },
86
+ "required": [
87
+ "open_list"
88
+ ],
89
+ "additionalProperties": false,
90
+ "$schema": "http://json-schema.org/draft-07/schema#"
91
+ }
92
+ }
93
+ },
94
+ {
95
+ "type": "function",
96
+ "function": {
97
+ "name": "find",
98
+ "description": "Find exact text pattern in pages.",
99
+ "parameters": {
100
+ "type": "object",
101
+ "properties": {
102
+ "find_list": {
103
+ "type": "array",
104
+ "items": {
105
+ "type": "object",
106
+ "properties": {
107
+ "pattern": {
108
+ "type": "string",
109
+ "description": ""
110
+ },
111
+ "cursor": {
112
+ "type": "integer",
113
+ "description": "",
114
+ "default": -1
115
+ }
116
+ },
117
+ "required": [
118
+ "pattern"
119
+ ],
120
+ "additionalProperties": false
121
+ },
122
+ "description": ""
123
+ }
124
+ },
125
+ "required": [
126
+ "find_list"
127
+ ],
128
+ "additionalProperties": false,
129
+ "$schema": "http://json-schema.org/draft-07/schema#"
130
+ }
131
+ }
132
+ }
133
+ ]
134
+ },
135
+ {
136
+ "role": "assistant",
137
+ "content": "",
138
+ "reasoning_content": "用户想知道小柴胡冲剂和布洛芬能否一起服用。",
139
+ "tool_calls": [
140
+ {
141
+ "type": "function",
142
+ "function": {
143
+ "name": "search",
144
+ "arguments": "{\"queries\": \"小柴胡冲剂 布洛芬 相互作用 一起吃\"}"
145
+ }
146
+ }
147
+ ]
148
+ },
149
+ {
150
+ "role": "tool",
151
+ "content": "[0]"
152
+ },
153
+ {
154
+ "role": "assistant",
155
+ "content": "请及时就医。",
156
+ "reasoning_content": "现在开始组织回答。",
157
+ "tool_calls": []
158
+ }
159
+ ]
encoding/tests/test_input_4.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "role": "system",
4
+ "content": "该助手为DeepSeek-V3,由深度求索公司创造。\n今天是2025年10月17日,星期五。"
5
+ },
6
+ {
7
+ "role": "latest_reminder",
8
+ "content": "2024-11-15,上海市,App,中文"
9
+ },
10
+ {
11
+ "role": "user",
12
+ "content": "热海大滚锅是世界著名温泉吗"
13
+ },
14
+ {
15
+ "role": "assistant",
16
+ "content": "关于热海大滚锅是否是世界著名温泉,可以这样总结:它在**中国乃至全球的地热奇观中占有重要地位**,但“世界著名”的称号更侧重于它作为**独特的地质现象和旅游景点**,而非普遍的温泉疗养体验。\n\n为了让你快速了解,我整理了一个简要的表格:\n\n| 维度 | 热海大滚锅的地位与特点 |\n| :--- | :--- |\n| **地质奇观** | **世界罕见**的高温地热系统,泉眼多、水温高、形态多样。 |\n| **国内声誉** | **中国三大地热区之一**,**国家5A级旅游景区**,自古闻名(徐霞客曾游历并记载)。 |\n| **国际知名度** | 在一些旅游平台被提及为“世界六大温泉”之一,但此说法流传不广,其国际声誉更多建立在地质独特性上。 |\n| **核心体验** | **观赏地热奇观**(如97℃沸腾的“大滚锅”)、**体验温泉煮鸡蛋**。 |\n\n### 💡 游玩攻略与温馨提示\n\n如果你计划前往热海大滚锅,这里有一些实用信息供你参考:\n\n- **门票与开放时间**:\n - **门票**:景区门票约为**50元/人**。如果选择包含温泉沐浴的套餐,价格会更高,例如约**288元**。\n - **开放时间**:景区一般**08:00-18:00**开放,但具体时间可能变动,建议提前核实。\n\n- **特色体验**:\n - **温泉煮鸡蛋**:这几乎是必试项目。可以在景区门口购买用草绳串起的生鸡蛋(约5-8元/串),然后到“大滚锅”旁的指定区域蒸煮,几分钟便可熟食,趣味十足。\n - **金汤足浴**:可以直接用从“大滚锅”流出的温泉水泡脚,缓解旅途疲劳。\n\n- **注意事项**:\n - **安全第一**:“大滚锅”水温极高,务必遵守游览规则,在指定区域内观赏,切勿随意触碰泉水。\n - **规划行程**:建议为热海景区预留**3-4小时**的游览时间。景区内步道不走回头路,出入口有观光车接送。\n\n希望这些信息能帮助你更好地了解热海大滚锅。如果你对腾冲的其他景点或者行程规划有更多疑问,我很乐意提供进一步的信息。",
17
+ "mask": 1
18
+ },
19
+ {
20
+ "role": "user",
21
+ "content": "世界著名温泉有哪些",
22
+ "task": "action"
23
+ },
24
+ {
25
+ "role": "assistant",
26
+ "content": "Search"
27
+ }
28
+ ]
encoding/tests/test_output_1.txt ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <|begin▁of▁sentence|>You are a helpful assistant.
2
+
3
+ ## Tools
4
+
5
+ You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<|DSML|tool_calls>" block like the following:
6
+
7
+ <|DSML|tool_calls>
8
+ <|DSML|invoke name="$TOOL_NAME">
9
+ <|DSML|parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE</|DSML|parameter>
10
+ ...
11
+ </|DSML|invoke>
12
+ <|DSML|invoke name="$TOOL_NAME2">
13
+ ...
14
+ </|DSML|invoke>
15
+ </|DSML|tool_calls>
16
+
17
+ String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.
18
+
19
+ If thinking_mode is enabled (triggered by <think>), you MUST output your complete reasoning inside <think>...</think> BEFORE any tool calls or final response.
20
+
21
+ Otherwise, output directly after </think> with tool calls or final response.
22
+
23
+ ### Available Tool Schemas
24
+
25
+ {"name": "get_weather", "description": "Get the weather for a specific location", "parameters": {"type": "object", "properties": {"location": {"type": "string", "description": "The city name"}, "unit": {"type": "string", "enum": ["celsius", "fahrenheit"], "description": "Temperature unit"}}, "required": ["location"]}}
26
+ {"name": "search", "description": "Search the web for information", "parameters": {"type": "object", "properties": {"query": {"type": "string", "description": "Search query"}, "num_results": {"type": "integer", "description": "Number of results to return"}}, "required": ["query"]}}
27
+
28
+ You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.
29
+ <|User|>What's the weather in Beijing?<|Assistant|><think>The user wants to know the weather in Beijing. I should use the get_weather tool.</think>
30
+
31
+ <|DSML|tool_calls>
32
+ <|DSML|invoke name="get_weather">
33
+ <|DSML|parameter name="location" string="true">Beijing</|DSML|parameter>
34
+ <|DSML|parameter name="unit" string="true">celsius</|DSML|parameter>
35
+ </|DSML|invoke>
36
+ </|DSML|tool_calls><|end▁of▁sentence|><|User|><tool_result>{"temperature": 22, "condition": "sunny", "humidity": 45}</tool_result><|Assistant|><think>Got the weather data. Let me format a nice response.</think>The weather in Beijing is currently sunny with a temperature of 22°C and 45% humidity.<|end▁of▁sentence|>
encoding/tests/test_output_2.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ <|begin▁of▁sentence|>You are a helpful assistant.<|User|>Hello<|Assistant|></think>Hi there! How can I help you?<|end▁of▁sentence|><|User|>What is the capital of France?<|Assistant|><think>The user asks about the capital of France. It is Paris.</think>The capital of France is Paris.<|end▁of▁sentence|>
encoding/tests/test_output_3.txt ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <|begin▁of▁sentence|>该助手为DeepSeek,由深度求索公司创造。<|latest_reminder|>2026-02-21,星期六,广州,App,中文<|User|>小柴胡冲剂和布洛芬能一起吃吗?
2
+
3
+ CITATION FORMAT: 【{cursor_id}†L{start_line_id}(-L{end_line_id})?】
4
+
5
+ ## Tools
6
+
7
+ You have access to a set of tools to help answer the user's question. You can invoke tools by writing a "<|DSML|tool_calls>" block like the following:
8
+
9
+ <|DSML|tool_calls>
10
+ <|DSML|invoke name="$TOOL_NAME">
11
+ <|DSML|parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE</|DSML|parameter>
12
+ ...
13
+ </|DSML|invoke>
14
+ <|DSML|invoke name="$TOOL_NAME2">
15
+ ...
16
+ </|DSML|invoke>
17
+ </|DSML|tool_calls>
18
+
19
+ String parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.
20
+
21
+ If thinking_mode is enabled (triggered by <think>), you MUST output your complete reasoning inside <think>...</think> BEFORE any tool calls or final response.
22
+
23
+ Otherwise, output directly after </think> with tool calls or final response.
24
+
25
+ ### Available Tool Schemas
26
+
27
+ {"name": "search", "description": "Web search. Split multiple queries with '||'.", "parameters": {"type": "object", "properties": {"queries": {"type": "string", "description": "query1||query2"}}, "required": ["queries"], "additionalProperties": false, "$schema": "http://json-schema.org/draft-07/schema#"}}
28
+ {"name": "open", "description": "Batch open IDs (format 【{id}†...】) or URLs.", "parameters": {"type": "object", "properties": {"open_list": {"type": "array", "items": {"type": "object", "properties": {"id": {"description": "ID or URL", "anyOf": [{"type": "integer"}, {"type": "string"}], "default": -1}, "cursor": {"type": "integer", "description": "", "default": -1}, "loc": {"type": "integer", "description": "Start line", "default": -1}, "num_lines": {"type": "integer", "description": "", "default": -1}, "view_source": {"type": "boolean", "description": "", "default": false}}, "additionalProperties": false}, "description": ""}}, "required": ["open_list"], "additionalProperties": false, "$schema": "http://json-schema.org/draft-07/schema#"}}
29
+ {"name": "find", "description": "Find exact text pattern in pages.", "parameters": {"type": "object", "properties": {"find_list": {"type": "array", "items": {"type": "object", "properties": {"pattern": {"type": "string", "description": ""}, "cursor": {"type": "integer", "description": "", "default": -1}}, "required": ["pattern"], "additionalProperties": false}, "description": ""}}, "required": ["find_list"], "additionalProperties": false, "$schema": "http://json-schema.org/draft-07/schema#"}}
30
+
31
+ You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.
32
+ <|Assistant|><think>用户想知道小柴胡冲剂和布洛芬能否一起服用。</think>
33
+
34
+ <|DSML|tool_calls>
35
+ <|DSML|invoke name="search">
36
+ <|DSML|parameter name="queries" string="true">小柴胡冲剂 布洛芬 相互作用 一起吃</|DSML|parameter>
37
+ </|DSML|invoke>
38
+ </|DSML|tool_calls><|end▁of▁sentence|><|User|><tool_result>[0]</tool_result><|Assistant|><think>现在开始组织回答。</think>请及时就医。<|end▁of▁sentence|>
encoding/tests/test_output_4.txt ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <|begin▁of▁sentence|>该助手为DeepSeek-V3,由深度求索公司创造。
2
+ 今天是2025年10月17日,星期五。<|latest_reminder|>2024-11-15,上海市,App,中文<|User|>热海大滚锅是世界著名温泉吗<|Assistant|></think>关于热海大滚锅是否是世界著名温泉,可以这样总结:它在**中国乃至全球的地热奇观中占有重要地位**,但“世界著名”的称号更侧重于它作为**独特的地质现象和旅游景点**,而非普遍的温泉疗养体验。
3
+
4
+ 为了让你快速了解,我整理了一个简要的表格:
5
+
6
+ | 维度 | 热海大滚锅的地位与特点 |
7
+ | :--- | :--- |
8
+ | **地质奇观** | **世界罕见**的高温地热系统,泉眼多、水温高、形态多样。 |
9
+ | **国内声誉** | **中国三大地热区之一**,**国家5A级旅游景区**,自古闻名(徐霞客曾游历并记载)。 |
10
+ | **国际知名度** | 在一些旅游平台被提及为“世界六大温泉”之一,但此说法流传不广,其国际声誉更多建立在地质独特性上。 |
11
+ | **核心体验** | **观赏地热奇观**(如97℃沸腾的“大滚锅”)、**体验温泉煮鸡蛋**。 |
12
+
13
+ ### 💡 游玩攻略与温馨提示
14
+
15
+ 如果你计划前往热海大滚锅,这里有一些实用信息供你参考:
16
+
17
+ - **门票与开放时间**:
18
+ - **门票**:景区门票约为**50元/人**。如果选择包含温泉沐浴的套餐,价格会更高,例如约**288元**。
19
+ - **开放时间**:景区一般**08:00-18:00**开放,但具体时间可能变动,建议提前核实。
20
+
21
+ - **特色体验**:
22
+ - **温泉煮鸡蛋**:这几乎是必试项目。可以在景区门口购买用草绳串起的生鸡蛋(约5-8元/串),然后到“大滚锅”旁的指定区域蒸煮,几分钟便可熟食,趣味十足。
23
+ - **金汤足浴**:可以直接用从“大滚锅”流出的温泉水泡脚,缓解旅途疲劳。
24
+
25
+ - **注意事项**:
26
+ - **安全第一**:“大滚锅”水温极高,务必遵守游览规则,在指定区域内观赏,切勿随意触碰泉水。
27
+ - **规划行程**:建议为热海景区预留**3-4小时**的游览时间。景区内步道不走回头路,出入口有观光车接送。
28
+
29
+ 希望这些信息能帮助你更好地了解热海大滚锅。如果你对腾冲的其他景点或者行程规划有更多疑问,我很乐意提供进一步的信息。<|end▁of▁sentence|><|User|>世界著名温泉有哪些<|Assistant|></think><|action|>Search<|end▁of▁sentence|>
generation_config.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 0,
4
+ "eos_token_id": 1,
5
+ "do_sample": true,
6
+ "temperature": 1.0,
7
+ "top_p": 1.0,
8
+ "transformers_version": "4.46.3"
9
+ }
inference/README.md ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Inference code for DeepSeek models
2
+
3
+ First convert huggingface model weight files to the format of this project.
4
+ ```bash
5
+ export EXPERTS=256
6
+ export MP=4
7
+ export CONFIG=config.json
8
+ python convert.py --hf-ckpt-path ${HF_CKPT_PATH} --save-path ${SAVE_PATH} --n-experts ${EXPERTS} --model-parallel ${MP}
9
+ ```
10
+
11
+ Then chat with DeepSeek model at will!
12
+ ```bash
13
+ torchrun --nproc-per-node ${MP} generate.py --ckpt-path ${SAVE_PATH} --config ${CONFIG} --interactive
14
+ ```
15
+
16
+ Or batch inference from file.
17
+ ```bash
18
+ torchrun --nproc-per-node ${MP} generate.py --ckpt-path ${SAVE_PATH} --config ${CONFIG} --input-file ${FILE}
19
+ ```
20
+
21
+ Or multi nodes inference.
22
+ ```bash
23
+ torchrun --nnodes ${NODES} --nproc-per-node $((MP / NODES)) --node-rank $RANK --master-addr $ADDR generate.py --ckpt-path ${SAVE_PATH} --config ${CONFIG} --input-file ${FILE}
24
+ ```
25
+
26
+ If you want to use fp8, just remove `"expert_dtype": "fp4"` in `config.json` and specify `--expert-dtype fp8` in `convert.py`.
inference/config.json ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "vocab_size": 129280,
3
+ "dim": 4096,
4
+ "moe_inter_dim": 2048,
5
+ "n_layers": 43,
6
+ "n_hash_layers": 3,
7
+ "n_mtp_layers": 3,
8
+ "dspark_block_size": 5,
9
+ "dspark_noise_token_id": 128799,
10
+ "dspark_target_layer_ids": [40, 41, 42],
11
+ "dspark_markov_rank": 256,
12
+ "n_heads": 64,
13
+ "n_routed_experts": 256,
14
+ "n_shared_experts": 1,
15
+ "n_activated_experts": 6,
16
+ "score_func": "sqrtsoftplus",
17
+ "route_scale": 1.5,
18
+ "swiglu_limit": 10.0,
19
+ "q_lora_rank": 1024,
20
+ "head_dim": 512,
21
+ "rope_head_dim": 64,
22
+ "o_groups": 8,
23
+ "o_lora_rank": 1024,
24
+ "window_size": 128,
25
+ "original_seq_len": 65536,
26
+ "rope_theta": 10000,
27
+ "rope_factor": 16,
28
+ "beta_fast": 32,
29
+ "beta_slow": 1,
30
+ "index_n_heads": 64,
31
+ "index_head_dim": 128,
32
+ "index_topk": 512,
33
+ "hc_mult": 4,
34
+ "hc_sinkhorn_iters": 20,
35
+ "dtype": "fp8",
36
+ "scale_fmt": "ue8m0",
37
+ "expert_dtype": "fp4",
38
+ "compress_rope_theta": 160000,
39
+ "compress_ratios": [0, 0, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 128, 4, 0, 0, 0]
40
+ }
inference/convert.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import shutil
3
+ from argparse import ArgumentParser
4
+ from glob import glob
5
+ from tqdm import tqdm, trange
6
+
7
+ import torch
8
+ from safetensors.torch import safe_open, save_file
9
+
10
+
11
+ FP4_TABLE = torch.tensor([
12
+ 0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
13
+ 0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0
14
+ ], dtype=torch.float32)
15
+
16
+
17
+ def cast_e2m1fn_to_e4m3fn(x: torch.Tensor, scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
18
+ """
19
+ Casts a tensor from e2m1fn to e4m3fn losslessly.
20
+ """
21
+ assert x.dtype == torch.int8
22
+ assert x.ndim == 2
23
+ out_dim, in_dim = x.size()
24
+ in_dim *= 2
25
+ fp8_block_size = 128
26
+ fp4_block_size = 32
27
+ assert in_dim % fp8_block_size == 0 and out_dim % fp8_block_size == 0
28
+ assert scale.size(0) == out_dim and scale.size(1) == in_dim // fp4_block_size
29
+
30
+ x = x.view(torch.uint8)
31
+ low = x & 0x0F
32
+ high = (x >> 4) & 0x0F
33
+ x = torch.stack([FP4_TABLE[low.long()], FP4_TABLE[high.long()]], dim=-1).flatten(2)
34
+
35
+ # max_fp4 (6.0) * MAX_OFFSET must fit in e4m3fn (max 448)
36
+ # 6.0 * 2^6 = 384 < 448; 6.0 * 2^7 = 768 > 448; so MAX_OFFSET_BITS = 6
37
+ MAX_OFFSET_BITS = 6
38
+
39
+ bOut = out_dim // fp8_block_size
40
+ bIn = in_dim // fp8_block_size
41
+ # bOut, bIn, 128, 128
42
+ x = x.view(bOut, fp8_block_size, bIn, fp8_block_size).transpose(1, 2)
43
+ # bOut, bIn, 128*4
44
+ scale = scale.float().view(bOut, fp8_block_size, bIn, -1).transpose(1, 2).flatten(2)
45
+ ## bOut, bIn, 1
46
+ scale_max_offset_bits = scale.amax(dim=-1, keepdim=True) / (2**MAX_OFFSET_BITS)
47
+ # bOut, bIn, 128*4
48
+ offset = scale / scale_max_offset_bits
49
+ # bOut, bIn, 128, 128
50
+ offset = offset.unflatten(-1, (fp8_block_size, -1)).repeat_interleave(fp4_block_size, dim=-1)
51
+ x = (x * offset).transpose(1, 2).reshape(out_dim, in_dim)
52
+ return x.to(torch.float8_e4m3fn), scale_max_offset_bits.squeeze(-1).to(torch.float8_e8m0fnu)
53
+
54
+
55
+ mapping = {
56
+ "embed": ("embed", 0),
57
+ "wq_b": ("wq_b", 0),
58
+ "wo_a": ("wo_a", 0),
59
+ "wo_b": ("wo_b", 1),
60
+ "head": ("head", 0),
61
+ "attn_sink": ("attn_sink", 0),
62
+ "weights_proj": ("weights_proj", 0),
63
+ "markov_w1": ("markov_w1", 0),
64
+ "markov_w2": ("markov_w2", 0),
65
+ }
66
+
67
+
68
+ def main(hf_ckpt_path, save_path, n_experts, mp, expert_dtype):
69
+ """
70
+ Converts and saves model checkpoint files into a specified format.
71
+
72
+ Args:
73
+ hf_ckpt_path (str): Path to the directory containing the input checkpoint files.
74
+ save_path (str): Path to the directory where the converted checkpoint files will be saved.
75
+ n_experts (int): Total number of experts in the model.
76
+ mp (int): Model parallelism factor.
77
+
78
+ Returns:
79
+ None
80
+ """
81
+ torch.set_num_threads(8)
82
+ n_local_experts = n_experts // mp
83
+ state_dicts = [{} for _ in range(mp)]
84
+
85
+ for file_path in tqdm(glob(os.path.join(hf_ckpt_path, "*.safetensors"))):
86
+ with safe_open(file_path, framework="pt", device="cpu") as f:
87
+ for name in f.keys():
88
+ param: torch.Tensor = f.get_tensor(name)
89
+ if name.startswith("model."):
90
+ name = name[len("model."):]
91
+ if name.startswith("mtp.") and ("emb" in name or name.endswith("head.weight")):
92
+ continue
93
+ name = name.replace("self_attn", "attn")
94
+ name = name.replace("mlp", "ffn")
95
+ name = name.replace("weight_scale_inv", "scale")
96
+ name = name.replace("e_score_correction_bias", "bias")
97
+ if any(x in name for x in ["hc", "attn_sink", "tie2eid", "ape"]): # without .weight
98
+ key = name.split(".")[-1]
99
+ else:
100
+ key = name.split(".")[-2]
101
+ if key in mapping:
102
+ new_key, dim = mapping[key]
103
+ else:
104
+ new_key, dim = key, None
105
+ name = name.replace(key, new_key)
106
+ for i in range(mp):
107
+ new_param = param
108
+ if "experts" in name and "shared_experts" not in name:
109
+ idx = int(name.split(".")[-3])
110
+ if idx < i * n_local_experts or idx >= (i + 1) * n_local_experts:
111
+ continue
112
+ elif dim is not None:
113
+ assert param.size(dim) % mp == 0, f"Dimension {dim} must be divisible by {mp}"
114
+ shard_size = param.size(dim) // mp
115
+ new_param = param.narrow(dim, i * shard_size, shard_size).contiguous()
116
+ state_dicts[i][name] = new_param
117
+
118
+ os.makedirs(save_path, exist_ok=True)
119
+
120
+ for i in trange(mp):
121
+ names = list(state_dicts[i].keys())
122
+ for name in names:
123
+ if name.endswith("wo_a.weight"):
124
+ weight = state_dicts[i][name]
125
+ scale = state_dicts[i].pop(name.replace("weight", "scale"))
126
+ weight = weight.unflatten(0, (-1, 128)).unflatten(-1, (-1, 128)).float() * scale[:, None, :, None].float()
127
+ state_dicts[i][name] = weight.flatten(2, 3).flatten(0, 1).bfloat16()
128
+ elif "experts" in name and state_dicts[i][name].dtype == torch.int8:
129
+ if expert_dtype == "fp8":
130
+ scale_name = name.replace("weight", "scale")
131
+ weight = state_dicts[i].pop(name)
132
+ scale = state_dicts[i].pop(scale_name)
133
+ state_dicts[i][name], state_dicts[i][scale_name] = cast_e2m1fn_to_e4m3fn(weight, scale)
134
+ else:
135
+ state_dicts[i][name] = state_dicts[i][name].view(torch.float4_e2m1fn_x2)
136
+ save_file(state_dicts[i], os.path.join(save_path, f"model{i}-mp{mp}.safetensors"))
137
+
138
+ for file in ["tokenizer.json", "tokenizer_config.json"]:
139
+ old_file_path = os.path.join(hf_ckpt_path, file)
140
+ new_file_path = os.path.join(save_path, file)
141
+ if os.path.exists(old_file_path):
142
+ shutil.copyfile(old_file_path, new_file_path)
143
+
144
+
145
+ if __name__ == "__main__":
146
+ parser = ArgumentParser()
147
+ parser.add_argument("--hf-ckpt-path", type=str, required=True)
148
+ parser.add_argument("--save-path", type=str, required=True)
149
+ parser.add_argument("--n-experts", type=int, required=True)
150
+ parser.add_argument("--model-parallel", type=int, required=True)
151
+ parser.add_argument("--expert-dtype", type=str, choices=["fp8", "fp4"], required=False, default=None)
152
+ args = parser.parse_args()
153
+ assert args.n_experts % args.model_parallel == 0, "Number of experts must be divisible by model parallelism"
154
+ main(args.hf_ckpt_path, args.save_path, args.n_experts, args.model_parallel, args.expert_dtype)
inference/generate.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import sys
4
+ from argparse import ArgumentParser
5
+ from typing import List
6
+
7
+ import torch
8
+ import torch.distributed as dist
9
+ from transformers import AutoTokenizer
10
+ from safetensors.torch import load_model
11
+
12
+ from model import Transformer, ModelArgs
13
+ current_dir = os.path.dirname(os.path.abspath(__file__))
14
+ encoding_dir = os.path.join(current_dir, '../encoding')
15
+ sys.path.insert(0, os.path.abspath(encoding_dir))
16
+ from encoding_dsv4 import encode_messages, parse_message_from_completion_text
17
+
18
+
19
+ @torch.inference_mode()
20
+ def generate(
21
+ model: Transformer,
22
+ prompt_tokens: List[List[int]],
23
+ max_new_tokens: int,
24
+ eos_id: int,
25
+ ) -> List[List[int]]:
26
+ """Batch generation with left-padded prompts.
27
+
28
+ The first forward pass processes [min_prompt_len:] tokens (prefill phase).
29
+ Subsequent passes generate one token at a time (decode phase). For positions
30
+ still within a prompt, the ground-truth token overrides the model's prediction.
31
+ """
32
+ prompt_lens = [len(t) for t in prompt_tokens]
33
+ assert max(prompt_lens) <= model.max_seq_len, f"Prompt length exceeds model maximum sequence length (max_seq_len={model.max_seq_len})"
34
+ total_len = min(model.max_seq_len, max_new_tokens + max(prompt_lens))
35
+ tokens = torch.full((len(prompt_tokens), total_len), -1, dtype=torch.long)
36
+ for i, t in enumerate(prompt_tokens):
37
+ tokens[i, :len(t)] = torch.tensor(t, dtype=torch.long)
38
+ prev_pos = 0
39
+ finished = torch.tensor([False] * len(prompt_tokens))
40
+ prompt_mask = tokens != -1
41
+ for cur_pos in range(min(prompt_lens), total_len):
42
+ next_token = model.forward(tokens[:, prev_pos:cur_pos], prev_pos)[0]
43
+ next_token = torch.where(prompt_mask[:, cur_pos], tokens[:, cur_pos], next_token)
44
+ tokens[:, cur_pos] = next_token
45
+ finished |= torch.logical_and(~prompt_mask[:, cur_pos], next_token == eos_id)
46
+ prev_pos = cur_pos
47
+ if finished.all():
48
+ break
49
+ completion_tokens = []
50
+ for i, toks in enumerate(tokens.tolist()):
51
+ toks = toks[prompt_lens[i]:prompt_lens[i]+max_new_tokens]
52
+ if eos_id in toks:
53
+ toks = toks[:toks.index(eos_id)]
54
+ toks.append(eos_id)
55
+ completion_tokens.append(toks)
56
+ return completion_tokens
57
+
58
+
59
+ def main(
60
+ ckpt_path: str,
61
+ config: str,
62
+ input_file: str = "",
63
+ interactive: bool = True,
64
+ max_new_tokens: int = 100,
65
+ temperature: float = 1.0,
66
+ ) -> None:
67
+ world_size = int(os.getenv("WORLD_SIZE", "1"))
68
+ rank = int(os.getenv("RANK", "0"))
69
+ local_rank = int(os.getenv("LOCAL_RANK", "0"))
70
+ if world_size > 1:
71
+ dist.init_process_group("nccl")
72
+ global print
73
+ if rank != 0:
74
+ print = lambda *_, **__: None
75
+ torch.cuda.set_device(local_rank)
76
+ torch.cuda.memory._set_allocator_settings("expandable_segments:True")
77
+ torch.set_default_dtype(torch.bfloat16)
78
+ torch.set_num_threads(8)
79
+ torch.manual_seed(33377335)
80
+ with open(config) as f:
81
+ args = ModelArgs(**json.load(f))
82
+ args.temperature = temperature
83
+ if interactive:
84
+ args.max_batch_size = 1
85
+ args.max_seq_len = 64 * 1024
86
+ print(args)
87
+ with torch.device("cuda"):
88
+ model = Transformer(args)
89
+ tokenizer = AutoTokenizer.from_pretrained(ckpt_path)
90
+ print("load model")
91
+ load_model(model, os.path.join(ckpt_path, f"model{rank}-mp{world_size}.safetensors"), strict=False)
92
+ torch.set_default_device("cuda")
93
+ print("I'm DeepSeek 👋")
94
+
95
+ if interactive:
96
+ messages = []
97
+ while True:
98
+ if world_size == 1:
99
+ prompt = input(">>> ")
100
+ elif rank == 0:
101
+ prompt = input(">>> ")
102
+ objects = [prompt]
103
+ dist.broadcast_object_list(objects, 0)
104
+ else:
105
+ objects = [None]
106
+ dist.broadcast_object_list(objects, 0)
107
+ prompt = objects[0]
108
+ if prompt == "/exit":
109
+ break
110
+ elif prompt == "/clear":
111
+ messages.clear()
112
+ continue
113
+ messages.append({"role": "user", "content": prompt})
114
+ prompt_tokens = tokenizer.encode(encode_messages(messages, thinking_mode="chat"))
115
+ completion_tokens = generate(model, [prompt_tokens], max_new_tokens, tokenizer.eos_token_id)
116
+ completion = tokenizer.decode(completion_tokens[0])
117
+ print(completion)
118
+ messages.append(parse_message_from_completion_text(completion, thinking_mode="chat"))
119
+ else:
120
+ with open(input_file) as f:
121
+ prompts = f.read().split("\n\n")
122
+ prompt_tokens = [tokenizer.encode(encode_messages([{"role": "user", "content": prompt}], thinking_mode="chat")) for prompt in prompts]
123
+ completion_tokens = generate(model, prompt_tokens, max_new_tokens, tokenizer.eos_token_id)
124
+ completions = tokenizer.batch_decode(completion_tokens)
125
+ for prompt, completion in zip(prompts, completions):
126
+ print("Prompt:", prompt)
127
+ print("Completion:", completion)
128
+ print()
129
+
130
+ if world_size > 1:
131
+ dist.destroy_process_group()
132
+
133
+
134
+ if __name__ == "__main__":
135
+ parser = ArgumentParser()
136
+ parser.add_argument("--ckpt-path", type=str, required=True)
137
+ parser.add_argument("--config", type=str, required=True)
138
+ parser.add_argument("--input-file", type=str, default="")
139
+ parser.add_argument("--interactive", action="store_true")
140
+ parser.add_argument("--max-new-tokens", type=int, default=300)
141
+ parser.add_argument("--temperature", type=float, default=1.0)
142
+ args = parser.parse_args()
143
+ assert args.input_file or args.interactive, "Either input-file or interactive mode must be specified"
144
+ main(args.ckpt_path, args.config, args.input_file, args.interactive, args.max_new_tokens, args.temperature)
inference/kernel.py ADDED
@@ -0,0 +1,536 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import tilelang
3
+ import tilelang.language as T
4
+ from typing import Tuple, Optional
5
+
6
+
7
+ tilelang.set_log_level("WARNING")
8
+
9
+ pass_configs = {
10
+ tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
11
+ tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
12
+ }
13
+
14
+ FP8 = "float8_e4m3"
15
+ FP4 = "float4_e2m1fn"
16
+ FE8M0 = "float8_e8m0fnu"
17
+ BF16 = "bfloat16"
18
+ FP32 = "float32"
19
+ INT32 = "int32"
20
+
21
+
22
+ def fast_log2_ceil(x):
23
+ """Compute ceil(log2(x)) via IEEE 754 bit manipulation. Avoids slow log/ceil intrinsics."""
24
+ bits_x = T.reinterpret("uint32", x)
25
+ exp_x = (bits_x >> 23) & 0xFF
26
+ man_bits = bits_x & ((1 << 23) - 1)
27
+ return T.Cast("int32", exp_x - 127 + T.if_then_else(man_bits != 0, 1, 0))
28
+
29
+
30
+ def fast_pow2(x):
31
+ """Compute 2^x for integer x via IEEE 754 bit manipulation."""
32
+ bits_x = (x + 127) << 23
33
+ return T.reinterpret("float32", bits_x)
34
+
35
+
36
+ def fast_round_scale(amax, fp8_max_inv):
37
+ return fast_pow2(fast_log2_ceil(amax * fp8_max_inv))
38
+
39
+
40
+ @tilelang.jit(pass_configs=pass_configs)
41
+ def act_quant_kernel(
42
+ N, block_size=128, in_dtype=BF16, out_dtype=FP8, scale_dtype=FP32,
43
+ round_scale=False, inplace=False
44
+ ):
45
+ """Block-wise FP8 quantization. inplace=True does fused quant+dequant back to BF16."""
46
+ M = T.symbolic("M")
47
+ fp8_min = -448.0
48
+ fp8_max = 448.0
49
+ fp8_max_inv = 1 / fp8_max
50
+ num_stages = 0 if round_scale or inplace else 2
51
+ blk_m = 32
52
+ group_size = block_size
53
+ # Internal computation in FP32; scale_dtype controls output storage format.
54
+ compute_dtype = FP32
55
+ out_dtype = in_dtype if inplace else out_dtype
56
+
57
+ @T.prim_func
58
+ def act_quant_kernel_(
59
+ X: T.Tensor[(M, N), in_dtype],
60
+ Y: T.Tensor[(M, N), out_dtype],
61
+ S: T.Tensor[(M, T.ceildiv(N, group_size)), scale_dtype],
62
+ ):
63
+ with T.Kernel(T.ceildiv(M, blk_m), T.ceildiv(N, group_size), threads=128) as (
64
+ pid_m,
65
+ pid_n,
66
+ ):
67
+ x_shared = T.alloc_shared((blk_m, group_size), in_dtype)
68
+ x_local = T.alloc_fragment((blk_m, group_size), in_dtype)
69
+ amax_local = T.alloc_fragment((blk_m,), compute_dtype)
70
+ s_local = T.alloc_fragment((blk_m,), compute_dtype)
71
+ y_local = T.alloc_fragment((blk_m, group_size), out_dtype)
72
+ y_shared = T.alloc_shared((blk_m, group_size), out_dtype)
73
+
74
+ for _ in T.Pipelined(1, num_stages=num_stages):
75
+ T.copy(X[pid_m * blk_m, pid_n * group_size], x_shared)
76
+ T.copy(x_shared, x_local)
77
+ T.reduce_absmax(x_local, amax_local, dim=1)
78
+ for i in T.Parallel(blk_m):
79
+ amax_local[i] = T.max(amax_local[i], 1e-4)
80
+ if round_scale:
81
+ s_local[i] = fast_round_scale(amax_local[i], fp8_max_inv)
82
+ else:
83
+ s_local[i] = amax_local[i] * fp8_max_inv
84
+ if inplace:
85
+ for i, j in T.Parallel(blk_m, group_size):
86
+ y_local[i, j] = T.Cast(
87
+ out_dtype,
88
+ T.Cast(compute_dtype, T.Cast(FP8, T.clamp(
89
+ x_local[i, j] / s_local[i], fp8_min, fp8_max
90
+ ))) * s_local[i],
91
+ )
92
+ else:
93
+ for i, j in T.Parallel(blk_m, group_size):
94
+ y_local[i, j] = T.clamp(
95
+ x_local[i, j] / s_local[i], fp8_min, fp8_max
96
+ )
97
+ for i in T.Parallel(blk_m):
98
+ S[pid_m * blk_m + i, pid_n] = T.Cast(scale_dtype, s_local[i])
99
+ T.copy(y_local, y_shared)
100
+ T.copy(y_shared, Y[pid_m * blk_m, pid_n * group_size])
101
+
102
+ return act_quant_kernel_
103
+
104
+
105
+ def act_quant(
106
+ x: torch.Tensor, block_size: int = 128, scale_fmt: Optional[str] = None,
107
+ scale_dtype: torch.dtype = torch.float32, inplace: bool = False,
108
+ ) -> torch.Tensor:
109
+ """Block-wise FP8 quantization. inplace=True does fused quant+dequant back to BF16.
110
+ When scale_fmt is set, scales are rounded to power-of-2 (MXFP)."""
111
+ N = x.size(-1)
112
+ assert N % block_size == 0
113
+ tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32
114
+ z = x.contiguous()
115
+ y = torch.empty_like(z) if inplace else torch.empty_like(z, dtype=torch.float8_e4m3fn)
116
+ s = z.new_empty(*z.size()[:-1], N // block_size, dtype=scale_dtype)
117
+ kernel = act_quant_kernel(
118
+ N, block_size, scale_dtype=tl_dtype,
119
+ round_scale=scale_fmt is not None, inplace=inplace,
120
+ )
121
+ kernel(z.view(-1, N), y.view(-1, N), s.view(-1, N // block_size))
122
+ if inplace:
123
+ x.copy_(y)
124
+ return x
125
+ return y, s
126
+
127
+
128
+ @tilelang.jit(pass_configs=pass_configs)
129
+ def fp4_quant_kernel(
130
+ N, block_size=32, in_dtype=BF16, scale_dtype=FE8M0, inplace=False
131
+ ):
132
+ """Block-wise FP4 quantization. Power-of-2 scale via bit ops. inplace=True does fused quant+dequant."""
133
+ M = T.symbolic("M")
134
+ fp4_max = 6.0
135
+ fp4_max_inv = 1.0 / fp4_max
136
+ blk_m = 32
137
+ group_size = block_size
138
+ compute_dtype = FP32
139
+ out_dtype = in_dtype if inplace else FP4
140
+
141
+ @T.prim_func
142
+ def fp4_quant_kernel_(
143
+ X: T.Tensor[(M, N), in_dtype],
144
+ Y: T.Tensor[(M, N), out_dtype],
145
+ S: T.Tensor[(M, T.ceildiv(N, group_size)), scale_dtype],
146
+ ):
147
+ with T.Kernel(T.ceildiv(M, blk_m), T.ceildiv(N, group_size), threads=128) as (
148
+ pid_m,
149
+ pid_n,
150
+ ):
151
+ x_shared = T.alloc_shared((blk_m, group_size), in_dtype)
152
+ x_local = T.alloc_fragment((blk_m, group_size), in_dtype)
153
+ amax_local = T.alloc_fragment((blk_m,), compute_dtype)
154
+ s_local = T.alloc_fragment((blk_m,), compute_dtype)
155
+ y_local = T.alloc_fragment((blk_m, group_size), out_dtype)
156
+ y_shared = T.alloc_shared((blk_m, group_size), out_dtype)
157
+
158
+ for _ in T.Pipelined(1, num_stages=2):
159
+ T.copy(X[pid_m * blk_m, pid_n * group_size], x_shared)
160
+ T.copy(x_shared, x_local)
161
+ T.reduce_absmax(x_local, amax_local, dim=1)
162
+ for i in T.Parallel(blk_m):
163
+ amax_local[i] = T.max(amax_local[i], 6 * (2**-126))
164
+ s_local[i] = fast_round_scale(amax_local[i], fp4_max_inv)
165
+ if inplace:
166
+ for i, j in T.Parallel(blk_m, group_size):
167
+ y_local[i, j] = T.Cast(
168
+ out_dtype,
169
+ T.Cast(compute_dtype, T.Cast(FP4, T.clamp(
170
+ x_local[i, j] / s_local[i], -fp4_max, fp4_max
171
+ ))) * s_local[i],
172
+ )
173
+ else:
174
+ for i, j in T.Parallel(blk_m, group_size):
175
+ y_local[i, j] = T.clamp(
176
+ x_local[i, j] / s_local[i], -fp4_max, fp4_max
177
+ )
178
+ for i in T.Parallel(blk_m):
179
+ S[pid_m * blk_m + i, pid_n] = T.Cast(scale_dtype, s_local[i])
180
+ T.copy(y_local, y_shared)
181
+ T.copy(y_shared, Y[pid_m * blk_m, pid_n * group_size])
182
+
183
+ return fp4_quant_kernel_
184
+
185
+
186
+ def fp4_act_quant(
187
+ x: torch.Tensor, block_size: int = 32, inplace: bool = False,
188
+ ) -> torch.Tensor:
189
+ """Block-wise FP4 quantization. inplace=True does fused quant+dequant back to BF16."""
190
+ N = x.size(-1)
191
+ assert N % block_size == 0
192
+ z = x.contiguous()
193
+ y = torch.empty_like(z) if inplace else z.new_empty(*z.shape[:-1], N // 2, dtype=torch.float4_e2m1fn_x2)
194
+ s = z.new_empty(*z.size()[:-1], N // block_size, dtype=torch.float8_e8m0fnu)
195
+ kernel = fp4_quant_kernel(N, block_size, inplace=inplace)
196
+ kernel(z.view(-1, N), y.view(-1, y.size(-1)), s.view(-1, N // block_size))
197
+ if inplace:
198
+ x.copy_(y)
199
+ return x
200
+ return y, s
201
+
202
+
203
+ @tilelang.jit(pass_configs=pass_configs)
204
+ def fp8_gemm_kernel(N, K, out_dtype=BF16, accum_dtype=FP32, scale_dtype=FP32):
205
+ assert out_dtype in [BF16, FP32]
206
+
207
+ M = T.symbolic("M")
208
+ group_size = 128
209
+ block_M = 32
210
+ block_N = 128
211
+ block_K = 128
212
+
213
+ @T.prim_func
214
+ def fp8_gemm_kernel_(
215
+ A: T.Tensor[(M, K), FP8],
216
+ B: T.Tensor[(N, K), FP8],
217
+ C: T.Tensor[(M, N), out_dtype],
218
+ scales_a: T.Tensor[(M, T.ceildiv(K, group_size)), scale_dtype],
219
+ scales_b: T.Tensor[(T.ceildiv(N, group_size), T.ceildiv(K, group_size)), scale_dtype],
220
+ ):
221
+ with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (
222
+ bx,
223
+ by,
224
+ ):
225
+ A_shared = T.alloc_shared((block_M, block_K), FP8)
226
+ B_shared = T.alloc_shared((block_N, block_K), FP8)
227
+ C_shared = T.alloc_shared((block_M, block_N), out_dtype)
228
+ Scale_C_shared = T.alloc_shared((block_M), FP32)
229
+ C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
230
+ C_local_accum = T.alloc_fragment((block_M, block_N), accum_dtype)
231
+
232
+ # Improve L2 Cache
233
+ T.use_swizzle(panel_size=10)
234
+ T.clear(C_local)
235
+ T.clear(C_local_accum)
236
+
237
+ K_iters = T.ceildiv(K, block_K)
238
+ for k in T.Pipelined(K_iters, num_stages=4):
239
+ T.copy(A[by * block_M, k * block_K], A_shared)
240
+ T.copy(B[bx * block_N, k * block_K], B_shared)
241
+ # Cast scales to FP32 for computation; scales_b has one value per block_N group
242
+ Scale_B = T.Cast(FP32, scales_b[bx * block_N // group_size, k])
243
+ for i in T.Parallel(block_M):
244
+ Scale_C_shared[i] = T.Cast(FP32, scales_a[by * block_M + i, k]) * Scale_B
245
+
246
+ T.gemm(A_shared, B_shared, C_local, transpose_B=True)
247
+ # Separate accumulator for scale-corrected results (2x accumulation precision)
248
+ for i, j in T.Parallel(block_M, block_N):
249
+ C_local_accum[i, j] += C_local[i, j] * Scale_C_shared[i]
250
+ T.clear(C_local)
251
+ T.copy(C_local_accum, C_shared)
252
+ T.copy(C_shared, C[by * block_M, bx * block_N])
253
+
254
+ return fp8_gemm_kernel_
255
+
256
+
257
+ def fp8_gemm(
258
+ a: torch.Tensor, a_s: torch.Tensor, b: torch.Tensor, b_s: torch.Tensor,
259
+ scale_dtype: torch.dtype = torch.float32,
260
+ ) -> torch.Tensor:
261
+ """C[M,N] = A[M,K] @ B[N,K]^T with per-128 block FP8 scaling on both A and B."""
262
+ assert a.is_contiguous() and b.is_contiguous(), "Input tensors must be contiguous"
263
+ assert a_s.is_contiguous() and b_s.is_contiguous(), (
264
+ "Scaling factor tensors must be contiguous"
265
+ )
266
+ tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32
267
+ K = a.size(-1)
268
+ M = a.numel() // K
269
+ N = b.size(0)
270
+ c = a.new_empty(*a.size()[:-1], N, dtype=torch.get_default_dtype())
271
+ kernel = fp8_gemm_kernel(N, K, scale_dtype=tl_dtype)
272
+ kernel(a.view(M, K), b, c.view(M, N), a_s.view(M, -1), b_s)
273
+ return c
274
+
275
+
276
+ @tilelang.jit(pass_configs=pass_configs)
277
+ def sparse_attn_kernel(h: int, d: int, scale=None):
278
+ """Sparse multi-head attention via index gathering + online softmax (FlashAttention-style).
279
+ For each (batch, seq_pos), gathers top-k KV positions by index, computes attention
280
+ with numerically stable running max/sum, and includes a learnable attn_sink bias."""
281
+ b = T.symbolic("b")
282
+ m = T.symbolic("m")
283
+ n = T.symbolic("n")
284
+ topk = T.symbolic("topk")
285
+ if scale is None:
286
+ scale = (1.0 / d) ** 0.5
287
+
288
+ num_stages = 2
289
+ threads = 256
290
+ block = 64
291
+ num_blocks = tilelang.cdiv(topk, block)
292
+
293
+ @T.prim_func
294
+ def sparse_attn_kernel_(
295
+ q: T.Tensor[(b, m, h, d), BF16],
296
+ kv: T.Tensor[(b, n, d), BF16],
297
+ o: T.Tensor[(b, m, h, d), BF16],
298
+ attn_sink: T.Tensor[(h,), FP32],
299
+ topk_idxs: T.Tensor[(b, m, topk), INT32],
300
+ ):
301
+ with T.Kernel(m, b, threads=threads) as (bx, by):
302
+ q_shared = T.alloc_shared((h, d), BF16)
303
+ kv_shared = T.alloc_shared((block, d), BF16)
304
+ o_shared = T.alloc_shared((h, d), BF16)
305
+ acc_s_cast = T.alloc_shared((h, block), BF16)
306
+
307
+ idxs = T.alloc_fragment(block, INT32)
308
+ acc_s = T.alloc_fragment((h, block), FP32)
309
+ acc_o = T.alloc_fragment((h, d), FP32)
310
+ scores_max = T.alloc_fragment(h, FP32)
311
+ scores_max_prev = T.alloc_fragment(h, FP32)
312
+ scores_scale = T.alloc_fragment(h, FP32)
313
+ scores_sum = T.alloc_fragment(h, FP32)
314
+ sum_exp = T.alloc_fragment(h, FP32)
315
+
316
+ T.clear(acc_o)
317
+ T.clear(sum_exp)
318
+ T.fill(scores_max, -T.infinity(FP32))
319
+ T.copy(q[by, bx, :, :], q_shared)
320
+
321
+ for t in T.Pipelined(num_blocks, num_stages=num_stages):
322
+ for i in T.Parallel(block):
323
+ idxs[i] = T.if_then_else(t * block + i < topk, topk_idxs[by, bx, t * block + i], -1)
324
+ for i, j in T.Parallel(block, d):
325
+ kv_shared[i, j] = T.if_then_else(idxs[i] != -1, kv[by, idxs[i], j], 0)
326
+ for i, j in T.Parallel(h, block):
327
+ acc_s[i, j] = T.if_then_else(idxs[j] != -1, 0, -T.infinity(FP32))
328
+ T.gemm(q_shared, kv_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow)
329
+ for i, j in T.Parallel(h, block):
330
+ acc_s[i, j] *= scale
331
+ T.copy(scores_max, scores_max_prev)
332
+ T.reduce_max(acc_s, scores_max, dim=1, clear=False)
333
+ for i in T.Parallel(h):
334
+ scores_scale[i] = T.exp(scores_max_prev[i] - scores_max[i])
335
+ for i, j in T.Parallel(h, block):
336
+ acc_s[i, j] = T.exp(acc_s[i, j] - scores_max[i])
337
+ T.reduce_sum(acc_s, scores_sum, dim=1)
338
+ for i in T.Parallel(h):
339
+ sum_exp[i] = sum_exp[i] * scores_scale[i] + scores_sum[i]
340
+ T.copy(acc_s, acc_s_cast)
341
+ for i, j in T.Parallel(h, d):
342
+ acc_o[i, j] *= scores_scale[i]
343
+ T.gemm(acc_s_cast, kv_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)
344
+
345
+ for i in T.Parallel(h):
346
+ sum_exp[i] += T.exp(attn_sink[i] - scores_max[i])
347
+ for i, j in T.Parallel(h, d):
348
+ acc_o[i, j] /= sum_exp[i]
349
+ T.copy(acc_o, o_shared)
350
+ T.copy(o_shared, o[by, bx, :, :])
351
+
352
+ return sparse_attn_kernel_
353
+
354
+
355
+ def sparse_attn(
356
+ q: torch.Tensor, kv: torch.Tensor, attn_sink: torch.Tensor, topk_idxs: torch.Tensor, softmax_scale: float
357
+ ) -> torch.Tensor:
358
+ b, s, h, d = q.size()
359
+ # Pad heads to 16 for kernel efficiency (stripped after)
360
+ if h < 16:
361
+ q = torch.cat([q, q.new_zeros(b, s, 16 - h, d)], dim=2)
362
+ attn_sink = torch.cat([attn_sink, attn_sink.new_zeros(16 - h)])
363
+ o = torch.empty_like(q)
364
+ kernel = sparse_attn_kernel(q.size(2), d, softmax_scale)
365
+ kernel(q, kv, o, attn_sink, topk_idxs)
366
+ if h < 16:
367
+ o = o.narrow(2, 0, h).contiguous()
368
+ return o
369
+
370
+
371
+ @tilelang.jit(pass_configs=pass_configs)
372
+ def hc_split_sinkhorn_kernel(hc: int, sinkhorn_iters: int, eps: float):
373
+ n = T.symbolic("n")
374
+ mix_hc = (2 + hc) * hc
375
+ threads = 64
376
+
377
+ @T.prim_func
378
+ def hc_split_sinkhorn_kernel_(
379
+ mixes: T.Tensor[(n, mix_hc), FP32],
380
+ hc_scale: T.Tensor[(3,), FP32],
381
+ hc_base: T.Tensor[(mix_hc,), FP32],
382
+ pre: T.Tensor[(n, hc), FP32],
383
+ post: T.Tensor[(n, hc), FP32],
384
+ comb: T.Tensor[(n, hc, hc), FP32],
385
+ ):
386
+ with T.Kernel(n, threads=threads) as i:
387
+ mixes_shared = T.alloc_shared(mix_hc, FP32)
388
+ comb_frag = T.alloc_fragment((hc, hc), FP32)
389
+ T.copy(mixes[i, :], mixes_shared)
390
+
391
+ for j in T.Parallel(hc):
392
+ pre[i, j] = T.sigmoid(mixes_shared[j] * hc_scale[0] + hc_base[j]) + eps
393
+ for j in T.Parallel(hc):
394
+ post[i, j] = 2 * T.sigmoid(mixes_shared[j + hc] * hc_scale[1] + hc_base[j + hc])
395
+ for j, k in T.Parallel(hc, hc):
396
+ comb_frag[j, k] = mixes_shared[j * hc + k + hc * 2] * hc_scale[2] + hc_base[j * hc + k + hc * 2]
397
+
398
+ row_sum = T.alloc_fragment(hc, FP32)
399
+ col_sum = T.alloc_fragment(hc, FP32)
400
+
401
+ # comb = comb.softmax(-1) + eps
402
+ row_max = T.alloc_fragment(hc, FP32)
403
+ T.reduce_max(comb_frag, row_max, dim=1)
404
+ for j, k in T.Parallel(hc, hc):
405
+ comb_frag[j, k] = T.exp(comb_frag[j, k] - row_max[j])
406
+ T.reduce_sum(comb_frag, row_sum, dim=1)
407
+ for j, k in T.Parallel(hc, hc):
408
+ comb_frag[j, k] = comb_frag[j, k] / row_sum[j] + eps
409
+
410
+ # comb = comb / (comb.sum(-2) + eps)
411
+ T.reduce_sum(comb_frag, col_sum, dim=0)
412
+ for j, k in T.Parallel(hc, hc):
413
+ comb_frag[j, k] = comb_frag[j, k] / (col_sum[k] + eps)
414
+
415
+ for _ in T.serial(sinkhorn_iters - 1):
416
+ # comb = comb / (comb.sum(-1) + eps)
417
+ T.reduce_sum(comb_frag, row_sum, dim=1)
418
+ for j, k in T.Parallel(hc, hc):
419
+ comb_frag[j, k] = comb_frag[j, k] / (row_sum[j] + eps)
420
+ # comb = comb / (comb.sum(-2) + eps)
421
+ T.reduce_sum(comb_frag, col_sum, dim=0)
422
+ for j, k in T.Parallel(hc, hc):
423
+ comb_frag[j, k] = comb_frag[j, k] / (col_sum[k] + eps)
424
+
425
+ T.copy(comb_frag, comb[i, :, :])
426
+
427
+ return hc_split_sinkhorn_kernel_
428
+
429
+
430
+ def hc_split_sinkhorn(mixes: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor, hc_mult: int = 4, sinkhorn_iters: int = 20, eps: float = 1e-6):
431
+ b, s, _ = mixes.size()
432
+ pre = mixes.new_empty(b, s, hc_mult)
433
+ post = mixes.new_empty(b, s, hc_mult)
434
+ comb = mixes.new_empty(b, s, hc_mult, hc_mult)
435
+ kernel = hc_split_sinkhorn_kernel(hc_mult, sinkhorn_iters, eps)
436
+ kernel(mixes.view(-1, (2 + hc_mult) * hc_mult), hc_scale, hc_base,
437
+ pre.view(-1, hc_mult), post.view(-1, hc_mult), comb.view(-1, hc_mult, hc_mult))
438
+ return pre, post, comb
439
+
440
+
441
+ @tilelang.jit(pass_configs=pass_configs)
442
+ def fp4_gemm_kernel(N, K, out_dtype=BF16, accum_dtype=FP32, scale_dtype=FP32):
443
+ """FP8 act x FP4 weight GEMM kernel.
444
+
445
+ C[M, N] = A_fp8[M, K] @ B_fp4[N, K]^T
446
+
447
+ Act: 1x128 quant on K (reduce dim), FP8 with configurable scale dtype
448
+ Weight: 1x32 quant on K (reduce dim), FP4 with E8M0 scale
449
+
450
+ B is stored as [N, K//2] in float4_e2m1fn_x2, logical [N, K] in fp4.
451
+ The FP4 values are packed along the K (last) dimension.
452
+
453
+ Strategy: load FP4 sub-blocks of size [block_N, sub_K] (sub_K=32),
454
+ cast FP4 to FP8 via float, then do FP8xFP8 GEMM.
455
+ Apply act scale (per 128 on K) and weight scale (per 32 on K) to the accumulator.
456
+ """
457
+ M = T.symbolic("M")
458
+ act_group_size = 128
459
+ weight_group_size = 32
460
+ block_M = 32
461
+ block_N = 128
462
+ block_K = 32 # matches weight_group_size for simple scale handling
463
+ n_sub = act_group_size // block_K # 4 sub-blocks per act scale group
464
+
465
+ @T.prim_func
466
+ def fp4_gemm_kernel_(
467
+ A: T.Tensor[(M, K), FP8],
468
+ B: T.Tensor[(N, K), FP4],
469
+ C: T.Tensor[(M, N), out_dtype],
470
+ scales_a: T.Tensor[(M, T.ceildiv(K, act_group_size)), scale_dtype],
471
+ scales_b: T.Tensor[(N, T.ceildiv(K, weight_group_size)), scale_dtype],
472
+ ):
473
+ with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (
474
+ bx,
475
+ by,
476
+ ):
477
+ A_shared = T.alloc_shared((block_M, block_K), FP8)
478
+ B_fp4_shared = T.alloc_shared((block_N, block_K), FP4)
479
+ B_shared = T.alloc_shared((block_N, block_K), FP8)
480
+ C_shared = T.alloc_shared((block_M, block_N), out_dtype)
481
+ C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
482
+ C_local_accum = T.alloc_fragment((block_M, block_N), accum_dtype)
483
+ scale_a_frag = T.alloc_fragment((block_M,), FP32)
484
+ scale_b_frag = T.alloc_fragment((block_N,), FP32)
485
+
486
+ T.use_swizzle(panel_size=10)
487
+ T.clear(C_local)
488
+ T.clear(C_local_accum)
489
+
490
+ K_iters = T.ceildiv(K, block_K)
491
+ for k in T.Pipelined(K_iters, num_stages=2):
492
+ T.copy(A[by * block_M, k * block_K], A_shared)
493
+ T.copy(B[bx * block_N, k * block_K], B_fp4_shared)
494
+ # FP4->FP8 cast must go through FP32 to avoid ambiguous C++ overload
495
+ for i, j in T.Parallel(block_N, block_K):
496
+ B_shared[i, j] = T.Cast(FP8, T.Cast(FP32, B_fp4_shared[i, j]))
497
+
498
+ # Weight scale: per 32 on K, indexed by k (each k is one block_K=32)
499
+ for i in T.Parallel(block_N):
500
+ scale_b_frag[i] = T.Cast(FP32, scales_b[bx * block_N + i, k])
501
+
502
+ # Act scale: per 128 on K, indexed by k // 4
503
+ for i in T.Parallel(block_M):
504
+ scale_a_frag[i] = T.Cast(FP32, scales_a[by * block_M + i, k // n_sub])
505
+
506
+ T.gemm(A_shared, B_shared, C_local, transpose_B=True)
507
+
508
+ for i, j in T.Parallel(block_M, block_N):
509
+ C_local_accum[i, j] += C_local[i, j] * scale_a_frag[i] * scale_b_frag[j]
510
+ T.clear(C_local)
511
+
512
+ T.copy(C_local_accum, C_shared)
513
+ T.copy(C_shared, C[by * block_M, bx * block_N])
514
+
515
+ return fp4_gemm_kernel_
516
+
517
+
518
+ def fp4_gemm(
519
+ a: torch.Tensor, a_s: torch.Tensor, b: torch.Tensor, b_s: torch.Tensor,
520
+ scale_dtype: torch.dtype = torch.float32,
521
+ ) -> torch.Tensor:
522
+ """C[M,N] = A_fp8[M,K] @ B_fp4[N,K]^T.
523
+ A has per-128 act scale; B has per-32 E8M0 weight scale.
524
+ B is stored as [N, K//2] in float4_e2m1fn_x2 (2 FP4 values per byte, packed along K)."""
525
+ assert a.is_contiguous() and b.is_contiguous(), "Input tensors must be contiguous"
526
+ assert a_s.is_contiguous() and b_s.is_contiguous(), (
527
+ "Scaling factor tensors must be contiguous"
528
+ )
529
+ tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32
530
+ K = a.size(-1)
531
+ M = a.numel() // K
532
+ N = b.size(0)
533
+ c = a.new_empty(*a.size()[:-1], N, dtype=torch.get_default_dtype())
534
+ kernel = fp4_gemm_kernel(N, K, scale_dtype=tl_dtype)
535
+ kernel(a.view(M, K), b, c.view(M, N), a_s.view(M, -1), b_s)
536
+ return c
inference/model.py ADDED
@@ -0,0 +1,961 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from dataclasses import dataclass
3
+ from typing import Tuple, Optional, Literal
4
+ from functools import lru_cache
5
+ from contextlib import contextmanager
6
+
7
+ import torch
8
+ from torch import nn
9
+ import torch.nn.functional as F
10
+ import torch.distributed as dist
11
+
12
+ from kernel import act_quant, fp4_act_quant, fp8_gemm, fp4_gemm, sparse_attn, hc_split_sinkhorn
13
+
14
+
15
+ world_size = 1
16
+ rank = 0
17
+ block_size = 128
18
+ fp4_block_size = 32
19
+ default_dtype = torch.bfloat16
20
+ scale_fmt = None
21
+ scale_dtype = torch.float32
22
+
23
+
24
+ @contextmanager
25
+ def set_dtype(dtype):
26
+ """Temporarily override torch default dtype, restoring it on exit (even if an exception occurs)."""
27
+ prev = torch.get_default_dtype()
28
+ torch.set_default_dtype(dtype)
29
+ try:
30
+ yield
31
+ finally:
32
+ torch.set_default_dtype(prev)
33
+
34
+ @dataclass
35
+ class ModelArgs:
36
+ """Model hyperparameters. Field names match the config JSON keys."""
37
+ max_batch_size: int = 4
38
+ max_seq_len: int = 4096
39
+ temperature: float = 1
40
+ dtype: Literal["bf16", "fp8"] = "fp8"
41
+ scale_fmt: Literal[None, "ue8m0"] = "ue8m0"
42
+ expert_dtype: Literal[None, "fp4"] = None
43
+ scale_dtype: Literal["fp32", "fp8"] = "fp8"
44
+ vocab_size: int = 129280
45
+ dim: int = 4096
46
+ moe_inter_dim: int = 4096
47
+ n_layers: int = 7
48
+ n_hash_layers: int = 0
49
+ n_mtp_layers: int = 1
50
+ n_heads: int = 64
51
+ # moe
52
+ n_routed_experts: int = 8
53
+ n_shared_experts: int = 1
54
+ n_activated_experts: int = 2
55
+ score_func: Literal["softmax", "sigmoid", "sqrtsoftplus"] = "sqrtsoftplus"
56
+ route_scale: float = 1.
57
+ swiglu_limit: float = 0.
58
+ # mqa
59
+ q_lora_rank: int = 1024
60
+ head_dim: int = 512
61
+ rope_head_dim: int = 64
62
+ norm_eps: float = 1e-6
63
+ o_groups: int = 8
64
+ o_lora_rank: int = 1024
65
+ window_size: int = 128
66
+ compress_ratios: Tuple[int] = (0, 0, 4, 128, 4, 128, 4, 0)
67
+ # yarn
68
+ compress_rope_theta: float = 40000.0
69
+ original_seq_len: int = 0
70
+ rope_theta: float = 10000.0
71
+ rope_factor: float = 40
72
+ beta_fast: int = 32
73
+ beta_slow: int = 1
74
+ # index
75
+ index_n_heads: int = 64
76
+ index_head_dim: int = 128
77
+ index_topk: int = 512
78
+ # hc
79
+ hc_mult: int = 4
80
+ hc_sinkhorn_iters: int = 20
81
+ hc_eps: float = 1e-6
82
+ # dspark
83
+ dspark_block_size: int = 0
84
+ dspark_noise_token_id: int = 0
85
+ dspark_target_layer_ids: Tuple[int] = tuple()
86
+ dspark_markov_rank: int = 256
87
+
88
+
89
+ class ParallelEmbedding(nn.Module):
90
+ """Embedding sharded along the vocab dimension. Each rank holds vocab_size // world_size rows.
91
+ Out-of-range indices are zero-masked before all_reduce to combine partial embeddings."""
92
+ def __init__(self, vocab_size: int, dim: int):
93
+ super().__init__()
94
+ self.vocab_size = vocab_size
95
+ self.dim = dim
96
+ assert vocab_size % world_size == 0, f"Vocabulary size must be divisible by world size (world_size={world_size})"
97
+ self.part_vocab_size = (vocab_size // world_size)
98
+ self.vocab_start_idx = rank * self.part_vocab_size
99
+ self.vocab_end_idx = self.vocab_start_idx + self.part_vocab_size
100
+ self.weight = nn.Parameter(torch.empty(self.part_vocab_size, self.dim))
101
+
102
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
103
+ if world_size > 1:
104
+ mask = (x < self.vocab_start_idx) | (x >= self.vocab_end_idx)
105
+ x = x - self.vocab_start_idx
106
+ x[mask] = 0
107
+ y = F.embedding(x, self.weight)
108
+ if world_size > 1:
109
+ y[mask] = 0
110
+ dist.all_reduce(y)
111
+ return y
112
+
113
+
114
+ def linear(x: torch.Tensor, weight: torch.Tensor, bias: Optional[torch.Tensor] = None) -> torch.Tensor:
115
+ """Dispatches to fp4_gemm / fp8_gemm / F.linear based on weight dtype.
116
+ For quantized weights, x is first quantized to FP8 via act_quant."""
117
+ assert bias is None
118
+
119
+ if weight.dtype == torch.float4_e2m1fn_x2:
120
+ x, s = act_quant(x, block_size, scale_fmt, scale_dtype)
121
+ return fp4_gemm(x, s, weight, weight.scale, scale_dtype)
122
+ elif weight.dtype == torch.float8_e4m3fn:
123
+ x, s = act_quant(x, block_size, scale_fmt, scale_dtype)
124
+ return fp8_gemm(x, s, weight, weight.scale, scale_dtype)
125
+ else:
126
+ return F.linear(x, weight)
127
+
128
+
129
+ class Linear(nn.Module):
130
+ """Linear layer supporting BF16, FP8, and FP4 weight formats with per-block scaling."""
131
+
132
+ def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None):
133
+ super().__init__()
134
+ self.in_features = in_features
135
+ self.out_features = out_features
136
+ dtype = dtype or default_dtype
137
+ if dtype == torch.float4_e2m1fn_x2:
138
+ # FP4: weight is [out, in//2] in float4_e2m1fn_x2, logically [out, in] in fp4
139
+ # Scale is [out, in//32] in float8_e8m0fnu (1 scale per 32 fp4 elements along K)
140
+ self.weight = nn.Parameter(torch.empty(out_features, in_features // 2, dtype=torch.float4_e2m1fn_x2))
141
+ scale_out_features = out_features
142
+ scale_in_features = in_features // fp4_block_size
143
+ self.weight.scale = self.scale = nn.Parameter(torch.empty(scale_out_features, scale_in_features, dtype=torch.float8_e8m0fnu))
144
+ elif dtype == torch.float8_e4m3fn:
145
+ self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype))
146
+ scale_out_features = (out_features + block_size - 1) // block_size
147
+ scale_in_features = (in_features + block_size - 1) // block_size
148
+ self.weight.scale = self.scale = nn.Parameter(torch.empty(scale_out_features, scale_in_features, dtype=torch.float8_e8m0fnu))
149
+ else:
150
+ self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype))
151
+ self.register_parameter("scale", None)
152
+ if bias:
153
+ self.bias = nn.Parameter(torch.empty(out_features))
154
+ else:
155
+ self.register_parameter("bias", None)
156
+
157
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
158
+ return linear(x, self.weight, self.bias)
159
+
160
+
161
+ class ColumnParallelLinear(Linear):
162
+ """Shards output dim across TP ranks. No all-reduce needed on output."""
163
+ def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None):
164
+ assert out_features % world_size == 0, f"Output features must be divisible by world size (world_size={world_size})"
165
+ self.part_out_features = out_features // world_size
166
+ super().__init__(in_features, self.part_out_features, bias, dtype)
167
+
168
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
169
+ return linear(x, self.weight, self.bias)
170
+
171
+
172
+ class RowParallelLinear(Linear):
173
+ """Shards input dim across TP ranks. All-reduce on output to sum partial results."""
174
+ def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None):
175
+ assert in_features % world_size == 0, f"Input features must be divisible by world size (world_size={world_size})"
176
+ self.part_in_features = in_features // world_size
177
+ super().__init__(self.part_in_features, out_features, bias, dtype)
178
+
179
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
180
+ y = linear(x, self.weight, None)
181
+ if world_size > 1:
182
+ y = y.float()
183
+ dist.all_reduce(y)
184
+ if self.bias is not None:
185
+ y += self.bias
186
+ return y.type_as(x)
187
+
188
+
189
+ class RMSNorm(nn.Module):
190
+ def __init__(self, dim: int, eps: float = 1e-6):
191
+ super().__init__()
192
+ self.dim = dim
193
+ self.eps = eps
194
+ # rmsnorm in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for convenient.
195
+ self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32))
196
+
197
+ def forward(self, x: torch.Tensor):
198
+ dtype = x.dtype
199
+ x = x.float()
200
+ var = x.square().mean(-1, keepdim=True)
201
+ x = x * torch.rsqrt(var + self.eps)
202
+ return (self.weight * x).to(dtype)
203
+
204
+
205
+ @lru_cache(2)
206
+ def precompute_freqs_cis(dim, seqlen, original_seq_len, base, factor, beta_fast, beta_slow) -> torch.Tensor:
207
+ """Precomputes complex exponentials for rotary embeddings with YaRN scaling.
208
+ When original_seq_len > 0, applies frequency interpolation with a smooth
209
+ linear ramp between beta_fast and beta_slow correction ranges."""
210
+
211
+ def find_correction_dim(num_rotations, dim, base, max_seq_len):
212
+ return dim * math.log(max_seq_len / (num_rotations * 2 * math.pi)) / (2 * math.log(base))
213
+
214
+ def find_correction_range(low_rot, high_rot, dim, base, max_seq_len):
215
+ low = math.floor(find_correction_dim(low_rot, dim, base, max_seq_len))
216
+ high = math.ceil(find_correction_dim(high_rot, dim, base, max_seq_len))
217
+ return max(low, 0), min(high, dim-1)
218
+
219
+ def linear_ramp_factor(min, max, dim):
220
+ if min == max:
221
+ max += 0.001
222
+ linear_func = (torch.arange(dim, dtype=torch.float32) - min) / (max - min)
223
+ ramp_func = torch.clamp(linear_func, 0, 1)
224
+ return ramp_func
225
+
226
+ freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
227
+ if original_seq_len > 0:
228
+ low, high = find_correction_range(beta_fast, beta_slow, dim, base, original_seq_len)
229
+ smooth = 1 - linear_ramp_factor(low, high, dim // 2)
230
+ freqs = freqs / factor * (1 - smooth) + freqs * smooth
231
+
232
+ t = torch.arange(seqlen)
233
+ freqs = torch.outer(t, freqs)
234
+ freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
235
+ return freqs_cis
236
+
237
+
238
+ def apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor, inverse: bool = False) -> torch.Tensor:
239
+ """Applies rotary positional embeddings in-place. Uses conjugate for inverse (de-rotation)."""
240
+ y = x
241
+ x = torch.view_as_complex(x.float().unflatten(-1, (-1, 2)))
242
+ if inverse:
243
+ freqs_cis = freqs_cis.conj()
244
+ if x.ndim == 3:
245
+ freqs_cis = freqs_cis.view(1, x.size(1), x.size(-1))
246
+ else:
247
+ freqs_cis = freqs_cis.view(1, x.size(1), 1, x.size(-1))
248
+ x = torch.view_as_real(x * freqs_cis).flatten(-2)
249
+ y.copy_(x)
250
+ return y
251
+
252
+
253
+ def rotate_activation(x: torch.Tensor) -> torch.Tensor:
254
+ """Applies randomized Hadamard rotation to spread information across dims before FP8 quant."""
255
+ assert x.dtype == torch.bfloat16
256
+ from fast_hadamard_transform import hadamard_transform
257
+ return hadamard_transform(x, scale=x.size(-1) ** -0.5)
258
+
259
+
260
+ @lru_cache(1)
261
+ def get_window_topk_idxs(window_size: int, bsz: int, seqlen: int, start_pos: int):
262
+ if start_pos >= window_size - 1:
263
+ start_pos %= window_size
264
+ matrix = torch.cat([torch.arange(start_pos + 1, window_size), torch.arange(0, start_pos + 1)], dim=0)
265
+ elif start_pos > 0:
266
+ matrix = F.pad(torch.arange(start_pos + 1), (0, window_size - start_pos - 1), value=-1)
267
+ else:
268
+ base = torch.arange(seqlen).unsqueeze(1)
269
+ matrix = (base - window_size + 1).clamp(0) + torch.arange(min(seqlen, window_size))
270
+ matrix = torch.where(matrix > base, -1, matrix)
271
+ return matrix.int().unsqueeze(0).expand(bsz, -1, -1).contiguous()
272
+
273
+
274
+ @lru_cache(2)
275
+ def get_compress_topk_idxs(ratio: int, bsz: int, seqlen: int, start_pos: int, offset: int):
276
+ if start_pos > 0:
277
+ matrix = torch.arange(0, (start_pos + 1) // ratio) + offset
278
+ else:
279
+ matrix = torch.arange(seqlen // ratio).repeat(seqlen, 1)
280
+ mask = matrix >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio
281
+ matrix = torch.where(mask, -1, matrix + offset)
282
+ return matrix.int().unsqueeze(0).expand(bsz, -1, -1).contiguous()
283
+
284
+
285
+ class Compressor(nn.Module):
286
+ """Compresses KV cache via learned gated pooling over `compress_ratio` consecutive tokens.
287
+ When overlap=True (ratio==4), uses overlapping windows for smoother compression boundaries."""
288
+
289
+ def __init__(self, args: ModelArgs, compress_ratio: int = 4, head_dim: int = 512, rotate: bool = False):
290
+ super().__init__()
291
+ self.dim = args.dim
292
+ self.head_dim = head_dim
293
+ self.rope_head_dim = args.rope_head_dim
294
+ self.nope_head_dim = head_dim - args.rope_head_dim
295
+ self.compress_ratio = compress_ratio
296
+ self.overlap = compress_ratio == 4
297
+ self.rotate = rotate
298
+ coff = 1 + self.overlap
299
+
300
+ self.ape = nn.Parameter(torch.empty(compress_ratio, coff * self.head_dim, dtype=torch.float32))
301
+ # wkv and wgate in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for convenient.
302
+ # When overlap, the first half of dims is for overlapping compression, second half for normal.
303
+ self.wkv = Linear(self.dim, coff * self.head_dim, dtype=torch.float32)
304
+ self.wgate = Linear(self.dim, coff * self.head_dim, dtype=torch.float32)
305
+ self.norm = RMSNorm(self.head_dim, args.norm_eps)
306
+ self.kv_cache: torch.Tensor = None # assigned lazily from Attention.kv_cache
307
+ # State buffers for decode-phase incremental compression.
308
+ # With overlap: state[:, :ratio] = overlapping window, state[:, ratio:] = current window.
309
+ self.register_buffer("kv_state", torch.zeros(args.max_batch_size, coff * compress_ratio, coff * self.head_dim, dtype=torch.float32), persistent=False)
310
+ self.register_buffer("score_state", torch.full((args.max_batch_size, coff * compress_ratio, coff * self.head_dim), float("-inf"), dtype=torch.float32), persistent=False)
311
+ self.freqs_cis: torch.Tensor = None
312
+
313
+ def overlap_transform(self, tensor: torch.Tensor, value=0):
314
+ # tensor: [b,s,r,2d]
315
+ b, s, _, _ = tensor.size()
316
+ ratio, d = self.compress_ratio, self.head_dim
317
+ new_tensor = tensor.new_full((b, s, 2 * ratio, d), value)
318
+ new_tensor[:, :, ratio:] = tensor[:, :, :, d:]
319
+ new_tensor[:, 1:, :ratio] = tensor[:, :-1, :, :d]
320
+ return new_tensor
321
+
322
+ def forward(self, x: torch.Tensor, start_pos: int):
323
+ assert self.kv_cache is not None
324
+ bsz, seqlen, _ = x.size()
325
+ ratio, overlap, d, rd = self.compress_ratio, self.overlap, self.head_dim, self.rope_head_dim
326
+ dtype = x.dtype
327
+ # compression need fp32
328
+ x = x.float()
329
+ kv = self.wkv(x)
330
+ score = self.wgate(x)
331
+ if start_pos == 0:
332
+ should_compress = seqlen >= ratio
333
+ remainder = seqlen % ratio
334
+ cutoff = seqlen - remainder
335
+ offset = ratio if overlap else 0
336
+ if overlap and cutoff >= ratio:
337
+ self.kv_state[:bsz, :ratio] = kv[:, cutoff-ratio : cutoff]
338
+ self.score_state[:bsz, :ratio] = score[:, cutoff-ratio : cutoff] + self.ape
339
+ if remainder > 0:
340
+ kv, self.kv_state[:bsz, offset : offset+remainder] = kv.split([cutoff, remainder], dim=1)
341
+ self.score_state[:bsz, offset : offset+remainder] = score[:, cutoff:] + self.ape[:remainder]
342
+ score = score[:, :cutoff]
343
+ kv = kv.unflatten(1, (-1, ratio))
344
+ score = score.unflatten(1, (-1, ratio)) + self.ape
345
+ if overlap:
346
+ kv = self.overlap_transform(kv, 0)
347
+ score = self.overlap_transform(score, float("-inf"))
348
+ kv = (kv * score.softmax(dim=2)).sum(dim=2)
349
+ else:
350
+ should_compress = (start_pos + 1) % self.compress_ratio == 0
351
+ score += self.ape[start_pos % ratio]
352
+ if overlap:
353
+ self.kv_state[:bsz, ratio + start_pos % ratio] = kv.squeeze(1)
354
+ self.score_state[:bsz, ratio + start_pos % ratio] = score.squeeze(1)
355
+ if should_compress:
356
+ kv_state = torch.cat([self.kv_state[:bsz, :ratio, :d], self.kv_state[:bsz, ratio:, d:]], dim=1)
357
+ score_state = torch.cat([self.score_state[:bsz, :ratio, :d], self.score_state[:bsz, ratio:, d:]], dim=1)
358
+ kv = (kv_state * score_state.softmax(dim=1)).sum(dim=1, keepdim=True)
359
+ self.kv_state[:bsz, :ratio] = self.kv_state[:bsz, ratio:]
360
+ self.score_state[:bsz, :ratio] = self.score_state[:bsz, ratio:]
361
+ else:
362
+ self.kv_state[:bsz, start_pos % ratio] = kv.squeeze(1)
363
+ self.score_state[:bsz, start_pos % ratio] = score.squeeze(1)
364
+ if should_compress:
365
+ kv = (self.kv_state[:bsz] * self.score_state[:bsz].softmax(dim=1)).sum(dim=1, keepdim=True)
366
+ if not should_compress:
367
+ return
368
+ kv = self.norm(kv.to(dtype))
369
+ if start_pos == 0:
370
+ freqs_cis = self.freqs_cis[:cutoff:ratio]
371
+ else:
372
+ freqs_cis = self.freqs_cis[start_pos + 1 - self.compress_ratio].unsqueeze(0)
373
+ apply_rotary_emb(kv[..., -rd:], freqs_cis)
374
+ if self.rotate:
375
+ kv = rotate_activation(kv)
376
+ fp4_act_quant(kv, fp4_block_size, True)
377
+ else:
378
+ act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
379
+ if start_pos == 0:
380
+ self.kv_cache[:bsz, :seqlen // ratio] = kv
381
+ else:
382
+ self.kv_cache[:bsz, start_pos // ratio] = kv.squeeze(1)
383
+ return kv
384
+
385
+
386
+ class Indexer(torch.nn.Module):
387
+ """Selects top-k compressed KV positions for sparse attention via learned scoring.
388
+ Has its own Compressor (with Hadamard rotation) to build compressed KV for scoring."""
389
+
390
+ def __init__(self, args: ModelArgs, compress_ratio: int = 4):
391
+ super().__init__()
392
+ self.dim = args.dim
393
+ self.n_heads = args.index_n_heads
394
+ self.n_local_heads = args.index_n_heads // world_size
395
+ self.head_dim = args.index_head_dim
396
+ self.rope_head_dim = args.rope_head_dim
397
+ self.index_topk = args.index_topk
398
+ self.q_lora_rank = args.q_lora_rank
399
+ self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.head_dim)
400
+ self.weights_proj = ColumnParallelLinear(self.dim, self.n_heads, dtype=torch.bfloat16)
401
+ self.softmax_scale = self.head_dim ** -0.5
402
+ self.compress_ratio = compress_ratio
403
+
404
+ self.compressor = Compressor(args, compress_ratio, self.head_dim, True)
405
+ self.register_buffer("kv_cache", torch.zeros(args.max_batch_size, args.max_seq_len // compress_ratio, self.head_dim), persistent=False)
406
+ self.freqs_cis = None
407
+
408
+ def forward(self, x: torch.Tensor, qr: torch.Tensor, start_pos: int, offset: int):
409
+ bsz, seqlen, _ = x.size()
410
+ freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen]
411
+ ratio = self.compress_ratio
412
+ rd = self.rope_head_dim
413
+ end_pos = start_pos + seqlen
414
+ if self.compressor.kv_cache is None:
415
+ self.compressor.kv_cache = self.kv_cache
416
+ self.compressor.freqs_cis = self.freqs_cis
417
+ q = self.wq_b(qr)
418
+ q = q.unflatten(-1, (self.n_local_heads, self.head_dim))
419
+ apply_rotary_emb(q[..., -rd:], freqs_cis)
420
+ q = rotate_activation(q)
421
+ # use fp4 simulation for q and kv in indexer
422
+ fp4_act_quant(q, fp4_block_size, True)
423
+ self.compressor(x, start_pos)
424
+ weights = self.weights_proj(x) * (self.softmax_scale * self.n_heads ** -0.5)
425
+ # We performed QAT here, kv could also use fp8 format, though current implementation uses bf16
426
+ index_score = torch.einsum("bshd,btd->bsht", q, self.kv_cache[:bsz, :end_pos // ratio])
427
+ index_score = (index_score.relu_() * weights.unsqueeze(-1)).sum(dim=2)
428
+ if world_size > 1:
429
+ dist.all_reduce(index_score)
430
+ if start_pos == 0:
431
+ mask = torch.arange(seqlen // ratio).repeat(seqlen, 1) >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio
432
+ index_score += torch.where(mask, float("-inf"), 0)
433
+ topk_idxs = index_score.topk(min(self.index_topk, end_pos // ratio), dim=-1)[1]
434
+ if start_pos == 0:
435
+ mask = topk_idxs >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio
436
+ topk_idxs = torch.where(mask, -1, topk_idxs + offset)
437
+ else:
438
+ topk_idxs += offset
439
+ return topk_idxs
440
+
441
+
442
+ class Attention(nn.Module):
443
+ """Multi-head Latent Attention (MLA) with sliding window + optional KV compression.
444
+ Uses low-rank Q projection (wq_a -> q_norm -> wq_b) and grouped low-rank O projection."""
445
+ def __init__(self, layer_id: int, args: ModelArgs):
446
+ super().__init__()
447
+ self.layer_id = layer_id
448
+ self.dim = args.dim
449
+ self.n_heads = args.n_heads
450
+ self.n_local_heads = args.n_heads // world_size
451
+ self.q_lora_rank = args.q_lora_rank
452
+ self.o_lora_rank = args.o_lora_rank
453
+ self.head_dim = args.head_dim
454
+ self.rope_head_dim = args.rope_head_dim
455
+ self.nope_head_dim = args.head_dim - args.rope_head_dim
456
+ self.n_groups = args.o_groups
457
+ self.n_local_groups = self.n_groups // world_size
458
+ self.window_size = args.window_size
459
+ self.compress_ratio = args.compress_ratios[layer_id]
460
+ self.eps = args.norm_eps
461
+
462
+ self.attn_sink = nn.Parameter(torch.empty(self.n_local_heads, dtype=torch.float32))
463
+ self.wq_a = Linear(self.dim, self.q_lora_rank)
464
+ self.q_norm = RMSNorm(self.q_lora_rank, self.eps)
465
+ self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.head_dim)
466
+ self.wkv = Linear(self.dim, self.head_dim)
467
+ self.kv_norm = RMSNorm(self.head_dim, self.eps)
468
+ self.wo_a = ColumnParallelLinear(self.n_heads * self.head_dim // self.n_groups, self.n_groups * args.o_lora_rank, dtype=torch.bfloat16)
469
+ self.wo_b = RowParallelLinear(self.n_groups * args.o_lora_rank, self.dim)
470
+ self.softmax_scale = self.head_dim ** -0.5
471
+
472
+ if self.compress_ratio:
473
+ self.compressor = Compressor(args, self.compress_ratio, self.head_dim)
474
+ if self.compress_ratio == 4:
475
+ self.indexer = Indexer(args, self.compress_ratio)
476
+ else:
477
+ self.indexer = None
478
+
479
+ kv_cache_size = args.window_size + (args.max_seq_len // self.compress_ratio if self.compress_ratio else 0)
480
+ self.register_buffer("kv_cache", torch.zeros(args.max_batch_size, kv_cache_size, self.head_dim), persistent=False)
481
+ if self.compress_ratio:
482
+ original_seq_len, rope_theta = args.original_seq_len, args.compress_rope_theta
483
+ else:
484
+ # disable YaRN and use base rope_theta in pure sliding-window attention
485
+ original_seq_len, rope_theta = 0, args.rope_theta
486
+ freqs_cis = precompute_freqs_cis(self.rope_head_dim, args.max_seq_len, original_seq_len,
487
+ rope_theta, args.rope_factor, args.beta_fast, args.beta_slow)
488
+ self.register_buffer("freqs_cis", freqs_cis, persistent=False)
489
+
490
+ def forward(self, x: torch.Tensor, start_pos: int):
491
+ bsz, seqlen, _ = x.size()
492
+ freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen]
493
+ win = self.window_size
494
+ ratio = self.compress_ratio
495
+ rd = self.rope_head_dim
496
+ if self.compress_ratio and self.compressor.kv_cache is None:
497
+ self.compressor.kv_cache = self.kv_cache[:, win:]
498
+ self.compressor.freqs_cis = self.freqs_cis
499
+ if self.indexer is not None:
500
+ self.indexer.freqs_cis = self.freqs_cis
501
+ # q
502
+ qr = q = self.q_norm(self.wq_a(x))
503
+ q = self.wq_b(q).unflatten(-1, (self.n_local_heads, self.head_dim))
504
+ q *= torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps)
505
+ apply_rotary_emb(q[..., -rd:], freqs_cis)
506
+
507
+ # win kv & topk_idxs
508
+ kv = self.wkv(x)
509
+ kv = self.kv_norm(kv)
510
+ apply_rotary_emb(kv[..., -rd:], freqs_cis)
511
+ # FP8-simulate non-rope dims to match QAT; rope dims stay bf16 for positional precision
512
+ act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
513
+ topk_idxs = get_window_topk_idxs(win, bsz, seqlen, start_pos)
514
+ if self.compress_ratio:
515
+ offset = kv.size(1) if start_pos == 0 else win
516
+ if self.indexer is not None:
517
+ compress_topk_idxs = self.indexer(x, qr, start_pos, offset).int()
518
+ else:
519
+ compress_topk_idxs = get_compress_topk_idxs(ratio, bsz, seqlen, start_pos, offset)
520
+ topk_idxs = torch.cat([topk_idxs, compress_topk_idxs], dim=-1)
521
+
522
+ # compress kv & attn
523
+ if start_pos == 0:
524
+ if seqlen <= win:
525
+ self.kv_cache[:bsz, :seqlen] = kv
526
+ else:
527
+ cutoff = seqlen % win
528
+ self.kv_cache[:bsz, cutoff: win], self.kv_cache[:bsz, :cutoff] = kv[:, -win:].split([win - cutoff, cutoff], dim=1)
529
+ if self.compress_ratio:
530
+ if (kv_compress := self.compressor(x, start_pos)) is not None:
531
+ kv = torch.cat([kv, kv_compress], dim=1)
532
+ # We performed QAT here, kv could also use fp8 format, though current implementation uses bf16
533
+ o = sparse_attn(q, kv, self.attn_sink, topk_idxs, self.softmax_scale)
534
+ else:
535
+ self.kv_cache[:bsz, start_pos % win] = kv.squeeze(1)
536
+ if self.compress_ratio:
537
+ self.compressor(x, start_pos)
538
+ o = sparse_attn(q, self.kv_cache[:bsz], self.attn_sink, topk_idxs, self.softmax_scale)
539
+ apply_rotary_emb(o[..., -rd:], freqs_cis, True)
540
+
541
+ # o
542
+ o = o.view(bsz, seqlen, self.n_local_groups, -1)
543
+ wo_a = self.wo_a.weight.view(self.n_local_groups, self.o_lora_rank, -1)
544
+ # NOTE: wo_a is FP8 in checkpoint; could do FP8 einsum here for better perf,
545
+ # but using BF16 for simplicity.
546
+ o = torch.einsum("bsgd,grd->bsgr", o, wo_a)
547
+ x = self.wo_b(o.flatten(2))
548
+ return x
549
+
550
+
551
+ class Gate(nn.Module):
552
+ """MoE gating: computes expert routing scores and selects top-k experts.
553
+ Supports hash-based routing (first n_hash_layers) where expert indices are
554
+ predetermined per token ID, and score-based routing (remaining layers)."""
555
+ def __init__(self, layer_id: int, args: ModelArgs):
556
+ super().__init__()
557
+ self.dim = args.dim
558
+ self.topk = args.n_activated_experts
559
+ self.score_func = args.score_func
560
+ self.route_scale = args.route_scale
561
+ self.hash = layer_id < args.n_hash_layers
562
+ self.weight = nn.Parameter(torch.empty(args.n_routed_experts, args.dim))
563
+ if self.hash:
564
+ self.tid2eid = nn.Parameter(torch.empty(args.vocab_size, args.n_activated_experts, dtype=torch.int32), requires_grad=False)
565
+ self.bias = None
566
+ else:
567
+ self.bias = nn.Parameter(torch.empty(args.n_routed_experts, dtype=torch.float32))
568
+
569
+ def forward(self, x: torch.Tensor, input_ids: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor]:
570
+ scores = linear(x.float(), self.weight.float())
571
+ if self.score_func == "softmax":
572
+ scores = scores.softmax(dim=-1)
573
+ elif self.score_func == "sigmoid":
574
+ scores = scores.sigmoid()
575
+ else:
576
+ scores = F.softplus(scores).sqrt()
577
+ original_scores = scores
578
+ # Bias shifts scores for expert selection (topk) but does not affect routing weights.
579
+ if self.bias is not None:
580
+ scores = scores + self.bias
581
+ if self.hash:
582
+ indices = self.tid2eid[input_ids]
583
+ else:
584
+ indices = scores.topk(self.topk, dim=-1)[1]
585
+ weights = original_scores.gather(1, indices)
586
+ if self.score_func != "softmax":
587
+ weights /= weights.sum(dim=-1, keepdim=True)
588
+ weights *= self.route_scale
589
+ return weights, indices
590
+
591
+
592
+ class Expert(nn.Module):
593
+ """Single MoE expert: SwiGLU FFN (w1, w2, w3). Computation in float32 for stability."""
594
+ def __init__(self, dim: int, inter_dim: int, dtype=None, swiglu_limit=0):
595
+ super().__init__()
596
+ self.w1 = Linear(dim, inter_dim, dtype=dtype)
597
+ self.w2 = Linear(inter_dim, dim, dtype=dtype)
598
+ self.w3 = Linear(dim, inter_dim, dtype=dtype)
599
+ self.swiglu_limit = swiglu_limit
600
+
601
+ def forward(self, x: torch.Tensor, weights: Optional[torch.Tensor] = None) -> torch.Tensor:
602
+ dtype = x.dtype
603
+ gate = self.w1(x).float()
604
+ up = self.w3(x).float()
605
+ if self.swiglu_limit > 0:
606
+ up = torch.clamp(up, min=-self.swiglu_limit, max=self.swiglu_limit)
607
+ gate = torch.clamp(gate, max=self.swiglu_limit)
608
+ x = F.silu(gate) * up
609
+ if weights is not None:
610
+ x = weights * x
611
+ return self.w2(x.to(dtype))
612
+
613
+
614
+ class MoE(nn.Module):
615
+ """Mixture-of-Experts: gate routes each token to top-k routed experts + 1 shared expert.
616
+ Experts are sharded across TP ranks; each rank handles n_routed_experts // world_size experts."""
617
+ def __init__(self, layer_id: int, args: ModelArgs):
618
+ super().__init__()
619
+ self.layer_id = layer_id
620
+ self.dim = args.dim
621
+ assert args.n_routed_experts % world_size == 0, f"Number of experts must be divisible by world size (world_size={world_size})"
622
+ self.n_routed_experts = args.n_routed_experts
623
+ self.n_local_experts = args.n_routed_experts // world_size
624
+ self.n_activated_experts = args.n_activated_experts
625
+ self.experts_start_idx = rank * self.n_local_experts
626
+ self.experts_end_idx = self.experts_start_idx + self.n_local_experts
627
+ self.gate = Gate(layer_id, args)
628
+ expert_dtype = torch.float4_e2m1fn_x2 if args.expert_dtype == "fp4" else None
629
+ self.experts = nn.ModuleList([Expert(args.dim, args.moe_inter_dim, dtype=expert_dtype, swiglu_limit=args.swiglu_limit) if self.experts_start_idx <= i < self.experts_end_idx else None
630
+ for i in range(self.n_routed_experts)])
631
+ assert args.n_shared_experts == 1
632
+ self.shared_experts = Expert(args.dim, args.moe_inter_dim, swiglu_limit=args.swiglu_limit)
633
+
634
+ def forward(self, x: torch.Tensor, input_ids: torch.Tensor) -> torch.Tensor:
635
+ shape = x.size()
636
+ x = x.view(-1, self.dim)
637
+ weights, indices = self.gate(x, input_ids.flatten())
638
+ y = torch.zeros_like(x, dtype=torch.float32)
639
+ counts = torch.bincount(indices.flatten(), minlength=self.n_routed_experts).tolist()
640
+ for i in range(self.experts_start_idx, self.experts_end_idx):
641
+ if counts[i] == 0:
642
+ continue
643
+ expert = self.experts[i]
644
+ idx, top = torch.where(indices == i)
645
+ y[idx] += expert(x[idx], weights[idx, top, None])
646
+ if world_size > 1:
647
+ dist.all_reduce(y)
648
+ y += self.shared_experts(x)
649
+ return y.type_as(x).view(shape)
650
+
651
+
652
+ class Block(nn.Module):
653
+ """Transformer block with Hyper-Connections (HC) mixing.
654
+ Instead of a simple residual, HC maintains `hc_mult` copies of the hidden state.
655
+ hc_pre: reduces hc copies -> 1 via learned weighted sum (pre-weights from Sinkhorn).
656
+ hc_post: expands 1 -> hc copies via learned post-weights + combination matrix."""
657
+ attention_cls = Attention
658
+
659
+ def __init__(self, layer_id: int, args: ModelArgs):
660
+ super().__init__()
661
+ self.layer_id = layer_id
662
+ self.norm_eps = args.norm_eps
663
+ self.attn = self.attention_cls(layer_id, args)
664
+ self.ffn = MoE(layer_id, args)
665
+ self.attn_norm = RMSNorm(args.dim, self.norm_eps)
666
+ self.ffn_norm = RMSNorm(args.dim, self.norm_eps)
667
+ self.hc_mult = hc_mult = args.hc_mult
668
+ self.hc_sinkhorn_iters = args.hc_sinkhorn_iters
669
+ self.hc_eps = args.hc_eps
670
+ mix_hc = (2 + hc_mult) * hc_mult
671
+ hc_dim = hc_mult * args.dim
672
+ with set_dtype(torch.float32):
673
+ self.hc_attn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim))
674
+ self.hc_ffn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim))
675
+ self.hc_attn_base = nn.Parameter(torch.empty(mix_hc))
676
+ self.hc_ffn_base = nn.Parameter(torch.empty(mix_hc))
677
+ self.hc_attn_scale = nn.Parameter(torch.empty(3))
678
+ self.hc_ffn_scale = nn.Parameter(torch.empty(3))
679
+
680
+ def hc_pre(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor):
681
+ # x: [b,s,hc,d], hc_fn: [mix_hc,hc*d], hc_scale: [3], hc_base: [mix_hc], y: [b,s,hc,d]
682
+ shape, dtype = x.size(), x.dtype
683
+ x = x.flatten(2).float()
684
+ rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps)
685
+ mixes = F.linear(x, hc_fn) * rsqrt
686
+ pre, post, comb = hc_split_sinkhorn(mixes, hc_scale, hc_base, self.hc_mult, self.hc_sinkhorn_iters, self.hc_eps)
687
+ y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=2)
688
+ return y.to(dtype), post, comb
689
+
690
+ def hc_post(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, comb: torch.Tensor):
691
+ # x: [b,s,d], residual: [b,s,hc,d], post: [b,s,hc], comb: [b,s,hc,hc], y: [b,s,hc,d]
692
+ y = post.unsqueeze(-1) * x.unsqueeze(-2) + torch.sum(comb.unsqueeze(-1) * residual.unsqueeze(-2), dim=2)
693
+ return y.type_as(x)
694
+
695
+ def forward(self, x: torch.Tensor, start_pos: int, input_ids: Optional[torch.Tensor], *attn_args) -> torch.Tensor:
696
+ residual = x
697
+ x, post, comb = self.hc_pre(x, self.hc_attn_fn, self.hc_attn_scale, self.hc_attn_base)
698
+ x = self.attn_norm(x)
699
+ x = self.attn(x, start_pos, *attn_args)
700
+ x = self.hc_post(x, residual, post, comb)
701
+
702
+ residual = x
703
+ x, post, comb = self.hc_pre(x, self.hc_ffn_fn, self.hc_ffn_scale, self.hc_ffn_base)
704
+ x = self.ffn_norm(x)
705
+ x = self.ffn(x, input_ids)
706
+ x = self.hc_post(x, residual, post, comb)
707
+ return x
708
+
709
+ def hc_head(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor):
710
+ shape, dtype = x.size(), x.dtype
711
+ x = x.flatten(2).float()
712
+ rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps)
713
+ mixes = F.linear(x, hc_fn) * rsqrt
714
+ pre = torch.sigmoid(mixes * hc_scale + hc_base) + self.hc_eps
715
+ y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=2)
716
+ return y.to(dtype)
717
+
718
+
719
+ class ParallelHead(nn.Module):
720
+
721
+ def __init__(self, vocab_size: int, dim: int, norm_eps: float = 1e-6, hc_eps: float = 1e-6):
722
+ super().__init__()
723
+ self.vocab_size = vocab_size
724
+ self.dim = dim
725
+ self.norm_eps = norm_eps
726
+ self.hc_eps = hc_eps
727
+ self.part_vocab_size = (vocab_size // world_size)
728
+ # lm_head in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for easier computation of logits later.
729
+ self.weight = nn.Parameter(torch.empty(self.part_vocab_size, self.dim, dtype=torch.float32))
730
+
731
+ def forward(self, x: torch.Tensor, full_logits=False):
732
+ # x: [b,s,hc,d]
733
+ if not full_logits:
734
+ x = x[:, -1]
735
+ logits = F.linear(x.float(), self.weight)
736
+ if world_size > 1:
737
+ all_logits = [torch.empty_like(logits) for _ in range(world_size)]
738
+ dist.all_gather(all_logits, logits)
739
+ logits = torch.cat(all_logits, dim=-1)
740
+ return logits
741
+
742
+
743
+ @lru_cache(1)
744
+ def get_dspark_topk_idxs(window_size: int, bsz: int, block_size: int, start_pos: int):
745
+ assert start_pos > 0
746
+ matrix = torch.cat([torch.arange(min(window_size, start_pos + 1)), window_size + torch.arange(block_size)])
747
+ return matrix.int().view(1, 1, -1).expand(bsz, block_size, -1).contiguous()
748
+
749
+
750
+ class DSparkAttention(Attention):
751
+
752
+ def forward(self, x: torch.Tensor, start_pos: int, main_x: torch.Tensor):
753
+ assert self.compress_ratio == 0
754
+ bsz, seqlen, _ = main_x.size()
755
+ win = self.window_size
756
+ rd = self.rope_head_dim
757
+
758
+ main_freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen]
759
+ main_kv = self.kv_norm(self.wkv(main_x))
760
+ apply_rotary_emb(main_kv[..., -rd:], main_freqs_cis)
761
+ act_quant(main_kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
762
+
763
+ if start_pos == 0:
764
+ if seqlen <= win:
765
+ self.kv_cache[:bsz, :seqlen] = main_kv
766
+ else:
767
+ cutoff = seqlen % win
768
+ self.kv_cache[:bsz, cutoff: win], self.kv_cache[:bsz, :cutoff] = main_kv[:, -win:].split([win - cutoff, cutoff], dim=1)
769
+ return x
770
+
771
+ bsz, block_size, _ = x.size()
772
+ freqs_cis = self.freqs_cis[start_pos+seqlen:start_pos+seqlen+block_size]
773
+
774
+ q = self.q_norm(self.wq_a(x))
775
+ q = self.wq_b(q).unflatten(-1, (self.n_local_heads, self.head_dim))
776
+ q *= torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps)
777
+ apply_rotary_emb(q[..., -rd:], freqs_cis)
778
+ kv = self.kv_norm(self.wkv(x))
779
+ apply_rotary_emb(kv[..., -rd:], freqs_cis)
780
+ act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
781
+
782
+ topk_idxs = get_dspark_topk_idxs(win, bsz, block_size, start_pos)
783
+ self.kv_cache[:bsz, start_pos % win] = main_kv.squeeze(1)
784
+ kv = torch.cat([self.kv_cache[:bsz], kv], dim=1)
785
+ o = sparse_attn(q, kv, self.attn_sink, topk_idxs, self.softmax_scale)
786
+ apply_rotary_emb(o[..., -rd:], freqs_cis, True)
787
+
788
+ o = o.view(bsz, block_size, self.n_local_groups, -1)
789
+ wo_a = self.wo_a.weight.view(self.n_local_groups, self.o_lora_rank, -1)
790
+ o = torch.einsum("bsgd,grd->bsgr", o, wo_a)
791
+ x = self.wo_b(o.flatten(2))
792
+ return x
793
+
794
+
795
+ class DSparkMarkovHead(nn.Module):
796
+ def __init__(self, vocab_size: int, dspark_markov_rank: int):
797
+ super().__init__()
798
+ self.markov_w1 = ParallelEmbedding(vocab_size, dspark_markov_rank)
799
+ self.markov_w2 = ParallelHead(vocab_size, dspark_markov_rank)
800
+
801
+ def forward(self, token_ids: torch.Tensor) -> torch.Tensor:
802
+ embed = self.markov_w1(token_ids)
803
+ logits = self.markov_w2(embed, full_logits=True)
804
+ return logits, embed
805
+
806
+
807
+ class DSparkConfidenceHead(nn.Module):
808
+ def __init__(self, input_dim: int):
809
+ super().__init__()
810
+ # proj in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for fp32 confidence score.
811
+ self.proj = Linear(input_dim, 1, dtype=torch.float32)
812
+
813
+ def forward(self, hidden: torch.Tensor, markov_embed: torch.Tensor):
814
+ hidden = torch.cat([hidden, markov_embed], dim=-1)
815
+ return self.proj(hidden.float()).squeeze(-1)
816
+
817
+
818
+ class DSparkBlock(Block):
819
+ """DSpark stage stored under the mtp.* checkpoint namespace."""
820
+ attention_cls = DSparkAttention
821
+
822
+ def __init__(self, layer_id: int, args: ModelArgs):
823
+ super().__init__(layer_id, args)
824
+ self.dim = args.dim
825
+ stage_id = layer_id - args.n_layers
826
+ self.block_size = args.dspark_block_size
827
+ self.noise_token_id = args.dspark_noise_token_id
828
+ self.temperature = args.temperature
829
+ hc_dim = self.hc_mult * args.dim
830
+ if stage_id == 0:
831
+ assert len(args.dspark_target_layer_ids) > 0, "DSpark needs target layers"
832
+ self.main_proj = Linear(args.dim * len(args.dspark_target_layer_ids), args.dim)
833
+ self.main_norm = RMSNorm(args.dim, args.norm_eps)
834
+ if stage_id == args.n_mtp_layers - 1:
835
+ self.norm = RMSNorm(args.dim, args.norm_eps)
836
+ self.markov_head = DSparkMarkovHead(args.vocab_size, args.dspark_markov_rank)
837
+ self.confidence_head = DSparkConfidenceHead(args.dim + args.dspark_markov_rank)
838
+ with set_dtype(torch.float32):
839
+ self.hc_head_fn = nn.Parameter(torch.empty(self.hc_mult, hc_dim))
840
+ self.hc_head_base = nn.Parameter(torch.empty(self.hc_mult))
841
+ self.hc_head_scale = nn.Parameter(torch.empty(1))
842
+ self.embed: ParallelEmbedding = None
843
+ self.head: ParallelHead = None
844
+
845
+ def forward(self, x: torch.Tensor, start_pos: int, input_ids: torch.Tensor, main_x: torch.Tensor) -> torch.Tensor:
846
+ if start_pos > 0:
847
+ return super().forward(x, start_pos, input_ids, main_x)
848
+ # only compute KV cache in prefill stage
849
+ return self.attn(x, start_pos, main_x)
850
+
851
+ def forward_embed(self, main_hidden: torch.Tensor, input_ids: torch.Tensor):
852
+ assert self.embed is not None
853
+ main_x = self.main_norm(self.main_proj(main_hidden))
854
+ draft_input_ids = input_ids.new_full([input_ids.size(0), self.block_size], self.noise_token_id)
855
+ draft_input_ids[:, 0] = input_ids
856
+ x = self.embed(draft_input_ids)
857
+ x = x.unsqueeze(2).repeat(1, 1, self.hc_mult, 1)
858
+ return x, main_x
859
+
860
+ def forward_head(self, x: torch.Tensor, input_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
861
+ assert self.head is not None
862
+ x = self.hc_head(x, self.hc_head_fn, self.hc_head_scale, self.hc_head_base)
863
+ logits = self.head(self.norm(x), full_logits=True)
864
+ output_ids = input_ids.new_empty(input_ids.size(0), self.block_size + 1)
865
+ output_ids[:, 0] = input_ids
866
+ markov_embeds = []
867
+ for i in range(self.block_size):
868
+ logits_bias, markov_embed = self.markov_head(output_ids[:, i])
869
+ logits[:, i].add_(logits_bias)
870
+ markov_embeds.append(markov_embed)
871
+ output_ids[:, i + 1] = sample(logits[:, i], self.temperature)
872
+ markov_embed = torch.stack(markov_embeds, dim=1)
873
+ confidence = self.confidence_head(x, markov_embed)
874
+ return output_ids, logits, confidence
875
+
876
+
877
+ class Transformer(nn.Module):
878
+ """Full DeepSeek-V4 model: embed -> HC-expand -> N blocks -> HC-head -> logits.
879
+ Sets global state (world_size, rank, default_dtype, scale_fmt, scale_dtype) in __init__."""
880
+ def __init__(self, args: ModelArgs):
881
+ global world_size, rank, default_dtype, scale_fmt, scale_dtype
882
+ world_size = dist.get_world_size() if dist.is_initialized() else 1
883
+ rank = dist.get_rank() if dist.is_initialized() else 0
884
+ default_dtype = torch.float8_e4m3fn if args.dtype == "fp8" else torch.bfloat16
885
+ scale_fmt = "ue8m0" if args.scale_dtype == "fp8" else args.scale_fmt
886
+ scale_dtype = torch.float8_e8m0fnu if args.scale_dtype == "fp8" else torch.float32
887
+ super().__init__()
888
+ self.max_seq_len = args.max_seq_len
889
+ self.temperature = args.temperature
890
+ self.norm_eps = args.norm_eps
891
+ self.hc_eps = args.hc_eps
892
+ self.embed = ParallelEmbedding(args.vocab_size, args.dim)
893
+ self.layers = torch.nn.ModuleList()
894
+ for layer_id in range(args.n_layers):
895
+ self.layers.append(Block(layer_id, args))
896
+ self.norm = RMSNorm(args.dim, self.norm_eps)
897
+ self.head = ParallelHead(args.vocab_size, args.dim, self.norm_eps, self.hc_eps)
898
+ self.mtp = torch.nn.ModuleList()
899
+ self.target_layer_ids = args.dspark_target_layer_ids
900
+ if args.dspark_block_size:
901
+ for layer_id in range(args.n_mtp_layers):
902
+ self.mtp.append(DSparkBlock(args.n_layers + layer_id, args))
903
+ self.mtp[-1].embed = self.embed
904
+ self.mtp[-1].head = self.head
905
+ self.hc_mult = hc_mult = args.hc_mult
906
+ hc_dim = hc_mult * args.dim
907
+ with set_dtype(torch.float32):
908
+ self.hc_head_fn = nn.Parameter(torch.empty(hc_mult, hc_dim))
909
+ self.hc_head_base = nn.Parameter(torch.empty(hc_mult))
910
+ self.hc_head_scale = nn.Parameter(torch.empty(1))
911
+
912
+ @torch.inference_mode()
913
+ def forward(self, input_ids: torch.Tensor, start_pos: int = 0):
914
+ h = self.embed(input_ids)
915
+ # Expand to hc_mult copies for Hyper-Connections
916
+ h = h.unsqueeze(2).repeat(1, 1, self.hc_mult, 1)
917
+ main_hiddens = []
918
+ for i, layer in enumerate(self.layers):
919
+ h = layer(h, start_pos, input_ids)
920
+ if i in self.target_layer_ids:
921
+ main_hiddens.append(h.mean(dim=2))
922
+ h = layer.hc_head(h, self.hc_head_fn, self.hc_head_scale, self.hc_head_base)
923
+ logits = self.head(self.norm(h))
924
+ output_ids = sample(logits, self.temperature)
925
+ main_hidden = torch.cat(main_hiddens, dim=-1) if main_hiddens else None
926
+ return output_ids, logits, main_hidden
927
+
928
+ @torch.inference_mode()
929
+ def forward_spec(self, input_ids: torch.Tensor, main_hidden: torch.Tensor, start_pos: int = 0):
930
+ h, main_x = self.mtp[0].forward_embed(main_hidden, input_ids)
931
+ for layer in self.mtp:
932
+ h = layer(h, start_pos, input_ids, main_x)
933
+ if start_pos == 0:
934
+ return
935
+ output_ids, logits, confidence = self.mtp[-1].forward_head(h, input_ids)
936
+ return output_ids, logits, confidence
937
+
938
+
939
+ def sample(logits, temperature: float = 1.0):
940
+ """Gumbel-max trick: equivalent to multinomial sampling but faster on GPU,
941
+ since it avoids the GPU-to-CPU sync in torch.multinomial."""
942
+ if temperature == 0:
943
+ return logits.argmax(dim=-1)
944
+ logits = logits / max(temperature, 1e-5)
945
+ probs = torch.softmax(logits, dim=-1, dtype=torch.float32)
946
+ return probs.div_(torch.empty_like(probs).exponential_(1)).argmax(dim=-1)
947
+
948
+
949
+ if __name__ == "__main__":
950
+ torch.set_default_dtype(torch.bfloat16)
951
+ torch.set_default_device("cuda")
952
+ torch.manual_seed(0)
953
+ args = ModelArgs(n_hash_layers=0, dspark_block_size=6, dspark_target_layer_ids=(5, 6))
954
+ x = torch.randint(0, args.vocab_size, (2, 150))
955
+ model = Transformer(args)
956
+
957
+ output_ids, logits, main_hidden = model(x[:, :128])
958
+ model.forward_spec(output_ids, main_hidden)
959
+ for i in range(128, 150):
960
+ output_ids, logits, main_hidden = model(x[:, i:i+1], i)
961
+ output_ids, logits, confidence = model.forward_spec(output_ids, main_hidden, i)
inference/requirements.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ torch>=2.10.0
2
+ transformers>=5.0.0
3
+ safetensors>=0.7.0
4
+ fast_hadamard_transform
5
+ tilelang==0.1.8
model-00004-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:30d5f13bcb5747469efa4b8a9f23c236b409c1433e6772b76ffe59e7983d2455
3
+ size 3596229272
model-00006-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a104d4b2b4a4112284ff262c79668d32b7074effded82f676288b1e20e835391
3
+ size 3590024776
model-00008-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:411e5e73a505bf7e52cfc648914a3a58973cfa6094a150d1a81147f3f3af1e7a
3
+ size 3590024776
model-00009-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ff9b2d66b6e57c6c23eb65f92b621619b4dff547b13a440474cf4d0c443d1375
3
+ size 3568768976
model-00010-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:10253bd784b438906fc18963f413a216001a8eaec5f0a4e27a1e9933b7bd50e1
3
+ size 3590024776
model-00013-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:16dba59cb7c3a29fe4acb3f2c96ba3030dec09e6c72d857369471b6e9820861e
3
+ size 3568770544
model-00014-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:49d38206befa801976b694753938ace28557b9bfc72e40af72420443fed60dce
3
+ size 3590026352
model-00017-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:456d02b3ee143a55761b396c44b6116f3a5bf3e39a3e38d031a04d6f182965e4
3
+ size 3568770544
model-00021-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ebc2e7e22616da6949f11ffbe1c3fb10c6f94bc701f70304d18ee2ac02ac9018
3
+ size 3568770544
model-00022-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2bac4a19c9b2b53d1da98a2e73e552f36911b135f8dec6616a7018587edeb193
3
+ size 3590026352
model-00023-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2ae193e3af0ef90ba496ceee66103a4259592f1b164dbc4c020cce3015e8df8a
3
+ size 3568770544
model-00024-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c618d923abb2b692ec2fcab459e25c7fe2d95f2dce6612a32d2111745be86da8
3
+ size 3590026352
model-00026-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1c1e74ca141034a0a425e85fbe827e812ad8cf2a34cbf59cee587bc4b3d4df3b
3
+ size 3590026352
model-00027-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4e9b62fd5c24b6d444aa178f5835f3871bdfbcc57c4695567587d418755da3d6
3
+ size 3568770544
model-00028-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:990106854454786e08b2fd138e9e63437296b6dfc90938a91efa1f61bdcf9f52
3
+ size 3590026352
model-00030-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fa2b5a21cf6578cb95db7cb4ec4557543a3a7c981350be58d168b064d244dca4
3
+ size 3590026352
model-00031-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b771091af9f13a05904876a30fae9d227fed3a3a7f42b64dd35e03ea535e816d
3
+ size 3568770544
model-00033-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e36993a76b05ae9e2b82215e08730d49808be49d841e53d9ca82e12ef9234aa1
3
+ size 3568770544
model-00034-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:beed84bfe81b090018e4882ce34bbf1c14a3ff5b781a889eb227a0834679a434
3
+ size 3590026352
model-00035-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5531f7d0dcb105ff28ff034de2f02fcceb041940d69d022dc818dfea795f125e
3
+ size 3568770544
model-00036-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8d9f8dd26a891f869c150b695b8990ad0d9ab2025995c8b58adce33bb9527035
3
+ size 3590026352
model-00037-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c332b14fe0c1466b86b8f67e68b98c5e44811f2dfd4b668ff371b61e02bd9ace
3
+ size 3568770544
model-00040-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e684541f002d60f34844beb70b14bba5c97cd6380e39dd04b2706f3f19e320ce
3
+ size 3590026352
model-00041-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8b77345a35686838feb020c256180d3e60b5c6e893ef1546c3de9aa758e30f73
3
+ size 3568770544
model-00047-of-00048.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9cee337c3b9f5e8b01718e5abd3d65ff1b5c08529f8ecaa33292b692b7872a4a
3
+ size 3560111960
model.safetensors.index.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_bos_token": false,
3
+ "add_eos_token": false,
4
+ "bos_token": {
5
+ "__type": "AddedToken",
6
+ "content": "<|begin▁of▁sentence|>",
7
+ "lstrip": false,
8
+ "normalized": true,
9
+ "rstrip": false,
10
+ "single_word": false
11
+ },
12
+ "clean_up_tokenization_spaces": false,
13
+ "eos_token": {
14
+ "__type": "AddedToken",
15
+ "content": "<|end▁of▁sentence|>",
16
+ "lstrip": false,
17
+ "normalized": true,
18
+ "rstrip": false,
19
+ "single_word": false
20
+ },
21
+ "legacy": true,
22
+ "model_max_length": 1048576,
23
+ "pad_token": {
24
+ "__type": "AddedToken",
25
+ "content": "<|end▁of▁sentence|>",
26
+ "lstrip": false,
27
+ "normalized": true,
28
+ "rstrip": false,
29
+ "single_word": false
30
+ },
31
+ "sp_model_kwargs": {},
32
+ "unk_token": null,
33
+ "tokenizer_class": "PreTrainedTokenizerFast"
34
+ }