Audio-Text-to-Text
Transformers
Safetensors
Chinese
English
edgeinstant
feature-extraction
audio
speech-recognition
speech-translation
audio-question-answering
custom_code
Instructions to use chenjz24/EdgeIn-v1 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use chenjz24/EdgeIn-v1 with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("chenjz24/EdgeIn-v1", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Upload folder using huggingface_hub
Browse files- .gitattributes +1 -0
- INFERENCE_OPTIMIZATION.md +85 -0
- LICENSE.codec +201 -0
- README.md +72 -0
- chat_template.jinja +154 -0
- codec_edgeinstant.py +233 -0
- compact_native_clock.py +174 -0
- compaction.json +23 -0
- config.json +0 -0
- configuration_edgeinstant.py +95 -0
- generation.py +22 -0
- generation_config.json +13 -0
- image_processor/preprocessor_config.json +26 -0
- image_processor/video_preprocessor_config.json +21 -0
- inference_edgeinstant.py +65 -0
- model-00001-of-00003.safetensors +3 -0
- model-00002-of-00003.safetensors +3 -0
- model-00003-of-00003.safetensors +3 -0
- model.safetensors.index.json +0 -0
- modeling_edgeinstant.py +392 -0
- native_clock_talker.py +257 -0
- preprocessor_config.json +14 -0
- processing_edgeinstant.py +228 -0
- processor_config.json +7 -0
- requirements.txt +7 -0
- residual_predictor.py +162 -0
- sampling.py +32 -0
- sequence.py +8 -0
- teacher.py +171 -0
- tokenizer.json +3 -0
- tokenizer_config.json +36 -0
.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 |
+
}
|