chenjz24 commited on
Commit
f74eb65
·
verified ·
1 Parent(s): a9201d4

Upload folder using huggingface_hub

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ tokenizer.json filter=lfs diff=lfs merge=lfs -text
INFERENCE_OPTIMIZATION.md ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 完整模型精简与推理加速
2
+
3
+ 精简模型位于 `runs/bilingual_s5_full/final_compact/`,由 `runs/bilingual_s5_full/final/` 导出,保留文本、图像、视频、音频理解和语音输出。原始目录保留。
4
+
5
+ ## 模型体积
6
+
7
+ | 项目 | 原模型 | 精简模型 |
8
+ |---|---:|---:|
9
+ | 独立参数量 | 2,031,972,995 | 2,031,710,851 |
10
+ | safetensors 权重字节数 | 5,677,248,858 | 4,115,251,826 |
11
+ | 权重文件大小(十进制 GB) | 5.677 | 4.115 |
12
+
13
+ 权重文件减少 27.5%,实际参数减少 262,144(约 0.013%)。参数统计覆盖完整多模态模型,共享参数只计算一次。
14
+
15
+ 删除的是 codec 两个从未参与解码的 `input_proj`。Thinker 按配置的 BF16 推理精度保存;音频投影器和特殊 token 参数保持 FP32。精简目录支持常规 `save_pretrained` / `from_pretrained`。
16
+
17
+ Talker 的 132,588 行文本 embedding 全部可达且无完全重复行;32 张声码 embedding/输出头也无可进一步共享的完整表或重复行。原有 15 张共享声码 embedding 保持共享。
18
+
19
+ ## 推理结构
20
+
21
+ - 音频分块长度在 CPU 上一次计算,SDPA 各层复用,避免反复从 GPU 取回长度。
22
+ - 不做音频池化或堆叠时,直接返回编码输出。
23
+ - Thinker 的线性注意力 decoder 层在 CUDA 上对单 token 解码使用 CUDA Graph,覆盖归一化、递归状态更新、MLP 和残差计算。运算、权重和动态 KV cache 保持原样。
24
+ - 每个请求保有独立缓存;训练、有梯度的 forward、CPU 和多 token prefill 执行常规模型路径。改变设备、精度或训练模式会释放已有 graph。
25
+ - `generate()` 使用公开的 `model.generation_config`,并在调用期间选择 Flash/memory-efficient/math SDPA 后端。
26
+
27
+ ## 验证结果
28
+
29
+ 完整 MMAU test-mini 1,000 条输入中,原模型与精简模型生成 token 和文本全部一致,均自然到 EOS;答案字母准确率均为 **64.9%(649/1000)**。sound、speech、music 分别为 240/333、210/333、199/334。清单为 `data/spoken_execution_eval/mmau_test_mini_1000.jsonl`,提示词保持该清单原样。
30
+
31
+ 音频问答时延在 sound/speech/music 各 4 条代表样本上各测 3 次,两模型交替执行,所有 token 一致。MMAU 每条输出 2 token(答案字母和 EOS)。
32
+
33
+ | 中位数指标 | 原模型 | 精简模型 |
34
+ |---|---:|---:|
35
+ | 整答时延 | 173.03 ms | 97.35 ms |
36
+ | TTFT | 129.12 ms | 86.43 ms |
37
+ | 解码吞吐量 | 30.71 token/s | 65.03 token/s |
38
+ | 端到端吞吐量 | 11.56 token/s | 20.54 token/s |
39
+
40
+ 同一样本成对加速比中位数 **1.82×**。TTFT/解码吞吐量来自独立 streaming 测量,普通整答时延未插入逐 token 同步。
41
+
42
+ 6 个中英文文本任务、每任务 2 次普通测量和 1 次独立 streaming 测量,生成 token 全部一致;回答长度为 10、35、58、110、256、256 token。整答成对加速比中位数 **2.44×**,整答时延中位数 **2.490 s → 0.944 s**,解码吞吐量中位数 **32.41 → 77.00 token/s**。两条 256-token 输出达到设定上限,两模型停止位置相同。
43
+
44
+ 加载后 1,574 个保留 state-dict tensor 逐元素一致;真实音频投影特征以及中英文语音的声码、波形逐元素一致。CPU/CUDA 测试覆盖 FP32/BF16、缓存独立性、保留 hidden states、图复用、训练梯度、保存加载及图像/视频路径,共 **126 项通过**。
45
+
46
+ 结果文件位于 `runs/hf_inference_optimization/`:
47
+
48
+ - `mmau_equivalence.json`:1,000 条 MMAU 输出及准确率对比。
49
+ - `mmau_verified_shard_0.json` 至 `mmau_verified_shard_3.json`:完整逐条记录。
50
+ - `mmau_latency.json`:12 条样本、3 次重复的最终音频问答计时。
51
+ - `text_decode_pair.json`:长回答计时、TTFT、吞吐量和 token 对比。
52
+ - `weight_and_speech_equivalence.json`:权重、音频特征和语音输出等价验证。
53
+ - `redundancy.json`:embedding 可达性及逐行去重统计。
54
+
55
+ 实测使用 NVIDIA H20、PyTorch 2.11.0+cu130、Transformers 5.12.1、fla-core 0.5.0、causal-conv1d 1.6.0;CPU 功能及独立 CUDA Graph 测试另外覆盖 PyTorch 2.6.0+cu124。
56
+
57
+ ## 使用
58
+
59
+ 上传的 `run_mmau.py` 只需把 `--model` 指向 `runs/bilingual_s5_full/final_compact`。默认自动使用优化路径,保留两次预热。
60
+
61
+ 重新导出:
62
+
63
+ ```bash
64
+ python -m hf.compact_edgeinstant \
65
+ --source runs/bilingual_s5_full/final \
66
+ --output runs/bilingual_s5_full/final_compact_new \
67
+ --device cuda:0
68
+ ```
69
+
70
+ 配对测量:
71
+
72
+ ```bash
73
+ OMP_NUM_THREADS=2 MKL_NUM_THREADS=2 python scripts/benchmark_hf_inference.py \
74
+ --source runs/bilingual_s5_full/final \
75
+ --candidate runs/bilingual_s5_full/final_compact \
76
+ --manifest runs/hf_inference_optimization/mmau_pair_manifest.jsonl \
77
+ --repeats 3 --device cuda:0 \
78
+ --output runs/hf_inference_optimization/mmau_latency.json
79
+ ```
80
+
81
+ 测量包含前处理、设备传输、贪心生成至 EOS/token 上限、文本解码和 CUDA 同步。音频文件读取及重采样在计时外完成。TTFT 在独立的 streaming 测��中从前处理开始,至第一个新 token 返回。解码吞吐量为 `(生成 token 数 − 1) / (最后 token 时间 − 首 token 时间)`,端到端吞吐量为 `生成 token 数 / 整答时延`;均包含 EOS。普通整答计时不插入 streamer。
82
+
83
+ ## 限制
84
+
85
+ 等价基准是原 checkpoint 按上传脚本的 `dtype="auto"` 加载后的模型;精简权重不保留原文件额外的 FP32 小数精度,原始目录可继续用于 FP32 实验。首次 CUDA Graph 建立和新 batch shape 需要预热。最终测量时机器有其他 GPU 工作负载,因此绝对时延受资源争用影响;同一样本的两模型在同一 GPU 上交替执行。绝对速度还取决于硬件、CUDA/PyTorch 和 kernel 依赖,实际运行环境及其他 GPU 进程均记录在测量 JSON 中。MMAU 短答案通常只有答案字母和 EOS,长回答吞吐量应参考独立的文本测试。
LICENSE.codec ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
6
+
7
+ 1. Definitions.
8
+
9
+ "License" shall mean the terms and conditions for use, reproduction,
10
+ and distribution as defined by Sections 1 through 9 of this document.
11
+
12
+ "Licensor" shall mean the copyright owner or entity authorized by
13
+ the copyright owner that is granting the License.
14
+
15
+ "Legal Entity" shall mean the union of the acting entity and all
16
+ other entities that control, are controlled by, or are under common
17
+ control with that entity. For the purposes of this definition,
18
+ "control" means (i) the power, direct or indirect, to cause the
19
+ direction or management of such entity, whether by contract or
20
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
21
+ outstanding shares, or (iii) beneficial ownership of such entity.
22
+
23
+ "You" (or "Your") shall mean an individual or Legal Entity
24
+ exercising permissions granted by this License.
25
+
26
+ "Source" form shall mean the preferred form for making modifications,
27
+ including but not limited to software source code, documentation
28
+ source, and configuration files.
29
+
30
+ "Object" form shall mean any form resulting from mechanical
31
+ transformation or translation of a Source form, including but
32
+ not limited to compiled object code, generated documentation,
33
+ and conversions to other media types.
34
+
35
+ "Work" shall mean the work of authorship, whether in Source or
36
+ Object form, made available under the License, as indicated by a
37
+ copyright notice that is included in or attached to the work
38
+ (an example is provided in the Appendix below).
39
+
40
+ "Derivative Works" shall mean any work, whether in Source or Object
41
+ form, that is based on (or derived from) the Work and for which the
42
+ editorial revisions, annotations, elaborations, or other modifications
43
+ represent, as a whole, an original work of authorship. For the purposes
44
+ of this License, Derivative Works shall not include works that remain
45
+ separable from, or merely link (or bind by name) to the interfaces of,
46
+ the Work and Derivative Works thereof.
47
+
48
+ "Contribution" shall mean any work of authorship, including
49
+ the original version of the Work and any modifications or additions
50
+ to that Work or Derivative Works thereof, that is intentionally
51
+ submitted to Licensor for inclusion in the Work by the copyright owner
52
+ or by an individual or Legal Entity authorized to submit on behalf of
53
+ the copyright owner. For the purposes of this definition, "submitted"
54
+ means any form of electronic, verbal, or written communication sent
55
+ to the Licensor or its representatives, including but not limited to
56
+ communication on electronic mailing lists, source code control systems,
57
+ and issue tracking systems that are managed by, or on behalf of, the
58
+ Licensor for the purpose of discussing and improving the Work, but
59
+ excluding communication that is conspicuously marked or otherwise
60
+ designated in writing by the copyright owner as "Not a Contribution."
61
+
62
+ "Contributor" shall mean Licensor and any individual or Legal Entity
63
+ on behalf of whom a Contribution has been received by Licensor and
64
+ subsequently incorporated within the Work.
65
+
66
+ 2. Grant of Copyright License. Subject to the terms and conditions of
67
+ this License, each Contributor hereby grants to You a perpetual,
68
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
69
+ copyright license to reproduce, prepare Derivative Works of,
70
+ publicly display, publicly perform, sublicense, and distribute the
71
+ Work and such Derivative Works in Source or Object form.
72
+
73
+ 3. Grant of Patent License. Subject to the terms and conditions of
74
+ this License, each Contributor hereby grants to You a perpetual,
75
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
76
+ (except as stated in this section) patent license to make, have made,
77
+ use, offer to sell, sell, import, and otherwise transfer the Work,
78
+ where such license applies only to those patent claims licensable
79
+ by such Contributor that are necessarily infringed by their
80
+ Contribution(s) alone or by combination of their Contribution(s)
81
+ with the Work to which such Contribution(s) was submitted. If You
82
+ institute patent litigation against any entity (including a
83
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
84
+ or a Contribution incorporated within the Work constitutes direct
85
+ or contributory patent infringement, then any patent licenses
86
+ granted to You under this License for that Work shall terminate
87
+ as of the date such litigation is filed.
88
+
89
+ 4. Redistribution. You may reproduce and distribute copies of the
90
+ Work or Derivative Works thereof in any medium, with or without
91
+ modifications, and in Source or Object form, provided that You
92
+ meet the following conditions:
93
+
94
+ (a) You must give any other recipients of the Work or
95
+ Derivative Works a copy of this License; and
96
+
97
+ (b) You must cause any modified files to carry prominent notices
98
+ stating that You changed the files; and
99
+
100
+ (c) You must retain, in the Source form of any Derivative Works
101
+ that You distribute, all copyright, patent, trademark, and
102
+ attribution notices from the Source form of the Work,
103
+ excluding those notices that do not pertain to any part of
104
+ the Derivative Works; and
105
+
106
+ (d) If the Work includes a "NOTICE" text file as part of its
107
+ distribution, then any Derivative Works that You distribute must
108
+ include a readable copy of the attribution notices contained
109
+ within such NOTICE file, excluding those notices that do not
110
+ pertain to any part of the Derivative Works, in at least one
111
+ of the following places: within a NOTICE text file distributed
112
+ as part of the Derivative Works; within the Source form or
113
+ documentation, if provided along with the Derivative Works; or,
114
+ within a display generated by the Derivative Works, if and
115
+ wherever such third-party notices normally appear. The contents
116
+ of the NOTICE file are for informational purposes only and
117
+ do not modify the License. You may add Your own attribution
118
+ notices within Derivative Works that You distribute, alongside
119
+ or as an addendum to the NOTICE text from the Work, provided
120
+ that such additional attribution notices cannot be construed
121
+ as modifying the License.
122
+
123
+ You may add Your own copyright statement to Your modifications and
124
+ may provide additional or different license terms and conditions
125
+ for use, reproduction, or distribution of Your modifications, or
126
+ for any such Derivative Works as a whole, provided Your use,
127
+ reproduction, and distribution of the Work otherwise complies with
128
+ the conditions stated in this License.
129
+
130
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
131
+ any Contribution intentionally submitted for inclusion in the Work
132
+ by You to the Licensor shall be under the terms and conditions of
133
+ this License, without any additional terms or conditions.
134
+ Notwithstanding the above, nothing herein shall supersede or modify
135
+ the terms of any separate license agreement you may have executed
136
+ with Licensor regarding such Contributions.
137
+
138
+ 6. Trademarks. This License does not grant permission to use the trade
139
+ names, trademarks, service marks, or product names of the Licensor,
140
+ except as required for reasonable and customary use in describing the
141
+ origin of the Work and reproducing the content of the NOTICE file.
142
+
143
+ 7. Disclaimer of Warranty. Unless required by applicable law or
144
+ agreed to in writing, Licensor provides the Work (and each
145
+ Contributor provides its Contributions) on an "AS IS" BASIS,
146
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
147
+ implied, including, without limitation, any warranties or conditions
148
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
149
+ PARTICULAR PURPOSE. You are solely responsible for determining the
150
+ appropriateness of using or redistributing the Work and assume any
151
+ risks associated with Your exercise of permissions under this License.
152
+
153
+ 8. Limitation of Liability. In no event and under no legal theory,
154
+ whether in tort (including negligence), contract, or otherwise,
155
+ unless required by applicable law (such as deliberate and grossly
156
+ negligent acts) or agreed to in writing, shall any Contributor be
157
+ liable to You for damages, including any direct, indirect, special,
158
+ incidental, or consequential damages of any character arising as a
159
+ result of this License or out of the use or inability to use the
160
+ Work (including but not limited to damages for loss of goodwill,
161
+ work stoppage, computer failure or malfunction, or any and all
162
+ other commercial damages or losses), even if such Contributor
163
+ has been advised of the possibility of such damages.
164
+
165
+ 9. Accepting Warranty or Additional Liability. While redistributing
166
+ the Work or Derivative Works thereof, You may choose to offer,
167
+ and charge a fee for, acceptance of support, warranty, indemnity,
168
+ or other liability obligations and/or rights consistent with this
169
+ License. However, in accepting such obligations, You may act only
170
+ on Your own behalf and on Your sole responsibility, not on behalf
171
+ of any other Contributor, and only if You agree to indemnify,
172
+ defend, and hold each Contributor harmless for any liability
173
+ incurred by, or claims asserted against, such Contributor by reason
174
+ of your accepting any such warranty or additional liability.
175
+
176
+ END OF TERMS AND CONDITIONS
177
+
178
+ APPENDIX: How to apply the Apache License to your work.
179
+
180
+ To apply the Apache License to your work, attach the following
181
+ boilerplate notice, with the fields enclosed by brackets "[]"
182
+ replaced with your own identifying information. (Don't include
183
+ the brackets!) The text should be enclosed in the appropriate
184
+ comment syntax for the file format. We also recommend that a
185
+ file or class name and description of purpose be included on the
186
+ same "printed page" as the copyright notice for easier
187
+ identification within third-party archives.
188
+
189
+ Copyright 2026 Alibaba Cloud
190
+
191
+ Licensed under the Apache License, Version 2.0 (the "License");
192
+ you may not use this file except in compliance with the License.
193
+ You may obtain a copy of the License at
194
+
195
+ http://www.apache.org/licenses/LICENSE-2.0
196
+
197
+ Unless required by applicable law or agreed to in writing, software
198
+ distributed under the License is distributed on an "AS IS" BASIS,
199
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
200
+ See the License for the specific language governing permissions and
201
+ limitations under the License.
README.md ADDED
@@ -0,0 +1,72 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ library_name: transformers
3
+ pipeline_tag: audio-text-to-text
4
+ language:
5
+ - zh
6
+ - en
7
+ tags:
8
+ - edgeinstant
9
+ - audio
10
+ - speech-recognition
11
+ - speech-translation
12
+ - audio-question-answering
13
+ ---
14
+
15
+ # EdgeInstant-1.5b S5 bilingual audio instruction
16
+
17
+ 中英文 audio-in 后训练模型,支持语音转写、语音翻译、音频问答和指令执行,同时保留 direct 与 thinking 两种输出模式。此目录是包含模型权重、处理器及自定义代码的独立 Hugging Face 包。
18
+
19
+ ## 训练
20
+
21
+ 起点为 S4 最终模型,先以相同数据进行 600 步回答控制训练,再进行 1,600 步八卡混合后训练。语言模型全部参数、音频投影层和音频编码器末两层参与主训练;峰值学习率分别为 1e-6、2e-7、1e-7。
22
+
23
+ 主训练清单包含 531,174 行,涵盖中英文 ASR、双向语音翻译、多来源文本指令、真实语音命令提取、语音语境问答、声音和音乐理解、直接回答及思考数据。正监督中 thinking 占 3.76%。通过 EOS 权重 4 和重复负样本 unlikelihood 权重 0.2 改善回答结束及循环输出;推理使用 greedy。
24
+
25
+ 交付权重为主训练第 1,600 步。模型依据独立开发集上的生成稳定性、准确率、ASR 和翻译表现选择。完整配置见仓库 `configs/train/bilingual_s5.yaml`,方法与数据说明见 `docs/BILINGUAL_S5.md`。测试样本仅用于评测。
26
+
27
+ ## 使用
28
+
29
+ 安装本目录 `requirements.txt`。音频输入为 16 kHz 单声道;保留完整波形,不设 30 秒截断。
30
+
31
+ ```python
32
+ import soundfile as sf
33
+ import torch
34
+ from transformers import AutoModel, AutoProcessor
35
+
36
+ path = "/default-filesys/workspace/chenjunzhe/EdgeInstant-1.5b/runs/bilingual_s5_full/final"
37
+ processor = AutoProcessor.from_pretrained(path, trust_remote_code=True)
38
+ model = AutoModel.from_pretrained(
39
+ path, trust_remote_code=True, dtype=torch.bfloat16,
40
+ ).to("cuda").eval()
41
+ torch.backends.cuda.enable_cudnn_sdp(False)
42
+
43
+ waveform, sample_rate = sf.read("speech.wav", dtype="float32")
44
+ inputs = processor(
45
+ audio=waveform, sampling_rate=sample_rate, task="asr",
46
+ text="请将语音转写为文字,只输出转写结果。", enable_thinking=False,
47
+ ).to("cuda")
48
+ with torch.inference_mode():
49
+ output = model.generate(**inputs, max_new_tokens=512, do_sample=False)
50
+ answer = output[0, inputs["input_ids"].shape[1]:]
51
+ print(processor.decode(answer, skip_special_tokens=True))
52
+ ```
53
+
54
+ 语音翻译或问答使用 `task="qa"`,并通过 `text` 指定请求;直接执行语音中的指令时,可用“请听取并完成语音中的请求。”。文本输入使用 `processor(text=...)`。需要思考模式时设置 `enable_thinking=True`,为思考和最终答案留足输出 token。要求只输出选项时使用 `enable_thinking=False` 并在请求中明确输出格式。
55
+
56
+ ## 评测
57
+
58
+ 评测结论及已知限制见 [RESULTS.md](../RESULTS.md),27 项同协议完整测试与 S4 对照见 [final_results.md](../final_results.md),固定开发集见 [dev_comparison.md](../dev_comparison.md)。CER/WER 和正确率在结果表中使用百分数,BLEU 保持原单位。W&B 评测仅记录数值指标,预测保留在本地。
59
+
60
+ 真人语音命令在官方未见说话人验证集上,每种输出视图各 3,118 条:JSON 对象完全正确率 93.49%,单字段准确率 96.79%。这些结果衡量 FSC 字段约定下的语音命令提取。公开语境问答采用数据集原有参考,部分来源由模型生成标注。
61
+
62
+ ## 能力限制
63
+
64
+ 开放式复杂语音推理、口述数字计算和中译英语音翻译仍有不足。输出终止、格式正确不等于内容正确;thinking 模式也不保证准确率更高。极低音量或信息不足的语音可能产生无依据转写。新增中文口述数字开发集仅一条,不能用于估计整体中文计算能力。
65
+
66
+ 本轮训练与评测针对 audio-in 和文本回答;Talker 与语音合成组件沿用起点权重。模型包含 Qwen 组件,各组件及训练数据的原许可仍适用;解码器代码许可见 `LICENSE.codec`。
67
+
68
+ ## 精简推理版本
69
+
70
+ 本目录保留全部多模态能力,采用 BF16 推理权重、FP32 音频投影器与特殊 token 参数,移除 codec 未使用的输入投影。CUDA 单 token 解码自动对线性注意力 decoder 层使用 CUDA Graph,模型加载与 `generate()` 调用方式不变。原始 checkpoint 按 `dtype="auto"` 加载时,MMAU 1,000 条生成 token 完全一致,准确率均为 64.9%。
71
+
72
+ 参数量、权重体积、TTFT、吞吐量、验证范围和复现命令见 [推理优化说明](INFERENCE_OPTIMIZATION.md)。
chat_template.jinja ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {%- set image_count = namespace(value=0) %}
2
+ {%- set video_count = namespace(value=0) %}
3
+ {%- macro render_content(content, do_vision_count, is_system_content=false) %}
4
+ {%- if content is string %}
5
+ {{- content }}
6
+ {%- elif content is iterable and content is not mapping %}
7
+ {%- for item in content %}
8
+ {%- if 'image' in item or 'image_url' in item or item.type == 'image' %}
9
+ {%- if is_system_content %}
10
+ {{- raise_exception('System message cannot contain images.') }}
11
+ {%- endif %}
12
+ {%- if do_vision_count %}
13
+ {%- set image_count.value = image_count.value + 1 %}
14
+ {%- endif %}
15
+ {%- if add_vision_id %}
16
+ {{- 'Picture ' ~ image_count.value ~ ': ' }}
17
+ {%- endif %}
18
+ {{- '<|vision_start|><|image_pad|><|vision_end|>' }}
19
+ {%- elif 'video' in item or item.type == 'video' %}
20
+ {%- if is_system_content %}
21
+ {{- raise_exception('System message cannot contain videos.') }}
22
+ {%- endif %}
23
+ {%- if do_vision_count %}
24
+ {%- set video_count.value = video_count.value + 1 %}
25
+ {%- endif %}
26
+ {%- if add_vision_id %}
27
+ {{- 'Video ' ~ video_count.value ~ ': ' }}
28
+ {%- endif %}
29
+ {{- '<|vision_start|><|video_pad|><|vision_end|>' }}
30
+ {%- elif 'text' in item %}
31
+ {{- item.text }}
32
+ {%- else %}
33
+ {{- raise_exception('Unexpected item type in content.') }}
34
+ {%- endif %}
35
+ {%- endfor %}
36
+ {%- elif content is none or content is undefined %}
37
+ {{- '' }}
38
+ {%- else %}
39
+ {{- raise_exception('Unexpected content type.') }}
40
+ {%- endif %}
41
+ {%- endmacro %}
42
+ {%- if not messages %}
43
+ {{- raise_exception('No messages provided.') }}
44
+ {%- endif %}
45
+ {%- if tools and tools is iterable and tools is not mapping %}
46
+ {{- '<|im_start|>system\n' }}
47
+ {{- "# Tools\n\nYou have access to the following functions:\n\n<tools>" }}
48
+ {%- for tool in tools %}
49
+ {{- "\n" }}
50
+ {{- tool | tojson }}
51
+ {%- endfor %}
52
+ {{- "\n</tools>" }}
53
+ {{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>' }}
54
+ {%- if messages[0].role == 'system' %}
55
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
56
+ {%- if content %}
57
+ {{- '\n\n' + content }}
58
+ {%- endif %}
59
+ {%- endif %}
60
+ {{- '<|im_end|>\n' }}
61
+ {%- else %}
62
+ {%- if messages[0].role == 'system' %}
63
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
64
+ {{- '<|im_start|>system\n' + content + '<|im_end|>\n' }}
65
+ {%- endif %}
66
+ {%- endif %}
67
+ {%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
68
+ {%- for message in messages[::-1] %}
69
+ {%- set index = (messages|length - 1) - loop.index0 %}
70
+ {%- if ns.multi_step_tool and message.role == "user" %}
71
+ {%- set content = render_content(message.content, false)|trim %}
72
+ {%- if not(content.startswith('<tool_response>') and content.endswith('</tool_response>')) %}
73
+ {%- set ns.multi_step_tool = false %}
74
+ {%- set ns.last_query_index = index %}
75
+ {%- endif %}
76
+ {%- endif %}
77
+ {%- endfor %}
78
+ {%- if ns.multi_step_tool %}
79
+ {{- raise_exception('No user query found in messages.') }}
80
+ {%- endif %}
81
+ {%- for message in messages %}
82
+ {%- set content = render_content(message.content, true)|trim %}
83
+ {%- if message.role == "system" %}
84
+ {%- if not loop.first %}
85
+ {{- raise_exception('System message must be at the beginning.') }}
86
+ {%- endif %}
87
+ {%- elif message.role == "user" %}
88
+ {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
89
+ {%- elif message.role == "assistant" %}
90
+ {%- set reasoning_content = '' %}
91
+ {%- if message.reasoning_content is string %}
92
+ {%- set reasoning_content = message.reasoning_content %}
93
+ {%- else %}
94
+ {%- if '</think>' in content %}
95
+ {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
96
+ {%- set content = content.split('</think>')[-1].lstrip('\n') %}
97
+ {%- endif %}
98
+ {%- endif %}
99
+ {%- set reasoning_content = reasoning_content|trim %}
100
+ {%- if loop.index0 > ns.last_query_index %}
101
+ {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content + '\n</think>\n\n' + content }}
102
+ {%- else %}
103
+ {{- '<|im_start|>' + message.role + '\n' + content }}
104
+ {%- endif %}
105
+ {%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
106
+ {%- for tool_call in message.tool_calls %}
107
+ {%- if tool_call.function is defined %}
108
+ {%- set tool_call = tool_call.function %}
109
+ {%- endif %}
110
+ {%- if loop.first %}
111
+ {%- if content|trim %}
112
+ {{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
113
+ {%- else %}
114
+ {{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
115
+ {%- endif %}
116
+ {%- else %}
117
+ {{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
118
+ {%- endif %}
119
+ {%- if tool_call.arguments is defined %}
120
+ {%- for args_name, args_value in tool_call.arguments|items %}
121
+ {{- '<parameter=' + args_name + '>\n' }}
122
+ {%- set args_value = args_value | tojson | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %}
123
+ {{- args_value }}
124
+ {{- '\n</parameter>\n' }}
125
+ {%- endfor %}
126
+ {%- endif %}
127
+ {{- '</function>\n</tool_call>' }}
128
+ {%- endfor %}
129
+ {%- endif %}
130
+ {{- '<|im_end|>\n' }}
131
+ {%- elif message.role == "tool" %}
132
+ {%- if loop.previtem and loop.previtem.role != "tool" %}
133
+ {{- '<|im_start|>user' }}
134
+ {%- endif %}
135
+ {{- '\n<tool_response>\n' }}
136
+ {{- content }}
137
+ {{- '\n</tool_response>' }}
138
+ {%- if not loop.last and loop.nextitem.role != "tool" %}
139
+ {{- '<|im_end|>\n' }}
140
+ {%- elif loop.last %}
141
+ {{- '<|im_end|>\n' }}
142
+ {%- endif %}
143
+ {%- else %}
144
+ {{- raise_exception('Unexpected message role.') }}
145
+ {%- endif %}
146
+ {%- endfor %}
147
+ {%- if add_generation_prompt %}
148
+ {{- '<|im_start|>assistant\n' }}
149
+ {%- if enable_thinking is defined and enable_thinking is true %}
150
+ {{- '<think>\n' }}
151
+ {%- else %}
152
+ {{- '<think>\n\n</think>\n\n' }}
153
+ {%- endif %}
154
+ {%- endif %}
codec_edgeinstant.py ADDED
@@ -0,0 +1,233 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 The Alibaba Qwen team.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """Qwen3-TTS 12.5 Hz waveform decoder for the EdgeInstant model."""
4
+
5
+ from __future__ import annotations
6
+
7
+ import json
8
+ import math
9
+ from pathlib import Path
10
+
11
+ import torch
12
+ from torch import nn
13
+ from torch.nn import functional as F
14
+ from transformers.models.qwen3_omni_moe.configuration_qwen3_omni_moe import (
15
+ Qwen3OmniMoeCode2WavConfig,
16
+ )
17
+ from transformers.models.qwen3_omni_moe.modeling_qwen3_omni_moe import (
18
+ Qwen3OmniMoeCausalConvNet,
19
+ Qwen3OmniMoeCode2WavDecoderResidualUnit,
20
+ Qwen3OmniMoeCode2WavTransformerModel,
21
+ Qwen3OmniMoeConvNeXtBlock,
22
+ SnakeBeta,
23
+ )
24
+
25
+
26
+ class _CausalTransConvNet(nn.Module):
27
+ def __init__(self, in_channels, out_channels, kernel_size, stride=1):
28
+ super().__init__()
29
+ self.conv = nn.ConvTranspose1d(in_channels, out_channels, kernel_size, stride=stride)
30
+ self.right_pad = kernel_size - stride
31
+
32
+ def forward(self, hidden):
33
+ hidden = self.conv(hidden)
34
+ if self.right_pad:
35
+ hidden = hidden[..., : -self.right_pad]
36
+ return hidden.contiguous()
37
+
38
+
39
+ class _Codebook(nn.Module):
40
+ def __init__(self, dim, size):
41
+ super().__init__()
42
+ self.cluster_usage = nn.Parameter(torch.ones(size))
43
+ self.embedding_sum = nn.Parameter(torch.zeros(size, dim))
44
+
45
+ def forward(self, codes):
46
+ embedding = self.embedding_sum / self.cluster_usage.clamp(min=1e-5)[:, None]
47
+ return F.embedding(codes, embedding)
48
+
49
+
50
+ class _VectorQuantization(nn.Module):
51
+ def __init__(self, dim, size):
52
+ super().__init__()
53
+ self._codebook = _Codebook(dim, size)
54
+
55
+ def forward(self, codes):
56
+ return self._codebook(codes).transpose(1, 2)
57
+
58
+
59
+ class _ResidualVectorQuantization(nn.Module):
60
+ def __init__(self, count, dim, size):
61
+ super().__init__()
62
+ self.layers = nn.ModuleList([_VectorQuantization(dim, size) for _ in range(count)])
63
+
64
+ def forward(self, codes):
65
+ quantized = torch.zeros([1], device=codes.device)[0]
66
+ for index, layer in enumerate(self.layers):
67
+ quantized = quantized + layer(codes[:, index])
68
+ return quantized
69
+
70
+
71
+ class _ResidualVectorQuantizer(nn.Module):
72
+ def __init__(self, count, config, *, inference_only=False):
73
+ super().__init__()
74
+ dim = config.codebook_dim // 2
75
+ self.input_proj = (
76
+ None if inference_only else nn.Conv1d(config.codebook_dim, dim, 1, bias=False)
77
+ )
78
+ self.output_proj = nn.Conv1d(dim, config.codebook_dim, 1, bias=False)
79
+ self.vq = _ResidualVectorQuantization(count, dim, config.codebook_size)
80
+
81
+ def forward(self, codes):
82
+ return self.output_proj(self.vq(codes))
83
+
84
+
85
+ class _SplitResidualVectorQuantizer(nn.Module):
86
+ def __init__(self, config, *, inference_only=False):
87
+ super().__init__()
88
+ self.rvq_first = _ResidualVectorQuantizer(1, config, inference_only=inference_only)
89
+ self.rvq_rest = _ResidualVectorQuantizer(
90
+ config.num_quantizers - 1, config, inference_only=inference_only,
91
+ )
92
+
93
+ def forward(self, codes):
94
+ quantized = self.rvq_first(codes[:, :1])
95
+ quantized += self.rvq_rest(codes[:, 1:])
96
+ return quantized
97
+
98
+
99
+ class _ProjectedTransformer(Qwen3OmniMoeCode2WavTransformerModel):
100
+ def __init__(self, config):
101
+ super().__init__(config)
102
+ self.input_proj = nn.Linear(config.latent_dim, config.hidden_size)
103
+ self.output_proj = nn.Linear(config.hidden_size, config.latent_dim)
104
+
105
+ def forward(self, hidden):
106
+ positions = torch.arange(hidden.shape[1], device=hidden.device)
107
+ distances = positions[:, None] - positions[None, :]
108
+ mask = ((distances >= 0) & (distances < self.config.sliding_window))[None, None]
109
+ output = super().forward(
110
+ inputs_embeds=self.input_proj(hidden),
111
+ attention_mask={"sliding_attention": mask.expand(hidden.shape[0], -1, -1, -1)},
112
+ use_cache=False,
113
+ )
114
+ return self.output_proj(output.last_hidden_state)
115
+
116
+
117
+ class _DecoderBlock(nn.Module):
118
+ def __init__(self, config, index):
119
+ super().__init__()
120
+ input_dim = config.decoder_dim // 2**index
121
+ output_dim = input_dim // 2
122
+ rate = config.upsample_rates[index]
123
+ self.block = nn.ModuleList(
124
+ [SnakeBeta(input_dim), _CausalTransConvNet(input_dim, output_dim, 2 * rate, rate)]
125
+ + [Qwen3OmniMoeCode2WavDecoderResidualUnit(output_dim, dilation) for dilation in (1, 3, 9)]
126
+ )
127
+
128
+ def forward(self, hidden):
129
+ for block in self.block:
130
+ hidden = block(hidden)
131
+ return hidden
132
+
133
+
134
+ class EdgeInstantCodecDecoder(nn.Module):
135
+ """Decode integer codes ``[batch, frames, 16]`` to float audio ``[batch, samples]``.
136
+
137
+ ``config`` is the complete Qwen3-TTS speech tokenizer configuration. Weights
138
+ match its ``decoder.*`` tensors after removing that prefix. Setting its
139
+ top-level ``inference_only`` field omits the unused quantizer input projections.
140
+ """
141
+
142
+ def __init__(self, config: dict):
143
+ super().__init__()
144
+ self.sample_rate = int(config["output_sample_rate"])
145
+ self.samples_per_frame = int(config["decode_upsample_rate"])
146
+ decoder_config = dict(config["decoder_config"])
147
+ decoder_config["rope_parameters"] = {
148
+ "rope_type": "default",
149
+ "rope_theta": decoder_config.pop("rope_theta", 10000),
150
+ }
151
+ self.config = Qwen3OmniMoeCode2WavConfig(**decoder_config)
152
+ self.config._attn_implementation = "sdpa"
153
+ cfg = self.config
154
+ if math.prod([*cfg.upsample_rates, *cfg.upsampling_ratios]) != self.samples_per_frame:
155
+ raise ValueError("decode_upsample_rate must match the waveform decoder upsampling factors")
156
+ self.pre_transformer = _ProjectedTransformer(cfg)
157
+ self.quantizer = _SplitResidualVectorQuantizer(
158
+ cfg, inference_only=bool(config.get("inference_only", False)),
159
+ )
160
+ self.pre_conv = Qwen3OmniMoeCausalConvNet(cfg.codebook_dim, cfg.latent_dim, 3)
161
+ self.upsample = nn.ModuleList(
162
+ [
163
+ nn.ModuleList(
164
+ [
165
+ _CausalTransConvNet(cfg.latent_dim, cfg.latent_dim, factor, factor),
166
+ Qwen3OmniMoeConvNeXtBlock(cfg.latent_dim),
167
+ ]
168
+ )
169
+ for factor in cfg.upsampling_ratios
170
+ ]
171
+ )
172
+ output_dim = cfg.decoder_dim // 2 ** len(cfg.upsample_rates)
173
+ self.decoder = nn.ModuleList(
174
+ [Qwen3OmniMoeCausalConvNet(cfg.latent_dim, cfg.decoder_dim, 7)]
175
+ + [_DecoderBlock(cfg, index) for index in range(len(cfg.upsample_rates))]
176
+ + [SnakeBeta(output_dim), Qwen3OmniMoeCausalConvNet(output_dim, 1, 7)]
177
+ )
178
+
179
+ def forward(self, codes: torch.Tensor) -> torch.Tensor:
180
+ if codes.ndim != 3 or codes.shape[-1] != self.config.num_quantizers:
181
+ raise ValueError(f"codes must be [batch, frames, {self.config.num_quantizers}]")
182
+ if codes.shape[1] == 0:
183
+ return self.pre_conv.conv.weight.new_empty((codes.shape[0], 0), dtype=torch.float32)
184
+ codes = codes.to(device=self.pre_conv.conv.weight.device, dtype=torch.long)
185
+ hidden = self.quantizer(codes.transpose(1, 2))
186
+ hidden = self.pre_conv(hidden).transpose(1, 2)
187
+ hidden = self.pre_transformer(hidden).transpose(1, 2)
188
+ for blocks in self.upsample:
189
+ for block in blocks:
190
+ hidden = block(hidden)
191
+ for block in self.decoder:
192
+ hidden = block(hidden)
193
+ return hidden.squeeze(1).clamp(min=-1, max=1).float()
194
+
195
+ def decode(
196
+ self, codes: torch.Tensor, *, chunk_size: int = 300, left_context_frames: int = 25
197
+ ) -> torch.Tensor:
198
+ """Decode with the speech tokenizer's chunking and left context convention."""
199
+ if chunk_size <= 0 or left_context_frames < 0:
200
+ raise ValueError("chunk_size must be positive and left_context_frames nonnegative")
201
+ if codes.shape[1] == 0:
202
+ return self(codes)
203
+ waves = []
204
+ for start in range(0, codes.shape[1], chunk_size):
205
+ context = min(start, left_context_frames)
206
+ wave = self(codes[:, start - context : start + chunk_size])
207
+ waves.append(wave[:, context * self.samples_per_frame :])
208
+ return torch.cat(waves, dim=-1)
209
+
210
+
211
+ def load_codec_decoder(
212
+ directory: str | Path, *, dtype: torch.dtype = torch.float32, device: str = "cpu"
213
+ ) -> EdgeInstantCodecDecoder:
214
+ """Load the decoder tensors from a local Qwen3-TTS speech tokenizer directory."""
215
+ from safetensors import safe_open
216
+
217
+ directory = Path(directory)
218
+ with (directory / "config.json").open() as handle:
219
+ config = json.load(handle)
220
+ decoder = EdgeInstantCodecDecoder(config).to(dtype=dtype, device=device)
221
+ rotary = decoder.pre_transformer.rotary_emb
222
+ rotary.inv_freq, rotary.attention_scaling = rotary.compute_default_rope_parameters(
223
+ rotary.config, torch.device(device)
224
+ )
225
+ rotary.original_inv_freq = rotary.inv_freq.clone()
226
+ with safe_open(directory / "model.safetensors", framework="pt", device="cpu") as checkpoint:
227
+ state = {
228
+ key.removeprefix("decoder."): checkpoint.get_tensor(key)
229
+ for key in checkpoint.keys()
230
+ if key.startswith("decoder.")
231
+ }
232
+ decoder.load_state_dict(state, strict=True)
233
+ return decoder.eval()
compact_native_clock.py ADDED
@@ -0,0 +1,174 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Compact inference representation of the frozen native-token acoustic model."""
2
+ from __future__ import annotations
3
+
4
+ import copy
5
+ from pathlib import Path
6
+
7
+ import torch
8
+ from torch import nn
9
+
10
+ from .native_clock_talker import NativeClockTalker
11
+ from .residual_predictor import Qwen3TTSResidualPredictor
12
+
13
+
14
+ class ProjectedTextEmbedding(nn.Module):
15
+ """Keep projected rows reachable from the complete native byte vocabulary."""
16
+
17
+ def __init__(self, weight: torch.Tensor, teacher_ids: torch.Tensor, vocabulary_size: int):
18
+ super().__init__()
19
+ self.embedding = nn.Embedding.from_pretrained(weight, freeze=True)
20
+ rows = torch.full((vocabulary_size,), -1, device=weight.device, dtype=torch.long)
21
+ rows[teacher_ids.to(weight.device)] = torch.arange(len(teacher_ids), device=weight.device)
22
+ self.register_buffer("teacher_to_row", rows)
23
+
24
+ def forward(self, teacher_ids: torch.Tensor) -> torch.Tensor:
25
+ return self.embedding(self.teacher_to_row[teacher_ids])
26
+
27
+
28
+ class CompactNativeClockTalker(NativeClockTalker):
29
+ """Fold frozen text projection and share identical codec input embeddings."""
30
+
31
+ serialization_buffers = ("prefix", "prefix_english", "prefix_auto", "tts_eos", "tts_pad")
32
+ tied_weight_keys = {
33
+ "residual_predictor.q0_embedding.weight": "codec_history_embeddings.0.weight",
34
+ **{
35
+ f"residual_predictor.code_embeddings.{group}.weight":
36
+ f"codec_history_embeddings.{group + 1}.weight"
37
+ for group in range(14)
38
+ },
39
+ }
40
+
41
+ def prepare_for_serialization(self) -> "CompactNativeClockTalker":
42
+ """Save speaker conditions and the runtime precision of rotary frequencies."""
43
+ self._non_persistent_buffers_set.difference_update(self.serialization_buffers)
44
+ for backbone in (self.backbone, self.residual_predictor.backbone):
45
+ backbone.rotary_emb._non_persistent_buffers_set.difference_update(
46
+ ("inv_freq", "original_inv_freq")
47
+ )
48
+ return self
49
+
50
+ def get_config(self) -> dict:
51
+ """Describe the compact model without paths to its initialization assets."""
52
+ residual_config = self.residual_predictor.backbone.config.to_dict()
53
+ residual_config["rope_theta"] = self.residual_predictor.backbone.config.rope_parameters["rope_theta"]
54
+ return {
55
+ "backbone_config": self.backbone.config.to_dict(),
56
+ "residual_predictor_config": residual_config,
57
+ "attn_implementation": self.backbone.config._attn_implementation,
58
+ "hidden_size": self.hidden_size,
59
+ "codebook_size": self.codebook_size,
60
+ "num_code_groups": len(self.codec_history_embeddings),
61
+ "projected_num_embeddings": self.text_embedding.embedding.num_embeddings,
62
+ "teacher_vocab_size": self.text_embedding.teacher_to_row.numel(),
63
+ "prefix_shapes": {
64
+ name: list(getattr(self, name).shape) for name in self.serialization_buffers
65
+ },
66
+ "native_teacher_ids": list(self.native_teacher_ids),
67
+ "native_teacher_offsets": list(self.native_teacher_offsets),
68
+ }
69
+
70
+ @classmethod
71
+ def from_config(
72
+ cls, config: dict, *, attn_implementation: str | None = None,
73
+ ) -> "CompactNativeClockTalker":
74
+ """Build the compact modules and shared weights for checkpoint loading."""
75
+ from transformers import Qwen3Config, Qwen3Model
76
+
77
+ model = cls()
78
+ model.hidden_size = int(config["hidden_size"])
79
+ model.codebook_size = int(config["codebook_size"])
80
+ model.eos_class = model.codebook_size
81
+ backbone_config = Qwen3Config(**config["backbone_config"])
82
+ backbone_config._attn_implementation = attn_implementation or config["attn_implementation"]
83
+ model.backbone = Qwen3Model(backbone_config)
84
+ model.backbone.embed_tokens = None
85
+ model.text_embedding = ProjectedTextEmbedding(
86
+ torch.empty(int(config["projected_num_embeddings"]), model.hidden_size),
87
+ torch.empty(0, dtype=torch.long), int(config["teacher_vocab_size"]),
88
+ )
89
+ model.text_projection = nn.Identity()
90
+ model.codec_history_embeddings = nn.ModuleList(
91
+ nn.Embedding(model.codebook_size, model.hidden_size)
92
+ for _ in range(int(config["num_code_groups"]))
93
+ )
94
+ model.codec_bos = nn.Parameter(torch.empty(model.hidden_size))
95
+ model.q0_head = nn.Linear(model.hidden_size, model.codebook_size, bias=False)
96
+ model.codec_eos_head = nn.Linear(model.hidden_size, 1, bias=False)
97
+ model.residual_predictor = Qwen3TTSResidualPredictor(
98
+ model.hidden_size, model.codebook_size, config["residual_predictor_config"],
99
+ )
100
+ model.residual_predictor.q0_embedding = model.codec_history_embeddings[0]
101
+ for group in range(len(model.residual_predictor.code_embeddings)):
102
+ model.residual_predictor.code_embeddings[group] = model.codec_history_embeddings[group + 1]
103
+ model.native_teacher_ids = list(config["native_teacher_ids"])
104
+ model.native_teacher_offsets = list(config["native_teacher_offsets"])
105
+ for name in cls.serialization_buffers:
106
+ model.register_buffer(name, torch.empty(config["prefix_shapes"][name]))
107
+ return model.prepare_for_serialization()
108
+
109
+ @classmethod
110
+ def from_teacher(
111
+ cls, teacher_path: str | Path, native_tokenizer_path: str | Path,
112
+ speaker_vector_path: str | Path, *, dtype: torch.dtype = torch.float32,
113
+ attn_implementation: str = "eager", device: str | torch.device = "cpu",
114
+ projection_batch_size: int = 1, projected_table_path: str | Path | None = None,
115
+ ) -> "CompactNativeClockTalker":
116
+ source = NativeClockTalker.from_teacher(
117
+ teacher_path, native_tokenizer_path, speaker_vector_path, dtype=dtype,
118
+ attn_implementation=attn_implementation,
119
+ ).eval().to(device)
120
+ return cls.from_native(source, projection_batch_size=projection_batch_size,
121
+ projected_table_path=projected_table_path)
122
+
123
+ @classmethod
124
+ @torch.no_grad()
125
+ def from_native(
126
+ cls, source: NativeClockTalker, *, projection_batch_size: int = 1,
127
+ projected_table_path: str | Path | None = None,
128
+ ) -> "CompactNativeClockTalker":
129
+ """Reuse acoustic weights; leave the source model's modules unmodified.
130
+
131
+ A batch size of one repeats the streaming projection's matrix shape.
132
+ Larger batches use GEMM and may change floating-point rounding.
133
+ """
134
+ model = cls()
135
+ model.hidden_size = source.hidden_size
136
+ for name in ("backbone", "codec_history_embeddings", "q0_head", "codec_eos_head"):
137
+ setattr(model, name, getattr(source, name))
138
+ model.codec_bos = source.codec_bos
139
+ model.native_teacher_ids = list(source.native_teacher_ids)
140
+ model.native_teacher_offsets = list(source.native_teacher_offsets)
141
+ for name, buffer in source.named_buffers(recurse=False):
142
+ model.register_buffer(name, buffer, persistent=name not in source._non_persistent_buffers_set)
143
+
144
+ device = source.text_embedding.weight.device
145
+ teacher_ids = torch.tensor(sorted(set(source.native_teacher_ids)), device=device)
146
+ table_path = Path(projected_table_path) if projected_table_path else None
147
+ if table_path is not None and table_path.exists():
148
+ from safetensors.torch import load_file
149
+ table = load_file(str(table_path), device=str(device))
150
+ if not torch.equal(table["teacher_ids"], teacher_ids):
151
+ raise ValueError("Projected table belongs to a different native vocabulary")
152
+ projected = table["projected"].to(dtype=source.text_embedding.weight.dtype)
153
+ else:
154
+ projected = source.q0_head.weight.new_empty((len(teacher_ids), source.hidden_size))
155
+ for start in range(0, len(teacher_ids), projection_batch_size):
156
+ ids = teacher_ids[start:start + projection_batch_size].reshape(-1, 1)
157
+ projected[start:start + len(ids)] = source.text_projection(source.text_embedding(ids))[:, 0]
158
+ if table_path is not None:
159
+ from safetensors.torch import save_file
160
+ table_path.parent.mkdir(parents=True, exist_ok=True)
161
+ save_file({"teacher_ids": teacher_ids.cpu(), "projected": projected.cpu()}, str(table_path))
162
+ model.text_embedding = ProjectedTextEmbedding(projected, teacher_ids, source.text_embedding.num_embeddings)
163
+ model.text_projection = nn.Identity()
164
+
165
+ # These fifteen input tables were copied from the same teacher tensors.
166
+ model.residual_predictor = copy.deepcopy(source.residual_predictor)
167
+ pairs = [(model.residual_predictor, "q0_embedding", model.codec_history_embeddings[0])]
168
+ pairs.extend((model.residual_predictor.code_embeddings, str(group), model.codec_history_embeddings[group + 1])
169
+ for group in range(len(model.residual_predictor.code_embeddings)))
170
+ for owner, name, history_embedding in pairs:
171
+ if not torch.equal(getattr(owner, name).weight, history_embedding.weight):
172
+ raise ValueError("Codec input embeddings have diverged and cannot be shared")
173
+ setattr(owner, name, history_embedding)
174
+ return model.eval()
compaction.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "source": "/default-filesys/workspace/chenjunzhe/EdgeInstant-1.5b/runs/bilingual_s5_full/final",
3
+ "output": "/default-filesys/workspace/chenjunzhe/EdgeInstant-1.5b/runs/bilingual_s5_full/final_compact",
4
+ "precision": "torch.bfloat16",
5
+ "preserved_modalities": [
6
+ "text",
7
+ "audio_input",
8
+ "image",
9
+ "video",
10
+ "speech_output"
11
+ ],
12
+ "source_parameters": 2031972995,
13
+ "exported_parameters": 2031710851,
14
+ "removed_weights": [
15
+ "codec.quantizer.rvq_first.input_proj.weight",
16
+ "codec.quantizer.rvq_rest.input_proj.weight"
17
+ ],
18
+ "source_weight_bytes": 5677248858,
19
+ "exported_weight_bytes": 4115251826,
20
+ "equivalence_reference": "Source loaded with dtype=auto, as in run_mmau.py; align and audio_special remain FP32.",
21
+ "cache_implementation": "dynamic",
22
+ "recurrent_decode": "cuda_graph"
23
+ }
config.json ADDED
The diff for this file is too large to render. See raw diff
 
configuration_edgeinstant.py ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Configuration for the AudioIn, Thinker, Talker and speech decoder."""
2
+ from __future__ import annotations
3
+
4
+ from transformers import PretrainedConfig
5
+ from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5Config
6
+ from transformers.models.qwen3_omni_moe.configuration_qwen3_omni_moe import (
7
+ Qwen3OmniMoeAudioEncoderConfig,
8
+ )
9
+
10
+
11
+ class EdgeInstantConfig(PretrainedConfig):
12
+ model_type = "edgeinstant"
13
+ sub_configs = {
14
+ "thinker_config": Qwen3_5Config,
15
+ "audio_config": Qwen3OmniMoeAudioEncoderConfig,
16
+ }
17
+ keys_to_ignore_at_inference = ["past_key_values"]
18
+
19
+ def __init__(
20
+ self,
21
+ thinker_config=None,
22
+ audio_config=None,
23
+ projector_config=None,
24
+ talker_config=None,
25
+ codec_config=None,
26
+ audio_tap="proj2",
27
+ audio_pool=1,
28
+ audio_stack=1,
29
+ audio_special_ids=None,
30
+ audio_pad_id=248076,
31
+ sampling_rate=16000,
32
+ output_sampling_rate=24000,
33
+ max_audio_seconds=None,
34
+ qa_prompt="Listen to the user's speech and answer.",
35
+ thinker_special_ids=None,
36
+ think_start_id=248068,
37
+ think_end_id=248069,
38
+ **kwargs,
39
+ ):
40
+ self.thinker_config = (
41
+ thinker_config if isinstance(thinker_config, Qwen3_5Config)
42
+ else Qwen3_5Config(**(thinker_config or {}))
43
+ )
44
+ self.audio_config = (
45
+ audio_config if isinstance(audio_config, Qwen3OmniMoeAudioEncoderConfig)
46
+ else Qwen3OmniMoeAudioEncoderConfig(**(audio_config or {}))
47
+ )
48
+ self.audio_tap = audio_tap
49
+ self.audio_pool = int(audio_pool)
50
+ self.audio_stack = int(audio_stack)
51
+ if audio_tap not in {"proj2", "ln_post"}:
52
+ raise ValueError(f"Unsupported AudioIn tower output: {audio_tap}")
53
+ if self.audio_pool < 1 or self.audio_stack < 1:
54
+ raise ValueError("audio_pool and audio_stack must be positive integers")
55
+ audio_size = self.audio_config.d_model if audio_tap == "ln_post" else self.audio_config.output_dim
56
+ hidden_size = self.thinker_config.text_config.hidden_size
57
+ self.projector_config = {
58
+ "norm_size": audio_size,
59
+ "input_size": audio_size * self.audio_stack,
60
+ "intermediate_size": None,
61
+ "output_size": hidden_size,
62
+ "align_mode": "replace",
63
+ "audio_out_norm": False,
64
+ **(projector_config or {}),
65
+ }
66
+ self.talker_config = dict(talker_config or {})
67
+ self.codec_config = dict(codec_config or {})
68
+ self.audio_special_ids = list(audio_special_ids if audio_special_ids is not None else [248070, 248071, 248076])
69
+ self.audio_pad_id = int(audio_pad_id)
70
+ self.sampling_rate = int(sampling_rate)
71
+ self.output_sampling_rate = int(output_sampling_rate)
72
+ self.max_audio_seconds = max_audio_seconds
73
+ self.qa_prompt = qa_prompt
74
+ self.thinker_special_ids = list(thinker_special_ids or [])
75
+ self.think_start_id = int(think_start_id)
76
+ self.think_end_id = int(think_end_id)
77
+ kwargs.setdefault("tie_word_embeddings", self.thinker_config.tie_word_embeddings)
78
+ kwargs.setdefault("pad_token_id", self.thinker_config.text_config.pad_token_id)
79
+ kwargs.setdefault("bos_token_id", self.thinker_config.text_config.bos_token_id)
80
+ kwargs.setdefault("eos_token_id", self.thinker_config.text_config.eos_token_id)
81
+ kwargs.setdefault("architectures", ["EdgeInstantForConditionalGeneration"])
82
+ kwargs.setdefault("processor_class", "EdgeInstantProcessor")
83
+ kwargs.setdefault("auto_map", {
84
+ "AutoConfig": "configuration_edgeinstant.EdgeInstantConfig",
85
+ "AutoModel": "modeling_edgeinstant.EdgeInstantForConditionalGeneration",
86
+ "AutoModelForCausalLM": "modeling_edgeinstant.EdgeInstantForConditionalGeneration",
87
+ "AutoProcessor": "processing_edgeinstant.EdgeInstantProcessor",
88
+ })
89
+ super().__init__(**kwargs)
90
+
91
+ def get_text_config(self, decoder=None, encoder=None):
92
+ return self.thinker_config.get_text_config(decoder=decoder, encoder=encoder)
93
+
94
+
95
+ EdgeInstantConfig.register_for_auto_class()
generation.py ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import dataclass
4
+ from enum import Enum
5
+
6
+ import torch
7
+
8
+
9
+ class AdvanceStatus(str, Enum):
10
+ WAIT_INPUT = "WAIT_INPUT"
11
+ PRODUCED = "PRODUCED"
12
+ END_AUDIO = "END_AUDIO"
13
+ FORCED_STOP = "FORCED_STOP"
14
+
15
+
16
+ @dataclass
17
+ class GenerationResult:
18
+ frames: list[torch.Tensor]
19
+ status: AdvanceStatus
20
+ actions: list[int]
21
+ consumed_text_tokens: int
22
+ text_end_consumed: bool
generation_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "cache_implementation": "dynamic",
4
+ "eos_token_id": [
5
+ 248046,
6
+ 248046
7
+ ],
8
+ "output_attentions": false,
9
+ "output_hidden_states": false,
10
+ "pad_token_id": 248044,
11
+ "transformers_version": "5.12.1",
12
+ "use_cache": true
13
+ }
image_processor/preprocessor_config.json ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "do_convert_rgb": true,
3
+ "do_normalize": true,
4
+ "do_rescale": true,
5
+ "do_resize": true,
6
+ "image_mean": [
7
+ 0.5,
8
+ 0.5,
9
+ 0.5
10
+ ],
11
+ "image_processor_type": "Qwen2VLImageProcessor",
12
+ "image_std": [
13
+ 0.5,
14
+ 0.5,
15
+ 0.5
16
+ ],
17
+ "merge_size": 2,
18
+ "patch_size": 16,
19
+ "resample": 3,
20
+ "rescale_factor": 0.00392156862745098,
21
+ "size": {
22
+ "longest_edge": 16777216,
23
+ "shortest_edge": 65536
24
+ },
25
+ "temporal_patch_size": 2
26
+ }
image_processor/video_preprocessor_config.json ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "size": {
3
+ "longest_edge": 25165824,
4
+ "shortest_edge": 4096
5
+ },
6
+ "patch_size": 16,
7
+ "temporal_patch_size": 2,
8
+ "merge_size": 2,
9
+ "image_mean": [
10
+ 0.5,
11
+ 0.5,
12
+ 0.5
13
+ ],
14
+ "image_std": [
15
+ 0.5,
16
+ 0.5,
17
+ 0.5
18
+ ],
19
+ "processor_class": "Qwen3VLProcessor",
20
+ "video_processor_type": "Qwen3VLVideoProcessor"
21
+ }
inference_edgeinstant.py ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """CUDA graph replay of Qwen3.5's unchanged recurrent decode operations."""
2
+ from __future__ import annotations
3
+
4
+ import torch
5
+ from transformers.cache_utils import DynamicCache
6
+ from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5DecoderLayer
7
+
8
+
9
+ class EdgeInstantDecoderLayer(Qwen3_5DecoderLayer):
10
+ def __init__(self, config, layer_idx):
11
+ super().__init__(config, layer_idx)
12
+ self.config = config
13
+ self.layer_idx = layer_idx
14
+ self._decode_graph = None
15
+
16
+ def train(self, mode=True):
17
+ self._decode_graph = None
18
+ return super().train(mode)
19
+
20
+ def _apply(self, fn, recurse=True):
21
+ self._decode_graph = None
22
+ return super()._apply(fn, recurse=recurse)
23
+
24
+ def forward(self, hidden_states, position_embeddings, attention_mask=None,
25
+ position_ids=None, past_key_values=None, **kwargs):
26
+ cache_params = past_key_values
27
+ if (self.training or torch.is_grad_enabled() or hidden_states.device.type != "cuda"
28
+ or hidden_states.shape[1] != 1 or cache_params is None
29
+ or attention_mask is not None or torch.compiler.is_compiling()
30
+ or not cache_params.has_previous_state(self.layer_idx)):
31
+ return super().forward(hidden_states, position_embeddings, attention_mask,
32
+ position_ids, past_key_values, **kwargs)
33
+ layer = cache_params.layers[self.layer_idx]
34
+ shape = (hidden_states.shape, hidden_states.dtype, hidden_states.device,
35
+ layer.conv_states.shape, layer.conv_states.dtype,
36
+ layer.recurrent_states.shape, layer.recurrent_states.dtype)
37
+ state = self._decode_graph
38
+ if state is None or state[0] != shape:
39
+ inputs = torch.empty_like(hidden_states)
40
+ graph_cache = DynamicCache(config=self.config)
41
+ graph_cache.update_conv_state(layer.conv_states, self.layer_idx)
42
+ graph_cache.update_recurrent_state(layer.recurrent_states, self.layer_idx)
43
+ inputs.copy_(hidden_states)
44
+ current = torch.cuda.current_stream(hidden_states.device)
45
+ stream = torch.cuda.Stream(device=hidden_states.device)
46
+ stream.wait_stream(current)
47
+ with torch.cuda.stream(stream):
48
+ for _ in range(3):
49
+ super().forward(inputs, None, past_key_values=graph_cache)
50
+ current.wait_stream(stream)
51
+ graph = torch.cuda.CUDAGraph()
52
+ with torch.cuda.graph(graph, stream=stream):
53
+ outputs = super().forward(inputs, None, past_key_values=graph_cache)
54
+ state = (shape, inputs, graph_cache, graph, outputs)
55
+ self._decode_graph = state
56
+ _, inputs, graph_cache, graph, outputs = state
57
+ # Keep each caller's cache independent, including caches returned by generate.
58
+ inputs.copy_(hidden_states)
59
+ graph_cache.update_conv_state(layer.conv_states, self.layer_idx)
60
+ graph_cache.update_recurrent_state(layer.recurrent_states, self.layer_idx)
61
+ graph.replay()
62
+ graph_layer = graph_cache.layers[self.layer_idx]
63
+ cache_params.update_conv_state(graph_layer.conv_states, self.layer_idx)
64
+ cache_params.update_recurrent_state(graph_layer.recurrent_states, self.layer_idx)
65
+ return outputs.clone()
model-00001-of-00003.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a659eaf53119fbb56590f3424e96e09dafeb9350d0f8b4633552db0f6d670b1f
3
+ size 1998027410
model-00002-of-00003.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1ecbfa3ea1abdfddc601ba3f87d041aeed8e78becc56f994cd6954ff560665e4
3
+ size 1996358360
model-00003-of-00003.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0dd9cb8877ae15c36adabc56b5cbf998333182334213a3dc89372e3b918ae025
3
+ size 120866056
model.safetensors.index.json ADDED
The diff for this file is too large to render. See raw diff
 
modeling_edgeinstant.py ADDED
@@ -0,0 +1,392 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Audio-conditioned Thinker and native-token speech generation."""
2
+ from __future__ import annotations
3
+
4
+ from typing import Iterator
5
+
6
+ import torch
7
+ from torch import nn
8
+ from torch.nn.attention import SDPBackend, sdpa_kernel
9
+ from transformers import GenerationMixin, PreTrainedModel, Qwen3_5ForConditionalGeneration
10
+ from transformers.modeling_outputs import CausalLMOutputWithPast
11
+ from transformers.models.qwen3_omni_moe.modeling_qwen3_omni_moe import (
12
+ Qwen3OmniMoeAudioEncoder,
13
+ _get_feat_extract_output_lengths,
14
+ )
15
+
16
+ from .configuration_edgeinstant import EdgeInstantConfig
17
+ from .codec_edgeinstant import EdgeInstantCodecDecoder
18
+ from .inference_edgeinstant import EdgeInstantDecoderLayer
19
+ from .compact_native_clock import CompactNativeClockTalker
20
+
21
+
22
+ class EdgeInstantThinker(Qwen3_5ForConditionalGeneration):
23
+ """Qwen3.5 with CUDA graph replay for cached recurrent decoding."""
24
+
25
+ _no_split_modules = ["Qwen3_5DecoderLayer", "EdgeInstantDecoderLayer", "Qwen3_5VisionBlock"]
26
+
27
+ def __init__(self, config):
28
+ super().__init__(config)
29
+ for index in range(len(self.model.language_model.layers)):
30
+ if config.text_config.layer_types[index] == "linear_attention":
31
+ self.model.language_model.layers[index] = EdgeInstantDecoderLayer(config.text_config, index)
32
+
33
+ class EdgeInstantAudioProjector(nn.Module):
34
+ def __init__(self, config: dict, hidden_size: int):
35
+ super().__init__()
36
+ self.hidden_size = hidden_size
37
+ self.mode = config.get("align_mode", "replace")
38
+ if self.mode == "residual" and int(config["output_size"]) != hidden_size:
39
+ raise ValueError("Residual audio projection requires audio_split=1 (output_size must equal Thinker hidden_size)")
40
+ self.out_norm = bool(config.get("audio_out_norm", False))
41
+ self.norm = nn.LayerNorm(int(config["norm_size"]))
42
+ middle = config.get("intermediate_size")
43
+ self.proj = nn.Linear(int(config["input_size"]), int(middle or config["output_size"]))
44
+ self.proj2 = nn.Linear(int(middle), int(config["output_size"])) if middle else None
45
+ self.gate = nn.Parameter(torch.zeros(()))
46
+ self.out_scale = nn.Parameter(torch.ones(()))
47
+
48
+ def forward(self, audio: torch.Tensor, pad_embedding: torch.Tensor) -> torch.Tensor:
49
+ shape = audio.shape
50
+ frames = audio.float().reshape(*shape[:-1], -1, self.norm.normalized_shape[0])
51
+ mapped = self.proj(self.norm(frames).reshape(shape))
52
+ if self.proj2 is not None:
53
+ mapped = self.proj2(nn.functional.gelu(mapped))
54
+ if self.mode == "residual":
55
+ return pad_embedding.to(mapped.dtype) + self.gate.clamp(0, 4) * mapped
56
+ if self.out_norm:
57
+ mapped = nn.functional.normalize(mapped, dim=-1, eps=1e-6) * self.hidden_size ** 0.5
58
+ return (mapped * self.out_scale).reshape(-1, self.hidden_size)
59
+
60
+
61
+ class EdgeInstantForConditionalGeneration(PreTrainedModel, GenerationMixin):
62
+ config_class = EdgeInstantConfig
63
+ base_model_prefix = "edgeinstant"
64
+ main_input_name = "input_ids"
65
+ _supports_sdpa = True
66
+ _keep_in_fp32_modules_strict = ["align", "audio_special"]
67
+ _no_split_modules = ["Qwen3OmniMoeAudioEncoder", "EdgeInstantAudioProjector", "EdgeInstantCodecDecoder"]
68
+
69
+ def __init__(self, config, *, thinker=None, audio_tower=None, align=None,
70
+ talker=None, codec=None, audio_special_delta=None):
71
+ super().__init__(config)
72
+ supplied = [component for component in (thinker, audio_tower, align, talker, codec) if component is not None]
73
+ self.thinker = thinker if thinker is not None else EdgeInstantThinker(config.thinker_config)
74
+ self.audio_tower = audio_tower if audio_tower is not None else Qwen3OmniMoeAudioEncoder(config.audio_config)
75
+ if config.audio_tap == "ln_post":
76
+ self.audio_tower.proj1 = nn.Identity()
77
+ self.audio_tower.act = nn.Identity()
78
+ self.audio_tower.proj2 = nn.Identity()
79
+ self.align = align if align is not None else EdgeInstantAudioProjector(
80
+ config.projector_config, config.thinker_config.text_config.hidden_size,
81
+ )
82
+ self.audio_special = nn.Embedding(len(config.audio_special_ids), config.thinker_config.text_config.hidden_size)
83
+ if audio_special_delta is not None:
84
+ self.audio_special.weight = nn.Parameter(audio_special_delta.detach().clone())
85
+ supplied.append(self.audio_special)
86
+ self.talker = talker if talker is not None else CompactNativeClockTalker.from_config(config.talker_config)
87
+ self.talker.prepare_for_serialization()
88
+ self.codec = codec if codec is not None else (
89
+ EdgeInstantCodecDecoder(config.codec_config) if config.codec_config else None
90
+ )
91
+ self._tied_weights_keys = {
92
+ f"talker.{target}": f"talker.{source}"
93
+ for target, source in self.talker.tied_weight_keys.items()
94
+ }
95
+ for component in supplied:
96
+ for module in component.modules():
97
+ module._is_hf_initialized = True
98
+ self.post_init()
99
+
100
+ def get_input_embeddings(self):
101
+ return self.thinker.get_input_embeddings()
102
+
103
+ def set_input_embeddings(self, embeddings):
104
+ self.thinker.set_input_embeddings(embeddings)
105
+
106
+ def get_output_embeddings(self):
107
+ return self.thinker.get_output_embeddings()
108
+
109
+ def _base_thinker(self):
110
+ return self.thinker.get_base_model() if hasattr(self.thinker, "get_base_model") else self.thinker
111
+
112
+ def encode_audio(self, input_features, feature_attention_mask=None, tower=None, feature_lengths=None):
113
+ """Return packed audio encoder features before the trainable projector."""
114
+ from transformers.models.qwen3_omni_moe.modeling_qwen3_omni_moe import (
115
+ chunk_and_pad_features, get_audio_cu_seqlens, get_valid_indices,
116
+ )
117
+ from transformers.utils.generic import is_flash_attention_requested
118
+
119
+ tower = self.audio_tower if tower is None else tower
120
+ tower_weight = next(tower.parameters())
121
+ if feature_lengths is not None:
122
+ lengths = torch.as_tensor(feature_lengths, dtype=torch.long).cpu()
123
+ elif feature_attention_mask is not None:
124
+ lengths = feature_attention_mask.sum(-1).to(device="cpu", dtype=torch.long)
125
+ else:
126
+ lengths = torch.full((input_features.shape[0],), input_features.shape[-1], dtype=torch.long)
127
+ if bool((lengths <= 0).any()):
128
+ raise ValueError("Audio must contain at least one valid mel frame")
129
+ packed = torch.cat([row[:, :int(length)] for row, length in zip(input_features, lengths)], dim=-1)
130
+ packed = packed.to(device=tower_weight.device, dtype=tower_weight.dtype)
131
+ padded_feature, chunk_lengths = chunk_and_pad_features(packed, lengths, tower.n_window)
132
+ valid_indices = get_valid_indices(chunk_lengths).to(tower_weight.device)
133
+ cu_seqlens = get_audio_cu_seqlens(chunk_lengths, lengths, tower.n_window_infer, tower.n_window)
134
+ # SDPA/eager split each attention window in Python; keep their boundaries on the CPU.
135
+ if is_flash_attention_requested(tower.config):
136
+ cu_seqlens = cu_seqlens.to(tower_weight.device)
137
+ hidden = tower(
138
+ input_features=packed, feature_lens=lengths,
139
+ padded_feature=padded_feature, chunk_lengths=chunk_lengths,
140
+ valid_indices=valid_indices, cu_seqlens=cu_seqlens,
141
+ ).last_hidden_state
142
+ return self.group_audio_features(hidden, lengths)
143
+
144
+ def group_audio_features(self, hidden, frame_lengths):
145
+ """Apply per-recording frame stacking or pooling to packed encoder output."""
146
+ if self.config.audio_stack == 1 and self.config.audio_pool == 1:
147
+ return hidden
148
+ counts = _get_feat_extract_output_lengths(frame_lengths).tolist()
149
+ processed = []
150
+ for row in hidden.split(counts):
151
+ if self.config.audio_stack > 1:
152
+ padding = -len(row) % self.config.audio_stack
153
+ if padding:
154
+ row = torch.cat((row, row[-1:].expand(padding, -1)))
155
+ row = row.reshape(len(row) // self.config.audio_stack, -1)
156
+ elif self.config.audio_pool > 1:
157
+ row = torch.stack([chunk.mean(0) for chunk in row.split(self.config.audio_pool)])
158
+ processed.append(row)
159
+ return torch.cat(processed)
160
+
161
+ def pad_embedding(self):
162
+ """Audio placeholder embedding including its learned special-token delta."""
163
+ embed = self.get_input_embeddings()
164
+ pad = embed.weight[self.config.audio_pad_id]
165
+ if self.config.audio_pad_id in self.config.audio_special_ids:
166
+ slot = self.config.audio_special_ids.index(self.config.audio_pad_id)
167
+ pad = pad + self.audio_special.weight[slot].to(pad)
168
+ return pad
169
+
170
+ def get_audio_features(self, input_features, feature_attention_mask, feature_lengths=None):
171
+ features = self.encode_audio(input_features, feature_attention_mask, feature_lengths=feature_lengths)
172
+ device = self.align.proj.weight.device
173
+ return self.align(features.to(device), self.pad_embedding().to(device))
174
+
175
+ def merge_audio_embeds(self, input_ids, input_features, feature_attention_mask, feature_lengths=None):
176
+ vectors = self.get_audio_features(input_features, feature_attention_mask, feature_lengths)
177
+ return self.merge_custom_audio(input_ids, vectors)
178
+
179
+ def merge_custom_audio(self, input_ids, audio_vectors):
180
+ """Insert projected or diagnostic audio vectors in the model's audio slots."""
181
+ embedding = self.get_input_embeddings()
182
+ ids = input_ids.to(embedding.weight.device)
183
+ embeds = embedding(ids)
184
+ for slot, token_id in enumerate(self.config.audio_special_ids):
185
+ embeds = embeds + (ids == token_id).unsqueeze(-1) * self.audio_special.weight[slot].to(embeds)
186
+ vectors = audio_vectors.to(embeds)
187
+ mask = ids == self.config.audio_pad_id
188
+ if int(mask.sum()) != len(vectors):
189
+ raise ValueError(f"Audio slots {int(mask.sum())} do not match encoded vectors {len(vectors)}")
190
+ return embeds.masked_scatter(mask.unsqueeze(-1).expand_as(embeds), vectors)
191
+
192
+ def forward(self, input_ids=None, attention_mask=None, input_features=None,
193
+ feature_attention_mask=None, inputs_embeds=None, labels=None,
194
+ audio_feature_lengths=None, answer_loss_reduction=None,
195
+ eos_loss_weight=1.0, unlikelihood_mask=None, unlikelihood_weight=1.0, **kwargs):
196
+ if input_features is not None:
197
+ if input_ids is None or feature_attention_mask is None:
198
+ raise ValueError("Audio forward requires input_ids and feature_attention_mask")
199
+ cache = kwargs.get("past_key_values")
200
+ if cache is None or cache.get_seq_length() == 0:
201
+ self._base_thinker().model.rope_deltas = None
202
+ inputs_embeds = self.merge_audio_embeds(
203
+ input_ids, input_features, feature_attention_mask, audio_feature_lengths,
204
+ )
205
+ if kwargs.get("position_ids") is None and (
206
+ kwargs.get("image_grid_thw") is not None or kwargs.get("video_grid_thw") is not None
207
+ ):
208
+ kwargs["position_ids"] = self._base_thinker().model.compute_3d_position_ids(
209
+ input_ids, inputs_embeds, attention_mask=attention_mask,
210
+ **{key: kwargs.get(key) for key in ("image_grid_thw", "video_grid_thw", "mm_token_type_ids")},
211
+ )
212
+ input_ids = None
213
+ if answer_loss_reduction is not None:
214
+ if labels is None or answer_loss_reduction not in {"token_sum", "example_sum"}:
215
+ raise ValueError("Answer loss requires labels and token_sum or example_sum reduction")
216
+ # LoRA lives on the base model's decoder linears; the vocabulary projection
217
+ # only needs hidden states that predict a supervised answer token or EOS.
218
+ thinker = self._base_thinker()
219
+ outputs = thinker.model(input_ids=input_ids, inputs_embeds=inputs_embeds,
220
+ attention_mask=attention_mask, **kwargs)
221
+ targets = labels[:, 1:]
222
+ mask = targets.ne(-100)
223
+ hidden = outputs.last_hidden_state[:, :-1][mask]
224
+ logits = thinker.lm_head(hidden)
225
+ losses = nn.functional.cross_entropy(logits.float(), targets[mask], reduction="none")
226
+ if unlikelihood_mask is not None:
227
+ negative = unlikelihood_mask[:, 1:][mask].bool()
228
+ if negative.any():
229
+ negative_logits = logits[negative].float()
230
+ negative_ids = targets[mask][negative, None]
231
+ selected = negative_logits.gather(1, negative_ids).squeeze(1)
232
+ alternatives = negative_logits.scatter(1, negative_ids, -torch.inf).logsumexp(dim=1)
233
+ losses = losses.index_put((negative,),
234
+ nn.functional.softplus(selected - alternatives) * unlikelihood_weight)
235
+ eos_ids = self.config.eos_token_id
236
+ if isinstance(eos_ids, int):
237
+ eos_ids = [eos_ids]
238
+ is_eos = torch.isin(targets[mask], targets.new_tensor(eos_ids))
239
+ weights = torch.where(is_eos, eos_loss_weight, 1.0)
240
+ losses = losses * weights
241
+ if answer_loss_reduction == "example_sum":
242
+ example_ids = mask.nonzero(as_tuple=True)[0]
243
+ counts = weights.new_zeros(targets.shape[0]).scatter_add(0, example_ids, weights)
244
+ losses = losses / counts[example_ids]
245
+ return CausalLMOutputWithPast(loss=losses.sum())
246
+ return self.thinker(input_ids=input_ids, inputs_embeds=inputs_embeds,
247
+ attention_mask=attention_mask, labels=labels, **kwargs)
248
+
249
+ @torch.no_grad()
250
+ @sdpa_kernel([SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH])
251
+ def generate(self, input_ids=None, attention_mask=None, input_features=None,
252
+ feature_attention_mask=None, **kwargs):
253
+ """Generate text token IDs, including the supplied prompt as in HF causal LMs."""
254
+ self._base_thinker().model.rope_deltas = None
255
+ if input_features is not None:
256
+ kwargs["inputs_embeds"] = self.merge_audio_embeds(input_ids, input_features, feature_attention_mask)
257
+ self.thinker.generation_config = self.generation_config
258
+ kwargs.setdefault("eos_token_id", self.config.eos_token_id)
259
+ return self.thinker.generate(input_ids=input_ids, attention_mask=attention_mask, **kwargs)
260
+
261
+ generate_text = generate
262
+
263
+ @torch.inference_mode()
264
+ def synthesize(self, token_ids, *, language="chinese", max_frames=1536, do_sample=False,
265
+ decode=True, **sampling):
266
+ """Speak one sequence of native Thinker token IDs."""
267
+ state = self.talker.new_state(language=language, max_frames=max_frames, do_sample=do_sample, **sampling)
268
+ for token in torch.as_tensor(token_ids).reshape(-1).tolist():
269
+ state.push_token(token)
270
+ state.end_input()
271
+ while not state.ended:
272
+ state.advance(max_new_frames=4)
273
+ codes = torch.stack(state.frames) if state.frames else torch.empty((0, 16), dtype=torch.long)
274
+ waveform = None
275
+ if decode:
276
+ if self.codec is None:
277
+ raise ValueError("The model package does not contain a codec decoder")
278
+ waveform = self.codec.decode(codes.unsqueeze(0))[0] if len(codes) else torch.empty(0)
279
+ return {"audio_codes": codes, "audio": waveform, "sampling_rate": self.config.output_sampling_rate,
280
+ "stop_reason": state.status.value}
281
+
282
+ @torch.inference_mode()
283
+ def generate_speech(self, input_ids=None, *, language="chinese", max_frames=1536,
284
+ audio_do_sample=False, **kwargs):
285
+ """Generate one reply and its waveform; text generation accepts standard HF options."""
286
+ if input_ids is None or input_ids.shape[0] != 1:
287
+ raise ValueError("Speech generation accepts one conversation at a time")
288
+ sequences = self.generate(input_ids=input_ids, **kwargs)
289
+ if not isinstance(sequences, torch.Tensor):
290
+ sequences = sequences.sequences
291
+ reply = sequences[0, input_ids.shape[1]:]
292
+ special = set(self.config.thinker_special_ids)
293
+ public, thinking = [], False
294
+ for token in reply.tolist():
295
+ if token == self.config.think_start_id:
296
+ thinking = True
297
+ elif token == self.config.think_end_id:
298
+ thinking = False
299
+ elif not thinking and token not in special:
300
+ public.append(token)
301
+ result = self.synthesize(public, language=language, max_frames=max_frames, do_sample=audio_do_sample)
302
+ return {"sequences": sequences, "text_token_ids": torch.tensor(public), **result}
303
+
304
+ @torch.inference_mode()
305
+ def stream_generate(self, input_ids, attention_mask=None, input_features=None,
306
+ feature_attention_mask=None, *, language="chinese", max_new_tokens=512,
307
+ max_frames=1536, first_packet_frames=1, chunk_frames=4,
308
+ left_context_frames=50, audio_do_sample=False, cancelled=None) -> Iterator[dict]:
309
+ """Yield native text tokens and continuous 24 kHz audio chunks for one reply."""
310
+ if input_ids.shape[0] != 1:
311
+ raise ValueError("Streaming accepts one conversation at a time")
312
+ state = self.talker.new_state(language=language, max_frames=max_frames, do_sample=audio_do_sample)
313
+ self._base_thinker().model.rope_deltas = None
314
+ inputs_embeds = None
315
+ if input_features is not None:
316
+ inputs_embeds = self.merge_audio_embeds(input_ids, input_features, feature_attention_mask)
317
+ if attention_mask is None:
318
+ attention_mask = torch.ones_like(input_ids)
319
+ sequences = input_ids
320
+ sent = 0
321
+ special = set(self.config.thinker_special_ids)
322
+ eos = self.config.eos_token_id
323
+ eos = {eos} if isinstance(eos, int) else set(eos)
324
+ thinking = False
325
+
326
+ def packet():
327
+ nonlocal sent
328
+ if len(state.frames) == sent:
329
+ return None
330
+ start = max(0, sent - left_context_frames)
331
+ codes = torch.stack(state.frames[start:]).unsqueeze(0)
332
+ pcm = self.codec(codes)[0]
333
+ pcm = pcm[(sent - start) * self.codec.samples_per_frame:]
334
+ event = {"type": "audio", "audio": pcm.detach().cpu(), "sampling_rate": self.codec.sample_rate,
335
+ "sample_start": sent * self.codec.samples_per_frame,
336
+ "audio_codes": torch.stack(state.frames[sent:])}
337
+ sent = len(state.frames)
338
+ return event
339
+
340
+ if self.codec is None:
341
+ raise ValueError("The model package does not contain a codec decoder")
342
+ try:
343
+ inputs = self.thinker.prepare_inputs_for_generation(
344
+ sequences, inputs_embeds=inputs_embeds, attention_mask=attention_mask,
345
+ is_first_iteration=True, use_cache=True, logits_to_keep=1,
346
+ )
347
+ output = self.thinker(**inputs)
348
+ cache = output.past_key_values
349
+ reached_eos = False
350
+ for _ in range(max_new_tokens):
351
+ if cancelled is not None and cancelled.is_set():
352
+ return
353
+ token = output.logits[:, -1].argmax(-1)
354
+ token_id = int(token.item())
355
+ if token_id in eos:
356
+ reached_eos = True
357
+ break
358
+ if token_id == self.config.think_start_id:
359
+ thinking = True
360
+ elif token_id == self.config.think_end_id:
361
+ thinking = False
362
+ elif not thinking and token_id not in special:
363
+ yield {"type": "text", "token_id": token_id}
364
+ state.push_token(token_id)
365
+ state.advance(max_new_frames=first_packet_frames if sent == 0 else chunk_frames)
366
+ event = packet()
367
+ if event is not None:
368
+ yield event
369
+ if state.ended:
370
+ break
371
+ sequences = torch.cat((sequences, token[:, None]), dim=-1)
372
+ attention_mask = torch.cat((attention_mask, torch.ones_like(token[:, None])), dim=-1)
373
+ inputs = self.thinker.prepare_inputs_for_generation(
374
+ sequences, past_key_values=cache, attention_mask=attention_mask,
375
+ is_first_iteration=False, next_sequence_length=1, use_cache=True, logits_to_keep=1,
376
+ )
377
+ output = self.thinker(**inputs)
378
+ cache = output.past_key_values
379
+ if not state.ended:
380
+ if not reached_eos:
381
+ raise RuntimeError(f"Thinker did not reach EOS within {max_new_tokens} tokens")
382
+ state.end_input()
383
+ while not state.ended:
384
+ if cancelled is not None and cancelled.is_set():
385
+ return
386
+ state.advance(max_new_frames=chunk_frames)
387
+ event = packet()
388
+ if event is not None:
389
+ yield event
390
+ yield {"type": "done", "stop_reason": state.status.value, "samples": sent * self.codec.samples_per_frame}
391
+ finally:
392
+ state.cancel()
native_clock_talker.py ADDED
@@ -0,0 +1,257 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from collections import deque
4
+ from dataclasses import dataclass, field
5
+ import json
6
+ from pathlib import Path
7
+
8
+ import torch
9
+ from safetensors import safe_open
10
+ from torch import nn
11
+
12
+ from .teacher import TeacherAcousticModel
13
+ from .sequence import Action
14
+ from .sampling import sample_token
15
+ from .generation import AdvanceStatus, GenerationResult
16
+
17
+
18
+ class NativeClockTalker(nn.Module):
19
+ """Teacher-initialized acoustic model with an incremental native text clock."""
20
+
21
+ codebook_size = 2048
22
+ hidden_size = 1024
23
+ eos_class = 2048
24
+ supported_languages = ("chinese", "english", "auto")
25
+
26
+ @classmethod
27
+ @torch.no_grad()
28
+ def from_teacher(
29
+ cls,
30
+ teacher_path: str | Path,
31
+ native_tokenizer_path: str | Path,
32
+ speaker_vector_path: str | Path,
33
+ *,
34
+ dtype: torch.dtype = torch.float32,
35
+ attn_implementation: str = "eager",
36
+ ) -> "NativeClockTalker":
37
+ teacher_path = Path(teacher_path)
38
+ source = TeacherAcousticModel.from_teacher(
39
+ teacher_path, native_tokenizer_path, speaker_vector_path,
40
+ dtype=dtype, attn_implementation=attn_implementation,
41
+ )
42
+ model = cls()
43
+ for name in ("backbone", "text_embedding", "text_projection", "codec_history_embeddings",
44
+ "residual_predictor", "q0_head"):
45
+ setattr(model, name, getattr(source, name))
46
+ model.codec_bos = source.codec_bos
47
+ model.native_teacher_ids = source.native_teacher_ids.tolist()
48
+ model.native_teacher_offsets = source.native_teacher_offsets.tolist()
49
+ model.codec_eos_head = nn.Linear(model.hidden_size, 1, bias=False, dtype=dtype)
50
+ config = json.loads((teacher_path / "config.json").read_text())
51
+ talker = config["talker_config"]
52
+ vocabulary = json.loads((teacher_path / "vocab.json").read_text())
53
+
54
+ def text(ids):
55
+ return model.text_projection(model.text_embedding(torch.tensor(ids)))
56
+
57
+ tts_bos, tts_eos, tts_pad = text([config["tts_bos_token_id"], config["tts_eos_token_id"],
58
+ config["tts_pad_token_id"]])
59
+ role = text([config["im_start_token_id"], config["assistant_token_id"], vocabulary["Ċ"]])
60
+ with safe_open(teacher_path / "model.safetensors", framework="pt", device="cpu") as weights:
61
+ codec = weights.get_tensor("talker.model.codec_embedding.weight").to(dtype=dtype)
62
+ language = [talker["codec_think_id"], talker["codec_think_bos_id"],
63
+ talker["codec_language_id"]["chinese"], talker["codec_think_eos_id"]]
64
+ prefix = torch.cat((role, codec[language] + tts_pad,
65
+ (source.speaker_vector + tts_pad).unsqueeze(0),
66
+ (codec[talker["codec_pad_id"]] + tts_bos).unsqueeze(0)), dim=0)
67
+ english = prefix.clone()
68
+ english[5] = codec[talker["codec_language_id"]["english"]] + tts_pad
69
+ auto = torch.cat((role, codec[[talker["codec_nothink_id"],
70
+ talker["codec_think_bos_id"], talker["codec_think_eos_id"]]] + tts_pad,
71
+ prefix[7:]), dim=0)
72
+ model.codec_eos_head.weight.copy_(
73
+ weights.get_tensor("talker.codec_head.weight")[talker["codec_eos_token_id"]].unsqueeze(0)
74
+ )
75
+ model.register_buffer("prefix", prefix.unsqueeze(0), persistent=False)
76
+ model.register_buffer("prefix_english", english.unsqueeze(0), persistent=False)
77
+ model.register_buffer("prefix_auto", auto.unsqueeze(0), persistent=False)
78
+ model.register_buffer("tts_eos", tts_eos.view(1, 1, -1), persistent=False)
79
+ model.register_buffer("tts_pad", tts_pad.view(1, 1, -1), persistent=False)
80
+ return model
81
+
82
+ def new_state(self, **kwargs) -> "NativeClockState":
83
+ return NativeClockState(model=self, **kwargs)
84
+
85
+ def prefix_for_language(self, language: str) -> torch.Tensor:
86
+ if language == "chinese":
87
+ return self.prefix
88
+ if language == "english":
89
+ return self.prefix_english
90
+ if language == "auto":
91
+ return self.prefix_auto
92
+ raise ValueError(f"Unsupported speech language: {language}")
93
+
94
+ def mapped_teacher_ids(self, native_id: int) -> list[int]:
95
+ if not 0 <= native_id < len(self.native_teacher_offsets) - 1:
96
+ return []
97
+ start, end = self.native_teacher_offsets[native_id:native_id + 2]
98
+ return self.native_teacher_ids[start:end]
99
+
100
+ def history_features(self, codes: torch.Tensor) -> torch.Tensor:
101
+ return sum(embedding(codes[..., group])
102
+ for group, embedding in enumerate(self.codec_history_embeddings))
103
+
104
+
105
+ @dataclass
106
+ class NativeClockState:
107
+ model: NativeClockTalker
108
+ language: str = "chinese"
109
+ max_frames: int = 1536
110
+ max_events: int = 4096
111
+ do_sample: bool = False
112
+ top_k: int = 50
113
+ top_p: float = 1.0
114
+ temperature: float = 0.9
115
+ repetition_penalty: float = 1.05
116
+ sampling_generator: torch.Generator | None = None
117
+ pending: deque = field(default_factory=deque)
118
+ native_token_ids: list[int] = field(default_factory=list)
119
+ native_remaining: list[int] = field(default_factory=list)
120
+ teacher_token_ids: list[int] = field(default_factory=list)
121
+ teacher_token_owners: list[int] = field(default_factory=list)
122
+ consumed_teacher_tokens: int = 0
123
+ consumed_tokens: int = 0
124
+ frames: list[torch.Tensor] = field(default_factory=list)
125
+ actions: list[int] = field(default_factory=list)
126
+ frame_conditioned_text_tokens: list[int] = field(default_factory=list)
127
+ frame_arrived_text_tokens: list[int] = field(default_factory=list)
128
+ frame_available_teacher_tokens: list[int] = field(default_factory=list)
129
+ frame_consumed_teacher_tokens: list[int] = field(default_factory=list)
130
+ cache: object | None = None
131
+ previous_codes: torch.Tensor | None = None
132
+ input_ended: bool = False
133
+ text_end_consumed: bool = False
134
+ text_end_at_codec_frame: int | None = None
135
+ ended: bool = False
136
+ cancelled: bool = False
137
+ event_count: int = 0
138
+ suppressed_early_eos: int = 0
139
+ q0_eos_generated: bool = False
140
+ status: AdvanceStatus = AdvanceStatus.WAIT_INPUT
141
+
142
+ @property
143
+ def consumed_native_tokens(self) -> int:
144
+ return self.consumed_tokens
145
+
146
+ def _update_consumed_native(self) -> None:
147
+ while self.consumed_tokens < len(self.native_remaining) and self.native_remaining[self.consumed_tokens] == 0:
148
+ self.consumed_tokens += 1
149
+
150
+ def push_token(self, native_id: int) -> None:
151
+ if self.input_ended or self.ended:
152
+ raise RuntimeError("Cannot append native text after the input has ended")
153
+ pieces = self.model.mapped_teacher_ids(int(native_id))
154
+ index = len(self.native_token_ids)
155
+ self.native_token_ids.append(int(native_id))
156
+ self.native_remaining.append(len(pieces))
157
+ self.teacher_token_ids.extend(pieces)
158
+ self.teacher_token_owners.extend([index] * len(pieces))
159
+ self.pending.extend((piece, index) for piece in pieces)
160
+ self._update_consumed_native()
161
+
162
+ def end_input(self) -> None:
163
+ self.input_ended = True
164
+
165
+ finish = end_input
166
+
167
+ def cancel(self) -> None:
168
+ self.cancelled = True
169
+ self.ended = True
170
+ self.status = AdvanceStatus.FORCED_STOP
171
+ self.cache = None
172
+ self.previous_codes = None
173
+ self.pending.clear()
174
+
175
+ def _result(self, status: AdvanceStatus) -> GenerationResult:
176
+ self.status = status
177
+ return GenerationResult(self.frames, status, self.actions, self.consumed_tokens, self.text_end_consumed)
178
+
179
+ @torch.inference_mode()
180
+ def advance(self, max_new_frames: int = 4) -> GenerationResult:
181
+ if self.ended:
182
+ return self._result(self.status)
183
+ produced = 0
184
+ while produced < max_new_frames:
185
+ if len(self.frames) >= self.max_frames or self.event_count >= self.max_events:
186
+ self.ended = True
187
+ return self._result(AdvanceStatus.FORCED_STOP)
188
+ if self.pending:
189
+ teacher_id, native_index = self.pending.popleft()
190
+ ids = torch.tensor([[teacher_id]], device=self.model.q0_head.weight.device)
191
+ text = self.model.text_projection(self.model.text_embedding(ids))
192
+ self.native_remaining[native_index] -= 1
193
+ self._update_consumed_native()
194
+ self.consumed_teacher_tokens += 1
195
+ elif not self.input_ended:
196
+ return self._result(AdvanceStatus.WAIT_INPUT)
197
+ elif not self.text_end_consumed:
198
+ text = self.model.tts_eos
199
+ self.text_end_consumed = True
200
+ self.text_end_at_codec_frame = len(self.frames)
201
+ else:
202
+ text = self.model.tts_pad
203
+ if self.cache is None:
204
+ inputs = torch.cat((self.model.prefix_for_language(self.language),
205
+ self.model.codec_bos.view(1, 1, -1) + text), dim=1)
206
+ else:
207
+ inputs = self.model.history_features(self.previous_codes).unsqueeze(1) + text
208
+ output = self.model.backbone(inputs_embeds=inputs, past_key_values=self.cache, use_cache=True)
209
+ self.cache = output.past_key_values
210
+ hidden = output.last_hidden_state[:, -1]
211
+ logits = torch.cat((self.model.q0_head(hidden), self.model.codec_eos_head(hidden)), dim=-1).float()
212
+ if self.frames and self.repetition_penalty != 1.0:
213
+ seen = torch.tensor(list({int(frame[0]) for frame in self.frames}), device=logits.device)
214
+ scores = logits[:, seen]
215
+ logits[:, seen] = torch.where(scores < 0, scores * self.repetition_penalty,
216
+ scores / self.repetition_penalty)
217
+ eos_allowed = self.input_ended and not self.pending and self.text_end_consumed
218
+ if not eos_allowed:
219
+ self.suppressed_early_eos += int(logits.argmax(-1).item() == self.model.eos_class)
220
+ logits[:, self.model.eos_class] = -torch.inf
221
+ q0 = sample_token(logits, do_sample=self.do_sample, top_k=self.top_k, top_p=self.top_p,
222
+ temperature=self.temperature, generator=self.sampling_generator)
223
+ self.event_count += 1
224
+ if int(q0.item()) == self.model.eos_class:
225
+ self.ended = True
226
+ self.q0_eos_generated = True
227
+ self.actions.append(int(Action.END))
228
+ return self._result(AdvanceStatus.END_AUDIO)
229
+ codes = self.model.residual_predictor.generate(hidden, q0, do_sample=self.do_sample,
230
+ top_k=self.top_k, top_p=self.top_p, temperature=self.temperature, generator=self.sampling_generator)
231
+ self.previous_codes = codes
232
+ self.frames.append(codes[0].detach().cpu())
233
+ self.actions.append(int(Action.EMIT))
234
+ self.frame_conditioned_text_tokens.append(self.consumed_tokens)
235
+ self.frame_arrived_text_tokens.append(len(self.native_token_ids))
236
+ self.frame_available_teacher_tokens.append(len(self.teacher_token_ids))
237
+ self.frame_consumed_teacher_tokens.append(self.consumed_teacher_tokens)
238
+ produced += 1
239
+ return self._result(AdvanceStatus.PRODUCED)
240
+
241
+ def speech_unit_report(self) -> dict:
242
+ return dict(native_token_ids=self.native_token_ids, native_token_count=len(self.native_token_ids),
243
+ teacher_token_ids=self.teacher_token_ids, teacher_token_count=len(self.teacher_token_ids),
244
+ teacher_token_owners=self.teacher_token_owners,
245
+ consumed_native_tokens=self.consumed_tokens, consumed_teacher_tokens=self.consumed_teacher_tokens)
246
+
247
+ def report(self) -> dict:
248
+ return dict(**self.speech_unit_report(), model_variant="teacher_initialized_native_clock_talker_v1",
249
+ language=self.language,
250
+ terminal_status=self.status.value, q0_eos_generated=self.q0_eos_generated,
251
+ q0_eos_teacher_id=2150, suppressed_early_eos=self.suppressed_early_eos,
252
+ text_end_consumed=self.text_end_consumed, text_end_at_codec_frame=self.text_end_at_codec_frame,
253
+ frame_arrived_text_tokens=self.frame_arrived_text_tokens,
254
+ frame_conditioned_text_tokens=self.frame_conditioned_text_tokens,
255
+ frame_available_teacher_tokens=self.frame_available_teacher_tokens,
256
+ frame_consumed_teacher_tokens=self.frame_consumed_teacher_tokens,
257
+ generated_codec_frames=len(self.frames), cancelled=self.cancelled)
preprocessor_config.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "chunk_length": 30,
3
+ "dither": 0.0,
4
+ "feature_extractor_type": "WhisperFeatureExtractor",
5
+ "feature_size": 128,
6
+ "hop_length": 160,
7
+ "n_fft": 400,
8
+ "n_samples": 480000,
9
+ "nb_max_frames": 3000,
10
+ "padding_side": "right",
11
+ "padding_value": 0.0,
12
+ "return_attention_mask": true,
13
+ "sampling_rate": 16000
14
+ }
processing_edgeinstant.py ADDED
@@ -0,0 +1,228 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Text and 16 kHz speech inputs for EdgeInstant."""
2
+ from __future__ import annotations
3
+
4
+ import json
5
+ from pathlib import Path
6
+
7
+ import numpy as np
8
+ from transformers import AutoImageProcessor, AutoTokenizer, ProcessorMixin, WhisperFeatureExtractor
9
+ from transformers.dynamic_module_utils import custom_object_save
10
+ from transformers.feature_extraction_utils import BatchFeature
11
+
12
+ from .configuration_edgeinstant import EdgeInstantConfig
13
+
14
+
15
+ class EdgeInstantProcessor(ProcessorMixin):
16
+ tokenizer_class = "AutoTokenizer"
17
+ feature_extractor_class = "WhisperFeatureExtractor"
18
+ model_input_names = ["input_ids", "attention_mask", "input_features", "feature_attention_mask"]
19
+
20
+ def __init__(self, tokenizer, feature_extractor, config, image_processor=None, video_processor_config=None):
21
+ self.tokenizer = tokenizer
22
+ self.feature_extractor = feature_extractor
23
+ self.config = config
24
+ self.image_processor = image_processor
25
+ self.video_processor_config = video_processor_config
26
+ self.chat_template = getattr(tokenizer, "chat_template", None)
27
+
28
+ def audio_token_count(self, frame_count):
29
+ """Count encoder frames, temporal grouping and projector output tokens."""
30
+ frame_count = int(frame_count)
31
+ # Qwen3-ASR uses three stride-2 convolutions per 100-frame block.
32
+ count = (frame_count // 100) * 13 + ((frame_count % 100) + 7) // 8
33
+ stride = self.config.audio_stack if self.config.audio_stack > 1 else self.config.audio_pool
34
+ count = (count + stride - 1) // stride
35
+ projector = self.config.projector_config
36
+ if projector["align_mode"] != "residual":
37
+ hidden_size = self.config.thinker_config.text_config.hidden_size
38
+ count = count * projector["output_size"] // hidden_size
39
+ return count
40
+
41
+ def _audio_batch(self, audio, sampling_rate):
42
+ if sampling_rate != self.config.sampling_rate:
43
+ raise ValueError(f"Audio must use {self.config.sampling_rate} Hz; received {sampling_rate} Hz")
44
+ if hasattr(audio, "detach"):
45
+ audio = audio.detach().cpu().numpy()
46
+ if isinstance(audio, np.ndarray):
47
+ batch = [audio] if audio.ndim == 1 else list(audio)
48
+ elif isinstance(audio, (list, tuple)):
49
+ batch = [audio] if audio and np.isscalar(audio[0]) else list(audio)
50
+ else:
51
+ raise TypeError("audio must be a mono waveform or a batch of mono waveforms")
52
+ if not batch:
53
+ raise ValueError("Audio must be a nonempty mono waveform")
54
+ result = []
55
+ for waveform in batch:
56
+ if hasattr(waveform, "detach"):
57
+ waveform = waveform.detach().cpu().numpy()
58
+ waveform = np.asarray(waveform, dtype=np.float32)
59
+ if waveform.ndim != 1 or not waveform.size:
60
+ raise ValueError("Audio must be a nonempty mono waveform")
61
+ result.append(waveform)
62
+ return result
63
+
64
+ def audio_features(self, audio, sampling_rate=16000, return_tensors="np"):
65
+ """Extract complete waveforms with zero context for their final mel frames."""
66
+ waveforms = self._audio_batch(audio, sampling_rate)
67
+ hop = self.feature_extractor.hop_length
68
+ longest = max(len(waveform) for waveform in waveforms)
69
+ aligned_samples = ((longest + hop - 1) // hop) * hop
70
+ original_padding = max(self.feature_extractor.n_samples, aligned_samples)
71
+ context_padding = ((aligned_samples + self.feature_extractor.n_fft + hop - 1) // hop) * hop
72
+ # Preserve the original STFT reflection boundary for full-length recordings.
73
+ padded_samples = min(original_padding, context_padding)
74
+ return self.feature_extractor(
75
+ waveforms, sampling_rate=self.config.sampling_rate,
76
+ padding="max_length", max_length=int(padded_samples), truncation=False,
77
+ return_attention_mask=True, return_tensors=return_tensors,
78
+ )
79
+
80
+ @staticmethod
81
+ def assistant_prefix(enable_thinking=False):
82
+ prefix = "<|im_start|>assistant\n<think>\n"
83
+ return prefix if enable_thinking else prefix + "\n</think>\n\n"
84
+
85
+ @staticmethod
86
+ def _batch_flags(value, batch_size, name):
87
+ flags = [value] * batch_size if isinstance(value, bool) else list(value)
88
+ if len(flags) != batch_size or any(not isinstance(flag, bool) for flag in flags):
89
+ raise ValueError(f"{name} must be a bool or one bool per input")
90
+ return flags
91
+
92
+ def __call__(
93
+ self,
94
+ text=None,
95
+ audio=None,
96
+ sampling_rate=16000,
97
+ task="qa",
98
+ system_prompt=None,
99
+ enable_thinking=False,
100
+ audio_first=True,
101
+ return_tensors="pt",
102
+ padding=True,
103
+ **kwargs,
104
+ ):
105
+ """Build a generation prompt for each text or speech input.
106
+
107
+ A string prompt is shared by all waveforms; a list supplies one prompt
108
+ per waveform. ``enable_thinking`` accepts a bool or one bool per input.
109
+ ``audio_first`` controls whether audio precedes the user text, with the
110
+ same scalar or per-input format.
111
+ Text padding defaults to the left for batched generation;
112
+ pass ``padding_side="right"`` for training batches.
113
+ """
114
+ if task not in {"qa", "asr"}:
115
+ raise ValueError(f"Unsupported audio task: {task}")
116
+ if text is None and audio is None:
117
+ raise ValueError("Provide text or audio")
118
+ audio_features = {}
119
+ if audio is not None:
120
+ waveforms = self._audio_batch(audio, sampling_rate)
121
+ features = self.audio_features(waveforms, sampling_rate=sampling_rate)
122
+ audio_features = {
123
+ "input_features": features["input_features"],
124
+ "feature_attention_mask": features["attention_mask"],
125
+ }
126
+ counts = [self.audio_token_count(length) for length in features["attention_mask"].sum(-1)]
127
+ prompts = [text or ""] * len(waveforms) if text is None or isinstance(text, str) else list(text)
128
+ if len(prompts) != len(waveforms):
129
+ raise ValueError("Provide one text prompt per waveform, or one shared string prompt")
130
+ thinking_flags = self._batch_flags(enable_thinking, len(prompts), "enable_thinking")
131
+ audio_first_flags = self._batch_flags(audio_first, len(prompts), "audio_first")
132
+ rendered = []
133
+ for prompt, count, thinking, first in zip(prompts, counts, thinking_flags, audio_first_flags):
134
+ user_prompt = prompt.strip() or ("Transcribe the speech." if task == "asr" else self.config.qa_prompt)
135
+ prefix = f"<|im_start|>system\n{system_prompt}<|im_end|>\n" if system_prompt and task != "asr" else ""
136
+ audio_prompt = f"<|audio_start|>{'<|audio_pad|>' * count}<|audio_end|>"
137
+ user_content = f"{audio_prompt}\n{user_prompt}" if first else f"{user_prompt}\n{audio_prompt}"
138
+ rendered.append(
139
+ f"{prefix}<|im_start|>user\n{user_content}<|im_end|>\n{self.assistant_prefix(thinking)}"
140
+ )
141
+ else:
142
+ prompts = [text] if isinstance(text, str) else list(text)
143
+ thinking_flags = self._batch_flags(enable_thinking, len(prompts), "enable_thinking")
144
+ rendered = []
145
+ for prompt, thinking in zip(prompts, thinking_flags):
146
+ messages = []
147
+ if system_prompt:
148
+ messages.append({"role": "system", "content": system_prompt})
149
+ messages.append({"role": "user", "content": prompt})
150
+ rendered.append(self.tokenizer.apply_chat_template(
151
+ messages, tokenize=False, add_generation_prompt=True, enable_thinking=thinking,
152
+ ))
153
+ kwargs.setdefault("padding_side", "left")
154
+ kwargs.setdefault("return_token_type_ids", False)
155
+ tokens = self.tokenizer(
156
+ rendered, add_special_tokens=False, padding=padding,
157
+ return_tensors=return_tensors, **kwargs,
158
+ )
159
+ return BatchFeature(data={**tokens, **audio_features}, tensor_type=return_tensors)
160
+
161
+ def decode(self, *args, **kwargs):
162
+ return self.tokenizer.decode(*args, **kwargs)
163
+
164
+ def batch_decode(self, *args, **kwargs):
165
+ return self.tokenizer.batch_decode(*args, **kwargs)
166
+
167
+ def to_dict(self):
168
+ return {
169
+ "processor_class": self.__class__.__name__,
170
+ "auto_map": {"AutoProcessor": "processing_edgeinstant.EdgeInstantProcessor"},
171
+ "image_processor_subfolder": "image_processor" if self.image_processor is not None else None,
172
+ }
173
+
174
+ def save_pretrained(self, save_directory, **kwargs):
175
+ directory = Path(save_directory)
176
+ directory.mkdir(parents=True, exist_ok=True)
177
+ self.config.save_pretrained(directory)
178
+ self.tokenizer.save_pretrained(directory, **kwargs)
179
+ self.feature_extractor.save_pretrained(directory)
180
+ if self.image_processor is not None:
181
+ self.image_processor.save_pretrained(directory / "image_processor")
182
+ if self.video_processor_config is not None:
183
+ (directory / "image_processor" / "video_preprocessor_config.json").write_text(
184
+ json.dumps(self.video_processor_config, indent=2) + "\n", encoding="utf-8",
185
+ )
186
+ processor_file = directory / "processor_config.json"
187
+ processor_file.write_text(json.dumps(self.to_dict(), indent=2) + "\n", encoding="utf-8")
188
+ custom_object_save(self, directory)
189
+ return [str(processor_file)]
190
+
191
+ @classmethod
192
+ def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
193
+ trust_remote_code = kwargs.pop("trust_remote_code", False)
194
+ kwargs.pop("_from_auto", None)
195
+ config = kwargs.pop("config", None)
196
+ hub_keys = {
197
+ "cache_dir", "force_download", "local_files_only", "token",
198
+ "revision", "subfolder", "proxies",
199
+ }
200
+ hub_kwargs = {key: value for key, value in kwargs.items() if key in hub_keys}
201
+ if config is None:
202
+ config = EdgeInstantConfig.from_pretrained(pretrained_model_name_or_path, **hub_kwargs)
203
+ tokenizer = AutoTokenizer.from_pretrained(
204
+ pretrained_model_name_or_path, config=config.thinker_config,
205
+ trust_remote_code=trust_remote_code, **kwargs,
206
+ )
207
+ feature_extractor = WhisperFeatureExtractor.from_pretrained(pretrained_model_name_or_path, **hub_kwargs)
208
+ metadata_path = Path(pretrained_model_name_or_path) / "processor_config.json"
209
+ if not metadata_path.is_file():
210
+ from transformers.utils.hub import cached_file
211
+ metadata_path = Path(cached_file(pretrained_model_name_or_path, "processor_config.json", **hub_kwargs))
212
+ metadata = json.loads(metadata_path.read_text())
213
+ image_processor = None
214
+ video_processor_config = None
215
+ if metadata.get("image_processor_subfolder"):
216
+ subfolder = str(Path(hub_kwargs.get("subfolder", "")) / metadata["image_processor_subfolder"])
217
+ image_kwargs = {**hub_kwargs, "subfolder": subfolder}
218
+ image_processor = AutoImageProcessor.from_pretrained(pretrained_model_name_or_path, **image_kwargs)
219
+ from transformers.utils.hub import cached_file
220
+ video_path = cached_file(pretrained_model_name_or_path, "video_preprocessor_config.json",
221
+ _raise_exceptions_for_missing_entries=False, **image_kwargs)
222
+ if video_path is not None:
223
+ video_processor_config = json.loads(Path(video_path).read_text())
224
+ return cls(tokenizer=tokenizer, feature_extractor=feature_extractor, config=config,
225
+ image_processor=image_processor, video_processor_config=video_processor_config)
226
+
227
+
228
+ EdgeInstantProcessor.register_for_auto_class()
processor_config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "processor_class": "EdgeInstantProcessor",
3
+ "auto_map": {
4
+ "AutoProcessor": "processing_edgeinstant.EdgeInstantProcessor"
5
+ },
6
+ "image_processor_subfolder": "image_processor"
7
+ }
requirements.txt ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ torch>=2.6
2
+ transformers==5.12.1
3
+ accelerate>=1.1.0
4
+ safetensors
5
+ numpy
6
+ soundfile
7
+ Pillow>=10
residual_predictor.py ADDED
@@ -0,0 +1,162 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ from torch import nn
7
+
8
+ from .sampling import sample_token
9
+
10
+
11
+ class Qwen3TTSResidualPredictor(nn.Module):
12
+ """The complete pretrained five-layer Qwen3-TTS code predictor."""
13
+
14
+ def __init__(self, hidden_size: int, codebook_size: int, config: dict[str, Any]):
15
+ super().__init__()
16
+ from transformers import Qwen3Config, Qwen3Model
17
+
18
+ if hidden_size != int(config["hidden_size"]) or codebook_size != int(
19
+ config["vocab_size"]
20
+ ):
21
+ raise ValueError("Qwen3-TTS residual predictor dimensions do not match")
22
+ backbone_config = Qwen3Config(
23
+ vocab_size=codebook_size,
24
+ hidden_size=hidden_size,
25
+ intermediate_size=int(config["intermediate_size"]),
26
+ num_hidden_layers=int(config["num_hidden_layers"]),
27
+ num_attention_heads=int(config["num_attention_heads"]),
28
+ num_key_value_heads=int(config["num_key_value_heads"]),
29
+ head_dim=int(config["head_dim"]),
30
+ hidden_act=config["hidden_act"],
31
+ max_position_embeddings=int(config["max_position_embeddings"]),
32
+ initializer_range=float(config["initializer_range"]),
33
+ rms_norm_eps=float(config["rms_norm_eps"]),
34
+ use_cache=True,
35
+ tie_word_embeddings=False,
36
+ rope_theta=float(config["rope_theta"]),
37
+ attention_bias=bool(config["attention_bias"]),
38
+ attention_dropout=float(config["attention_dropout"]),
39
+ )
40
+ backbone_config._attn_implementation = "eager"
41
+ self.backbone = Qwen3Model(backbone_config)
42
+ self.backbone.embed_tokens = None
43
+ self.q0_embedding = nn.Embedding(codebook_size, hidden_size)
44
+ self.code_embeddings = nn.ModuleList(
45
+ [nn.Embedding(codebook_size, hidden_size) for _ in range(14)]
46
+ )
47
+ self.heads = nn.ModuleList(
48
+ [nn.Linear(hidden_size, codebook_size, bias=False) for _ in range(15)]
49
+ )
50
+ self._greedy_graph = None
51
+ self._sampling_graph = None
52
+
53
+ @torch.no_grad()
54
+ def prepare_cuda_graph(self) -> None:
55
+ """Capture single-session greedy inference after loading the checkpoint."""
56
+ self._greedy_graph = self._capture_generation_graph(do_sample=False)
57
+
58
+ @torch.no_grad()
59
+ def prepare_sampling_cuda_graph(self) -> None:
60
+ """Capture top-k 50, top-p 1, temperature 0.9 with the default CUDA RNG."""
61
+ self._sampling_graph = self._capture_generation_graph(do_sample=True)
62
+
63
+ def _capture_generation_graph(self, *, do_sample: bool) -> tuple:
64
+ weight = self.q0_embedding.weight
65
+ if self.training or weight.device.type != "cuda":
66
+ raise ValueError("Residual CUDA graph requires a CUDA model in eval mode")
67
+ with torch.cuda.device(weight.device):
68
+ hidden = weight.new_zeros((1, weight.shape[1]))
69
+ q0 = torch.zeros(1, device=weight.device, dtype=torch.long)
70
+ masks = [weight.new_zeros((1, 1, 2, 2))]
71
+ masks[0][0, 0, 0, 1] = torch.finfo(weight.dtype).min
72
+ masks.extend(weight.new_zeros((1, 1, 1, n)) for n in range(3, 17))
73
+ stream = torch.cuda.Stream(device=weight.device)
74
+ stream.wait_stream(torch.cuda.current_stream())
75
+ with torch.cuda.stream(stream):
76
+ for _ in range(3):
77
+ self._generate(hidden, q0, do_sample=do_sample, attention_masks=masks)
78
+ torch.cuda.current_stream().wait_stream(stream)
79
+ graph = torch.cuda.CUDAGraph()
80
+ with torch.cuda.graph(graph):
81
+ output = self._generate(hidden, q0, do_sample=do_sample, attention_masks=masks)
82
+ # Keep all captured inputs alive for subsequent replays.
83
+ return graph, hidden, q0, output, masks
84
+
85
+ def forward(self, hidden: torch.Tensor, frame_prefix: torch.Tensor) -> torch.Tensor:
86
+ if frame_prefix.ndim != 2 or frame_prefix.shape[1] != 16:
87
+ raise ValueError("frame_prefix must be [N, 16]")
88
+ inputs = [hidden.unsqueeze(1), self.q0_embedding(frame_prefix[:, :1])]
89
+ inputs.extend(
90
+ embedding(frame_prefix[:, group : group + 1])
91
+ for group, embedding in enumerate(self.code_embeddings, start=1)
92
+ )
93
+ outputs = self.backbone(
94
+ inputs_embeds=torch.cat(inputs, dim=1), use_cache=False
95
+ ).last_hidden_state
96
+ return torch.stack(
97
+ [head(outputs[:, group + 1]) for group, head in enumerate(self.heads)],
98
+ dim=1,
99
+ )
100
+
101
+ @torch.no_grad()
102
+ def generate(
103
+ self,
104
+ hidden: torch.Tensor,
105
+ q0: torch.Tensor,
106
+ *,
107
+ do_sample: bool = False,
108
+ top_k: int = 50,
109
+ top_p: float = 1.0,
110
+ temperature: float = 0.9,
111
+ generator: torch.Generator | None = None,
112
+ ) -> torch.Tensor:
113
+ captured = self._greedy_graph if not do_sample else (
114
+ self._sampling_graph
115
+ if generator is None and top_k == 50 and top_p == 1.0 and temperature == 0.9
116
+ else None
117
+ )
118
+ if captured is not None and hidden.shape[0] == 1:
119
+ graph, static_hidden, static_q0, output, _ = captured
120
+ static_hidden.copy_(hidden)
121
+ static_q0.copy_(q0)
122
+ graph.replay()
123
+ return output.clone()
124
+ return self._generate(
125
+ hidden, q0, do_sample=do_sample, top_k=top_k, top_p=top_p,
126
+ temperature=temperature, generator=generator,
127
+ )
128
+
129
+ def _generate(
130
+ self, hidden: torch.Tensor, q0: torch.Tensor, *,
131
+ do_sample: bool = False, top_k: int = 50, top_p: float = 1.0,
132
+ temperature: float = 0.9, generator: torch.Generator | None = None,
133
+ attention_masks: list[torch.Tensor] | None = None,
134
+ ) -> torch.Tensor:
135
+ inputs = torch.cat(
136
+ [hidden.unsqueeze(1), self.q0_embedding(q0).unsqueeze(1)], dim=1
137
+ )
138
+ output = self.backbone(
139
+ inputs_embeds=inputs, use_cache=True,
140
+ attention_mask={"full_attention": attention_masks[0]} if attention_masks else None,
141
+ )
142
+ cache = output.past_key_values
143
+ codes = [q0]
144
+ for group, head in enumerate(self.heads):
145
+ code = sample_token(
146
+ head(output.last_hidden_state[:, -1]),
147
+ do_sample=do_sample,
148
+ top_k=top_k,
149
+ top_p=top_p,
150
+ temperature=temperature,
151
+ generator=generator,
152
+ )
153
+ codes.append(code)
154
+ if group < len(self.code_embeddings):
155
+ output = self.backbone(
156
+ inputs_embeds=self.code_embeddings[group](code).unsqueeze(1),
157
+ past_key_values=cache,
158
+ use_cache=True,
159
+ attention_mask={"full_attention": attention_masks[group + 1]} if attention_masks else None,
160
+ )
161
+ cache = output.past_key_values
162
+ return torch.stack(codes, dim=-1)
sampling.py ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import torch
4
+
5
+
6
+ def sample_token(
7
+ logits: torch.Tensor,
8
+ *,
9
+ do_sample: bool = False,
10
+ top_k: int = 50,
11
+ top_p: float = 1.0,
12
+ temperature: float = 0.9,
13
+ generator: torch.Generator | None = None,
14
+ ) -> torch.Tensor:
15
+ if not do_sample:
16
+ return logits.argmax(dim=-1)
17
+ scores = logits.float() / temperature
18
+ if 0 < top_k < scores.shape[-1]:
19
+ threshold = scores.topk(top_k, dim=-1).values[..., -1, None]
20
+ scores = scores.masked_fill(scores < threshold, -torch.inf)
21
+ if top_p < 1.0:
22
+ sorted_scores, sorted_indices = scores.sort(dim=-1, descending=True)
23
+ remove = sorted_scores.softmax(dim=-1).cumsum(dim=-1) > top_p
24
+ remove[..., 1:] = remove[..., :-1].clone()
25
+ remove[..., 0] = False
26
+ sorted_scores = sorted_scores.masked_fill(remove, -torch.inf)
27
+ scores = torch.full_like(scores, -torch.inf).scatter(
28
+ -1, sorted_indices, sorted_scores
29
+ )
30
+ return torch.multinomial(
31
+ scores.softmax(dim=-1), num_samples=1, generator=generator
32
+ ).squeeze(-1)
sequence.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ """Actions reported by native-token speech generation."""
2
+ from enum import IntEnum
3
+
4
+
5
+ class Action(IntEnum):
6
+ READ = 0
7
+ EMIT = 1
8
+ END = 2
teacher.py ADDED
@@ -0,0 +1,171 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from pathlib import Path
5
+
6
+ import numpy as np
7
+ import torch
8
+ from safetensors import safe_open
9
+ from torch import nn
10
+ import torch.nn.functional as F
11
+
12
+ from .residual_predictor import Qwen3TTSResidualPredictor
13
+
14
+
15
+ def native_teacher_mapping(
16
+ native_tokenizer_path: str | Path,
17
+ teacher_path: str | Path,
18
+ native_vocabulary_size: int | None = None,
19
+ ) -> tuple[torch.Tensor, torch.Tensor]:
20
+ """Map exact ByteLevel pieces to teacher IDs, with empty special-token bags."""
21
+ native_path = Path(native_tokenizer_path)
22
+ if native_path.is_dir():
23
+ native_path = native_path / "tokenizer.json"
24
+ native = json.loads(native_path.read_text())["model"]["vocab"]
25
+ teacher = json.loads((Path(teacher_path) / "vocab.json").read_text())
26
+ visible = list(range(33, 127)) + list(range(161, 173)) + list(range(174, 256))
27
+ inverse = {chr(byte): byte for byte in visible}
28
+ inverse.update({chr(256 + index): byte for index, byte in
29
+ enumerate(byte for byte in range(256) if byte not in visible)})
30
+
31
+ def raw(piece: str) -> bytes:
32
+ return bytes(inverse[character] for character in piece)
33
+
34
+ teacher_bytes = {raw(piece): index for piece, index in teacher.items()}
35
+ by_first_byte: dict[int, dict[bytes, int]] = {}
36
+ for piece, index in teacher_bytes.items():
37
+ by_first_byte.setdefault(piece[0], {})[piece] = index
38
+ widths = {first: sorted({len(piece) for piece in pieces}, reverse=True)
39
+ for first, pieces in by_first_byte.items()}
40
+ native_by_id = {index: raw(piece) for piece, index in native.items()}
41
+ size = native_vocabulary_size or (max(native_by_id) + 1)
42
+ flat, offsets = [], [0]
43
+ for native_id in range(size):
44
+ piece = native_by_id.get(native_id, b"")
45
+ if piece in teacher_bytes:
46
+ flat.append(teacher_bytes[piece])
47
+ else:
48
+ position = 0
49
+ while position < len(piece):
50
+ candidates = by_first_byte[piece[position]]
51
+ for width in widths[piece[position]]:
52
+ if width > len(piece) - position:
53
+ continue
54
+ candidate = piece[position:position + width]
55
+ if candidate in candidates:
56
+ flat.append(candidates[candidate])
57
+ position += width
58
+ break
59
+ else:
60
+ raise ValueError(f"Teacher vocabulary cannot represent native token {native_id}")
61
+ offsets.append(len(flat))
62
+ return torch.tensor(flat, dtype=torch.long), torch.tensor(offsets, dtype=torch.long)
63
+
64
+
65
+ class TeacherTextProjection(nn.Module):
66
+ def __init__(self, text_hidden_size: int, hidden_size: int):
67
+ super().__init__()
68
+ self.linear_fc1 = nn.Linear(text_hidden_size, text_hidden_size)
69
+ self.linear_fc2 = nn.Linear(text_hidden_size, hidden_size)
70
+
71
+ def forward(self, hidden: torch.Tensor) -> torch.Tensor:
72
+ return self.linear_fc2(F.silu(self.linear_fc1(hidden)))
73
+
74
+
75
+ class TeacherAcousticModel(nn.Module):
76
+ """Pretrained text, acoustic and residual modules shared by speech models."""
77
+
78
+ codebook_size = 2048
79
+ num_code_groups = 16
80
+
81
+ def __init__(
82
+ self,
83
+ config: dict,
84
+ native_teacher_ids: torch.Tensor,
85
+ native_teacher_offsets: torch.Tensor,
86
+ speaker_vector: torch.Tensor,
87
+ *,
88
+ attn_implementation: str = "eager",
89
+ ):
90
+ super().__init__()
91
+ from transformers import Qwen3Config, Qwen3Model
92
+
93
+ self.config = dict(config)
94
+ self.hidden_size = int(config["hidden_size"])
95
+ backbone_config = Qwen3Config(
96
+ vocab_size=self.codebook_size,
97
+ hidden_size=self.hidden_size,
98
+ intermediate_size=int(config["intermediate_size"]),
99
+ num_hidden_layers=int(config["num_hidden_layers"]),
100
+ num_attention_heads=int(config["num_attention_heads"]),
101
+ num_key_value_heads=int(config["num_key_value_heads"]),
102
+ head_dim=int(config["head_dim"]),
103
+ hidden_act=config["hidden_act"],
104
+ max_position_embeddings=int(config["max_position_embeddings"]),
105
+ rms_norm_eps=float(config["rms_norm_eps"]),
106
+ rope_theta=float(config["rope_theta"]),
107
+ attention_bias=bool(config["attention_bias"]),
108
+ attention_dropout=float(config.get("attention_dropout", 0.0)),
109
+ tie_word_embeddings=False,
110
+ use_cache=True,
111
+ )
112
+ backbone_config._attn_implementation = attn_implementation
113
+ self.backbone = Qwen3Model(backbone_config)
114
+ self.backbone.embed_tokens = None
115
+ self.text_embedding = nn.Embedding(int(config["text_vocab_size"]), int(config["text_hidden_size"]))
116
+ self.text_projection = TeacherTextProjection(int(config["text_hidden_size"]), self.hidden_size)
117
+ self.text_embedding.requires_grad_(False)
118
+ self.text_projection.requires_grad_(False)
119
+ self.register_buffer("native_teacher_ids", native_teacher_ids, persistent=False)
120
+ self.register_buffer("native_teacher_offsets", native_teacher_offsets, persistent=False)
121
+ self.register_buffer("speaker_vector", speaker_vector.reshape(self.hidden_size), persistent=False)
122
+ self.codec_history_embeddings = nn.ModuleList(
123
+ [nn.Embedding(self.codebook_size, self.hidden_size) for _ in range(self.num_code_groups)]
124
+ )
125
+ self.codec_bos = nn.Parameter(torch.zeros(self.hidden_size))
126
+ self.q0_head = nn.Linear(self.hidden_size, self.codebook_size, bias=False)
127
+ self.residual_predictor = Qwen3TTSResidualPredictor(
128
+ self.hidden_size, self.codebook_size, config["code_predictor_config"]
129
+ )
130
+
131
+ @classmethod
132
+ def from_teacher(
133
+ cls,
134
+ teacher_path: str | Path,
135
+ native_tokenizer_path: str | Path,
136
+ speaker_vector_path: str | Path,
137
+ *,
138
+ dtype: torch.dtype = torch.float32,
139
+ attn_implementation: str = "eager",
140
+ ) -> "TeacherAcousticModel":
141
+ teacher_path, native_tokenizer_path = Path(teacher_path), Path(native_tokenizer_path)
142
+ config = json.loads((teacher_path / "config.json").read_text())["talker_config"]
143
+ native_root = native_tokenizer_path if native_tokenizer_path.is_dir() else native_tokenizer_path.parent
144
+ native_config = json.loads((native_root / "config.json").read_text())
145
+ native_size = int(native_config.get("text_config", native_config)["vocab_size"])
146
+ ids, offsets = native_teacher_mapping(native_tokenizer_path, teacher_path, native_size)
147
+ speaker = torch.as_tensor(np.load(speaker_vector_path), dtype=torch.float32)
148
+ model = cls(config, ids, offsets, speaker, attn_implementation=attn_implementation).to(dtype=dtype)
149
+ with safe_open(teacher_path / "model.safetensors", framework="pt", device="cpu") as weights:
150
+ with torch.no_grad():
151
+ for name, parameter in model.backbone.named_parameters():
152
+ parameter.copy_(weights.get_tensor(f"talker.model.{name}"))
153
+ model.text_embedding.weight.copy_(weights.get_tensor("talker.model.text_embedding.weight"))
154
+ for name, parameter in model.text_projection.named_parameters():
155
+ parameter.copy_(weights.get_tensor(f"talker.text_projection.{name}"))
156
+ codec = weights.get_tensor("talker.model.codec_embedding.weight")
157
+ model.codec_history_embeddings[0].weight.copy_(codec[:model.codebook_size])
158
+ model.codec_bos.copy_(codec[int(config["codec_bos_id"])])
159
+ model.q0_head.weight.copy_(weights.get_tensor("talker.codec_head.weight")[:model.codebook_size])
160
+ model.residual_predictor.q0_embedding.weight.copy_(codec[:model.codebook_size])
161
+ for name, parameter in model.residual_predictor.backbone.named_parameters():
162
+ parameter.copy_(weights.get_tensor(f"talker.code_predictor.model.{name}"))
163
+ for group in range(15):
164
+ embedding = weights.get_tensor(f"talker.code_predictor.model.codec_embedding.{group}.weight")
165
+ model.codec_history_embeddings[group + 1].weight.copy_(embedding)
166
+ if group < 14:
167
+ model.residual_predictor.code_embeddings[group].weight.copy_(embedding)
168
+ model.residual_predictor.heads[group].weight.copy_(
169
+ weights.get_tensor(f"talker.code_predictor.lm_head.{group}.weight")
170
+ )
171
+ return model
tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6f32ce20dc35f57a7f9ad1eac03525bd7d30f9df8cea6507e958279cc3657706
3
+ size 19989492
tokenizer_config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "audio_bos_token": "<|audio_start|>",
4
+ "audio_eos_token": "<|audio_end|>",
5
+ "audio_token": "<|audio_pad|>",
6
+ "backend": "tokenizers",
7
+ "bos_token": null,
8
+ "clean_up_tokenization_spaces": false,
9
+ "eos_token": "<|im_end|>",
10
+ "errors": "replace",
11
+ "image_token": "<|image_pad|>",
12
+ "is_local": true,
13
+ "local_files_only": true,
14
+ "max_length": null,
15
+ "model_max_length": 262144,
16
+ "model_specific_special_tokens": {
17
+ "audio_bos_token": "<|audio_start|>",
18
+ "audio_eos_token": "<|audio_end|>",
19
+ "audio_token": "<|audio_pad|>",
20
+ "image_token": "<|image_pad|>",
21
+ "video_token": "<|video_pad|>",
22
+ "vision_bos_token": "<|vision_start|>",
23
+ "vision_eos_token": "<|vision_end|>"
24
+ },
25
+ "pad_to_multiple_of": null,
26
+ "pad_token": "<|endoftext|>",
27
+ "pad_token_type_id": 0,
28
+ "padding_side": "left",
29
+ "pretokenize_regex": "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
30
+ "split_special_tokens": false,
31
+ "tokenizer_class": "Qwen2Tokenizer",
32
+ "unk_token": null,
33
+ "video_token": "<|video_pad|>",
34
+ "vision_bos_token": "<|vision_start|>",
35
+ "vision_eos_token": "<|vision_end|>"
36
+ }