TaichuAI commited on
Commit
f5f6734
·
verified ·
1 Parent(s): 1ca15dc

Upload folder using huggingface_hub

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ tokenizer.json filter=lfs diff=lfs merge=lfs -text
.ms_upload_cache ADDED
@@ -0,0 +1 @@
 
 
1
+ {"version": 3, "repo_id": "TaichuAI/ZDTaichu5.0-9B-NVFP4", "files": {"README.md|1789460761.0|23258": {"hash": "b11f22ed8c86f05feb2e7ce747b69ee89f42b28b407d4b158cc299cd1ad489c1", "size": 23258, "status": "c"}, ".vscode/settings.json|1789460406.0|580": {"hash": "1d1a47d4cdda5f67a9302d196d91f6ed8d209972d1fb7d2d7618a9d8e49aee14", "size": 580, "status": "c"}, "assets/logos/gemini.png|1789460699.0|19333": {"hash": "10c628f55d22a9725b9f9fccce7cf062b9fb68da5f7736e87010e0594d7ba6db", "size": 19333, "status": "c"}, "__init__.py|1789455331.0|487": {"hash": "854a695a9b3732f6a4841bacfa83463f6b0e51ffbba692bb3ada8e647e3e37aa", "size": 487, "status": "c"}, "assets/logos/gemma.png|1789460698.0|16826": {"hash": "5f64463f9e00b29595c3ce7222100e263b67c3f29035c304d1249f30c99d2913", "size": 16826, "status": "c"}, "assets/logos/grok.png|1789460699.0|6883": {"hash": "f10ec0b2322f9d65a973d3d48fcdc41270d8b123a5d36e2826dabf9dcefb0762", "size": 6883, "status": "c"}, "assets/logos/qwen.png|1789460699.0|10324": {"hash": "de9bc7e285164e0d284a7b9555511c7f7767af7699d6272b3b67b376eeddbfb3", "size": 10324, "status": "c"}, "config.json|1789455731.0|24594": {"hash": "699af1e41b2054269f2b698f11950c5f35f46d02679a7b9d55c0301e8a52e823", "size": 24594, "status": "c"}, "assets/logos/taichu.png|1789460699.0|46894": {"hash": "90da889f0994e50a92677bf4141e6d58bc24c221547e0828d91e2889fc1bc47c", "size": 46894, "status": "c"}, "cradio_config.py|1789455331.0|2049": {"hash": "4971a6d6c0abb55eaddd0cb94753f4941f50928804b1766c09d7b9bb98f51dfe", "size": 2049, "status": "c"}, "assets/logos/stepfun.png|1789460699.0|6937": {"hash": "49e54bcda77c23f9a8ceacc0d3c1c07da09be4c2534d453ebcad7da201b25b96", "size": 6937, "status": "c"}, "chat_template.jinja|1789455319.0|7756": {"hash": "a4aee8afcf2e0711942cf848899be66016f8d14a889ff9ede07bca099c28f715", "size": 7756, "status": "c"}, "assets/taichu-vs-closed-models.svg|1789460699.0|403989": {"hash": "d80d5e9c375d397d268bfba0db7000cdb9ee558a890ee3fe6b6c38a49b3cffed", "size": 403989, "status": "c"}, "cradio_model.py|1789455331.0|26514": {"hash": "4c4d3bafb13cd47203ffacb78e1932c809cb48cc254b0f042cbb90121f8827e8", "size": 26514, "status": "c"}, "assets/taichu-release-benchmark-comparison.svg|1789460698.0|569738": {"hash": "7109868f2f09641ee86a0b12dba6f56f56ffa3b7c71bbb8fa0db655b21876fe0", "size": 569738, "status": "c"}, "configuration.py|1789455331.0|9699": {"hash": "57a04f680bcbaede0791bfc8d2653a7f82c11cbf8e38c68b00452b315a4ce0f0", "size": 9699, "status": "c"}, "generation_config.json|1789455327.0|284": {"hash": "fc84e407226adee495fc06cb354223d4912fea7ba6e8cee1b0b5f2c08e04badc", "size": 284, "status": "c"}, "image_processing.py|1789455331.0|10095": {"hash": "a68a34f5426b1886f5904f0b75940737a67caec2006a3ddf9387a199b6b10303", "size": 10095, "status": "c"}, "modeling.py|1789455331.0|50552": {"hash": "b3f46b9f1635663778eb30bff8b90c474d45b3186f469367df058a5f6a7117fa", "size": 50552, "status": "c"}, "processor_config.json|1789455327.0|724": {"hash": "ca2ea029533bcad96c9e8650e1c32d7412af9d51d934f1358153d20c4fdf5c01", "size": 724, "status": "c"}, "preprocessor_config.json|1789455327.0|528": {"hash": "a9bd6d2e6a32ed98335351904c6f129613f7c14f22bd556d4bf660b461112b9b", "size": 528, "status": "c"}, "processing.py|1789455331.0|27232": {"hash": "0e1fe5d7706886a4520604f3cdb0d7402866158db3636997ee7bb798587ad5f6", "size": 27232, "status": "c"}, "recipe.yaml|1789455224.0|3205": {"hash": "a396a0d93e140200410815c55411335b3a901f11dbe68672526dc7561641af98", "size": 3205, "status": "c"}, "tokenizer_config.json|1789455327.0|930": {"hash": "78311f543bf15a7153eff52ca8d692dbaf5bdc094f265d082ec3909807cc3db7", "size": 930, "status": "c"}, "vision_utils.py|1789455331.0|22068": {"hash": "f4eb3b4b28886f358691976d7910833c811c92f594a3c247985d1040f3bac6b6", "size": 22068, "status": "c"}, "tokenizer.json|1789455327.0|19989343": {"hash": "87a7830d63fcf43bf241c3c5242e96e62dd3fdc29224ca26fed8ea333db72de4", "size": 19989343, "status": "c"}, "model.safetensors|1789455210.0|9811274448": {"hash": "d26e50e66812bdb9563368991b4ad9be1eeaab8aa59b5220df5efbc89e202f3b", "size": 9811274448, "status": "c"}}}
.vscode/settings.json ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "workbench.colorCustomizations": {
3
+ "activityBar.background": "#1E1E5C",
4
+ "titleBar.activeBackground": "#2A2A81",
5
+ "titleBar.activeForeground": "#FCFCFE",
6
+ "titleBar.inactiveBackground": "#1E1E5C",
7
+ "titleBar.inactiveForeground": "#FCFCFE",
8
+ "statusBar.background": "#1E1E5C",
9
+ "statusBar.foreground": "#FCFCFE",
10
+ "statusBar.debuggingBackground": "#1E1E5C",
11
+ "statusBar.debuggingForeground": "#FCFCFE",
12
+ "statusBar.noFolderBackground": "#1E1E5C",
13
+ "statusBar.noFolderForeground": "#FCFCFE"
14
+ }
15
+ }
README.md CHANGED
@@ -1,3 +1,655 @@
1
  ---
2
- license: apache-2.0
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ language:
3
+ - en
4
+ - zh
5
+ pipeline_tag: image-text-to-text
6
+ tags:
7
+ - multimodal
8
+ - vision-language-model
9
+ - spatial-reasoning
10
+ - agent
11
+ - video-understanding
12
  ---
13
+
14
+ # ZDTaichu5.0-9B-NVFP4
15
+
16
+ > [!Note]
17
+ > This repository contains NVFP4-quantized model weights and configuration files for the post-trained model in the compressed-tensors format.
18
+ >
19
+ > These artifacts are compatible with Hugging Face Transformers, vLLM, etc.
20
+ >
21
+ > The quantization method is mixed NVFP4/FP8 quantization, with NVFP4 (group size 16, dynamic per-group activations) for selected MLP projections and FP8 (static per-channel weights, dynamic per-token activations) for the other linear layers, and its performance metrics are nearly identical to those of the original model.
22
+
23
+ [Project Page](https://taichu-ai.github.io/ZDTaichu5.0-9B/) | [GitHub](https://github.com/Taichu-AI/ZDTaichu5.0-9B) | [ModelScope](https://www.modelscope.cn/models/TaichuAI/ZDTaichu5.0-9B)
24
+
25
+ ZDTaichu5.0-9B is a multimodal foundation model for general visual understanding, spatial reasoning, agentic tool use, and embodied-AI research. It combines a Qwen3.5-9B language backbone with a C-RADIOv4-H vision encoder, supports text, images and videos with any-resolution visual input.
26
+
27
+ Within the 10B-scale general-purpose VLMs compared in this release blog, ZDTaichu5.0-9B retains first-tier general visual understanding while supporting spatial reasoning, high-level embodied VLM reasoning, and agent tasks under the reported evaluation settings. Rather than trading broad visual competence for specialization, it layers a more comprehensive spatial, embodied, and agent capability profile on top of a strong general-vision foundation.
28
+
29
+ The model accepts text, one or more images, and video. It is designed for:
30
+
31
+ - general image, document, chart, diagram, and OCR understanding;
32
+ - visual mathematics and knowledge-grounded visual question answering;
33
+ - fine-grained 2D relations, multi-view association, 3D scene understanding, perspective taking, and mental transformation;
34
+ - multi-step and multi-turn tool use;
35
+ - spatial perception, affordance understanding, and planning for VLA and embodied-AI adaptation.
36
+
37
+ More demos and showcases are provided at [Project Page](https://taichu-ai.github.io/ZDTaichu5.0-9B/).
38
+
39
+ ## Highlights
40
+
41
+ - **Strong general vision and broad capabilities:** remains in the leading group of 10B-scale general-purpose VLMs across images, documents, charts, diagrams, OCR, visual mathematics, multiple images and video, while extending to spatial reasoning, high-level embodied understanding and multi-step agent tasks.
42
+ - **Leading spatial reasoning and embodied understanding:** leads spatial capability among the compared 10B-scale general-purpose VLMs, with strong results on SparBench, ViewSpatial, MMSI-Bench and MindCube-tiny. Scores of 48 on ERQA and 56 on RoboSpatial cover scene reasoning, affordances and interaction-oriented understanding.
43
+ - **Strongest agent capability among the compared 10B-scale general-purpose VLMs:** leads the reported TAU2-Bench (87.7) and Claw-Eval (71.4) comparisons, and reaches 93.7 on IFEval.
44
+ - **Entropy-Gated Adaptive Recurrent Reasoning:** Dynamically allocates additional recurrent refinement steps in latent space to more challenging tokens, enabling greater computational depth where needed and improving reasoning performance on complex tasks.
45
+
46
+ ## Model Overview
47
+
48
+ | Item | Specification |
49
+ |---|---|
50
+ | Model type | Multimodal causal language model with vision encoder |
51
+ | Language backbone | Qwen3.5-9B LLM Decoder|
52
+ | Vision backbone | C-RADIOv4-H |
53
+ | Context length | Up to 128K tokens |
54
+ | Vision resolution | Any-resolution visual input |
55
+ | Input modalities | Text, single image, multiple images, and video |
56
+
57
+ ## Capabilities
58
+
59
+ ### General visual understanding
60
+
61
+ The model can recognize objects, attributes, and scenes; read text in natural images and documents; interpret tables, forms, plots, and diagrams; and answer questions that combine visual evidence with language and world knowledge.
62
+
63
+ ### Spatial perception and reasoning
64
+
65
+ Spatial training covers:
66
+
67
+ - left/right, above/below, front/behind, occlusion, containment, and relative distance;
68
+ - dense counting, fine-grained localization, points, coordinates, and bounding boxes;
69
+ - association across images and viewpoints;
70
+ - camera motion, relative pose, depth ordering, and room-scale layout;
71
+ - egocentric and allocentric perspective taking;
72
+ - 2D/3D rotation, paper folding, three-view projection, cross-sections, and part-motion reasoning;
73
+ - embodied affordances, manipulation semantics, and high-level action planning.
74
+
75
+ ### Multiple images and video
76
+
77
+ ZDTaichu5.0-9B compares and reasons across multiple images and supports video understanding, including event tracking and detail retrieval from long footage within its 128K-token context window.
78
+
79
+ ### Agentic tool use
80
+
81
+ The model is designed for multi-step and multi-turn tool-use tasks. Tool execution must be implemented, validated, and secured by the surrounding application; the model does not execute tools by itself.
82
+
83
+ ## Benchmark Results
84
+
85
+
86
+ The two figures compare ZDTaichu5.0-9B with open and closed models across general visual understanding, spatial and embodied capabilities, and agent and text capabilities.
87
+
88
+ **Comparison with open models**
89
+
90
+ ![ZDTaichu5.0-9B benchmark comparison with open models](assets/taichu-release-benchmark-comparison.svg)
91
+
92
+ **Comparison with closed models**
93
+
94
+ ![ZDTaichu5.0-9B benchmark comparison with closed models](assets/taichu-vs-closed-models.svg)
95
+
96
+
97
+
98
+ ### Spatial and embodied reasoning
99
+
100
+ <table>
101
+ <thead>
102
+ <tr>
103
+ <th align="left">Area</th>
104
+ <th align="left">Benchmark</th>
105
+ <th align="right">ZDTaichu5.0-9B</th>
106
+ <th align="right">Qwen3.5-9B</th>
107
+ <th align="right">STEP3-VL-10B</th>
108
+ <th align="right">gemma4-8B-E4B</th>
109
+ <th align="right">Gemini 3 Pro</th>
110
+ <th align="right">Grok 4</th>
111
+ <th align="right">GPT-5.2</th>
112
+ </tr>
113
+ </thead>
114
+ <tbody>
115
+ <tr>
116
+ <td rowspan="3" align="left" valign="middle">Basic spatial perception</td>
117
+ <td align="left">CV-Bench</td>
118
+ <td align="right">86.82</td>
119
+ <td align="right"><strong>87.19</strong></td>
120
+ <td align="right">83.49</td>
121
+ <td align="right">68.10</td>
122
+ <td align="right"><ins>90.07</ins></td>
123
+ <td align="right">—</td>
124
+ <td align="right">86.84</td>
125
+ </tr>
126
+ <tr>
127
+ <td align="left">3DSRBench</td>
128
+ <td align="right"><strong>60.96</strong></td>
129
+ <td align="right">56.78</td>
130
+ <td align="right">55.01</td>
131
+ <td align="right">53.62</td>
132
+ <td align="right"><ins>68.92</ins></td>
133
+ <td align="right">54.93</td>
134
+ <td align="right">60.20</td>
135
+ </tr>
136
+ <tr>
137
+ <td align="left">SparBench</td>
138
+ <td align="right"><strong>51.82</strong></td>
139
+ <td align="right">50.79</td>
140
+ <td align="right">45.68</td>
141
+ <td align="right">28.50</td>
142
+ <td align="right">48.74</td>
143
+ <td align="right">44.76</td>
144
+ <td align="right"><ins>55.07</ins></td>
145
+ </tr>
146
+ <tr>
147
+ <td rowspan="3" align="left" valign="middle">Complex spatial reasoning</td>
148
+ <td align="left">ViewSpatial</td>
149
+ <td align="right"><strong><ins>62.50</ins></strong></td>
150
+ <td align="right">48.20</td>
151
+ <td align="right">46.14</td>
152
+ <td align="right">41.68</td>
153
+ <td align="right">50.36</td>
154
+ <td align="right">43.23</td>
155
+ <td align="right">47.30</td>
156
+ </tr>
157
+ <tr>
158
+ <td align="left">MMSI-Bench</td>
159
+ <td align="right"><strong><ins>47.20</ins></strong></td>
160
+ <td align="right">38.70</td>
161
+ <td align="right">32.18</td>
162
+ <td align="right">29.20</td>
163
+ <td align="right">45.20</td>
164
+ <td align="right">37.80</td>
165
+ <td align="right">41.30</td>
166
+ </tr>
167
+ <tr>
168
+ <td align="left">MindCube-tiny</td>
169
+ <td align="right"><strong><ins>78.27</ins></strong></td>
170
+ <td align="right">57.60</td>
171
+ <td align="right">62.81</td>
172
+ <td align="right">48.85</td>
173
+ <td align="right">70.87</td>
174
+ <td align="right">63.56</td>
175
+ <td align="right">60.38</td>
176
+ </tr>
177
+ <tr>
178
+ <td rowspan="3" align="left" valign="middle">Embodied interaction</td>
179
+ <td align="left">ERQA</td>
180
+ <td align="right"><strong>48.00</strong></td>
181
+ <td align="right">41.50</td>
182
+ <td align="right">47.75</td>
183
+ <td align="right">30.20</td>
184
+ <td align="right"><ins>66.00</ins></td>
185
+ <td align="right">—</td>
186
+ <td align="right">59.80</td>
187
+ </tr>
188
+ <tr>
189
+ <td align="left">RoboSpatial</td>
190
+ <td align="right"><strong>56.00</strong></td>
191
+ <td align="right">54.10</td>
192
+ <td align="right">52.86</td>
193
+ <td align="right">49.43</td>
194
+ <td align="right"><ins>57.40</ins></td>
195
+ <td align="right">—</td>
196
+ <td align="right">43.78</td>
197
+ </tr>
198
+ <tr>
199
+ <td align="left">VSI-Bench</td>
200
+ <td align="right"><strong><ins>59.69</ins></strong></td>
201
+ <td align="right">55.68</td>
202
+ <td align="right">42.42</td>
203
+ <td align="right">32.91</td>
204
+ <td align="right">52.51</td>
205
+ <td align="right">47.92</td>
206
+ <td align="right">54.49</td>
207
+ </tr>
208
+ </tbody>
209
+ </table>
210
+
211
+ ### General visual understanding
212
+
213
+ <table>
214
+ <thead>
215
+ <tr>
216
+ <th align="left">Area</th>
217
+ <th align="left">Benchmark</th>
218
+ <th align="right">ZDTaichu5.0-9B</th>
219
+ <th align="right">Qwen3.5-9B</th>
220
+ <th align="right">STEP3-VL-10B</th>
221
+ <th align="right">gemma4-8B-E4B</th>
222
+ <th align="right">Gemini 3 Pro</th>
223
+ <th align="right">Grok 4</th>
224
+ <th align="right">GPT-5.2</th>
225
+ </tr>
226
+ </thead>
227
+ <tbody>
228
+ <tr>
229
+ <td align="left" rowspan="3" valign="middle">Multi modal Reasoning</td>
230
+ <td align="left">MathVista Mini</td>
231
+ <td align="right">84.50</td>
232
+ <td align="right"><strong>85.70</strong></td>
233
+ <td align="right">83.97</td>
234
+ <td align="right">65.30</td>
235
+ <td align="right"><ins>87.90</ins></td>
236
+ <td align="right">72.50</td>
237
+ <td align="right">83.10</td>
238
+ </tr>
239
+ <tr>
240
+ <td align="left">WeMath</td>
241
+ <td align="right"><strong>75.90</strong></td>
242
+ <td align="right">75.20</td>
243
+ <td align="right">73.03</td>
244
+ <td align="right">50.19</td>
245
+ <td align="right"><ins>86.90</ins></td>
246
+ <td align="right">—</td>
247
+ <td align="right">79.00</td>
248
+ </tr>
249
+ <tr>
250
+ <td align="left">MathVerse Mini Vision Only</td>
251
+ <td align="right">76.40</td>
252
+ <td align="right"><strong><ins>84.14</ins></strong></td>
253
+ <td align="right">74.60</td>
254
+ <td align="right">53.55</td>
255
+ <td align="right">—</td>
256
+ <td align="right">—</td>
257
+ <td align="right">—</td>
258
+ </tr>
259
+ <tr>
260
+ <td align="left" rowspan="3" valign="middle">General VQA</td>
261
+ <td align="left">MMStar</td>
262
+ <td align="right">76.80</td>
263
+ <td align="right"><strong>79.70</strong></td>
264
+ <td align="right">77.48</td>
265
+ <td align="right">62.00</td>
266
+ <td align="right"><ins>83.10</ins></td>
267
+ <td align="right">69.60</td>
268
+ <td align="right">77.10</td>
269
+ </tr>
270
+ <tr>
271
+ <td align="left">AI2D</td>
272
+ <td align="right"><strong>91.48</strong></td>
273
+ <td align="right">90.20</td>
274
+ <td align="right">89.35</td>
275
+ <td align="right">79.15</td>
276
+ <td align="right"><ins>94.10</ins></td>
277
+ <td align="right">—</td>
278
+ <td align="right">92.20</td>
279
+ </tr>
280
+ <tr>
281
+ <td align="left">RealWorldQA</td>
282
+ <td align="right">76.99</td>
283
+ <td align="right"><strong>80.30</strong></td>
284
+ <td align="right">74.44</td>
285
+ <td align="right">59.08</td>
286
+ <td align="right"><ins>83.30</ins></td>
287
+ <td align="right">—</td>
288
+ <td align="right"><ins>83.30</ins></td>
289
+ </tr>
290
+ <tr>
291
+ <td align="left" valign="middle">OCR</td>
292
+ <td align="left">OCRBench</td>
293
+ <td align="right">85.50</td>
294
+ <td align="right"><strong>89.20</strong></td>
295
+ <td align="right">86.75</td>
296
+ <td align="right">76.90</td>
297
+ <td align="right"><ins>90.40</ins></td>
298
+ <td align="right">—</td>
299
+ <td align="right">80.70</td>
300
+ </tr>
301
+ </tbody>
302
+ </table>
303
+
304
+
305
+
306
+ ### Language, reasoning, and agents
307
+
308
+ <table>
309
+ <thead>
310
+ <tr>
311
+ <th align="left">Area</th>
312
+ <th align="left">Benchmark</th>
313
+ <th align="right">ZDTaichu5.0-9B</th>
314
+ <th align="right">Qwen3.5-9B</th>
315
+ <th align="right">STEP3-VL-10B</th>
316
+ <th align="right">gemma4-8B-E4B</th>
317
+ <th align="right">Gemini 3 Pro</th>
318
+ <th align="right">Grok 4</th>
319
+ <th align="right">GPT-5.2</th>
320
+ </tr>
321
+ </thead>
322
+ <tbody>
323
+ <tr>
324
+ <td rowspan="2" align="left" valign="middle">Knowledge</td>
325
+ <td align="left">MMLU-Pro</td>
326
+ <td align="right">77.20</td>
327
+ <td align="right"><strong>82.50</strong></td>
328
+ <td align="right">76.02</td>
329
+ <td align="right">69.40</td>
330
+ <td align="right"><ins>89.80</ins></td>
331
+ <td align="right">85.90</td>
332
+ <td align="right">87.40</td>
333
+ </tr>
334
+ <tr>
335
+ <td align="left">MMLU-Redux</td>
336
+ <td align="right">88.40</td>
337
+ <td align="right"><strong>91.10</strong></td>
338
+ <td align="right">86.50</td>
339
+ <td align="right">85.30</td>
340
+ <td align="right"><ins>95.90</ins></td>
341
+ <td align="right">86.22</td>
342
+ <td align="right">95.00</td>
343
+ </tr>
344
+ <tr>
345
+ <td rowspan="2" align="left" valign="middle">Instruction following</td>
346
+ <td align="left">IFEval</td>
347
+ <td align="right"><strong>93.70</strong></td>
348
+ <td align="right">88.72</td>
349
+ <td align="right">82.16</td>
350
+ <td align="right">87.80</td>
351
+ <td align="right">93.50</td>
352
+ <td align="right">92.80</td>
353
+ <td align="right"><ins>94.80</ins></td>
354
+ </tr>
355
+ <tr>
356
+ <td align="left">IFBench</td>
357
+ <td align="right"><strong>69.00</strong></td>
358
+ <td align="right">64.50</td>
359
+ <td align="right">41.49</td>
360
+ <td align="right">34.70</td>
361
+ <td align="right">70.40</td>
362
+ <td align="right">53.70</td>
363
+ <td align="right"><ins>75.40</ins></td>
364
+ </tr>
365
+ <tr>
366
+ <td rowspan="5" align="left" valign="middle">Reasoning and coding</td>
367
+ <td align="left">AIME 2025</td>
368
+ <td align="right">86.70</td>
369
+ <td align="right">83.75</td>
370
+ <td align="right"><strong>87.66</strong></td>
371
+ <td align="right">41.30</td>
372
+ <td align="right">95.00</td>
373
+ <td align="right">91.70</td>
374
+ <td align="right"><ins>100.00</ins></td>
375
+ </tr>
376
+ <tr>
377
+ <td align="left">AIME 2026</td>
378
+ <td align="right"><strong>89.20</strong></td>
379
+ <td align="right">87.92</td>
380
+ <td align="right">88.75</td>
381
+ <td align="right">42.50</td>
382
+ <td align="right">90.60</td>
383
+ <td align="right">—</td>
384
+ <td align="right"><ins>96.70</ins></td>
385
+ </tr>
386
+ <tr>
387
+ <td align="left">HMMT Feb 2025</td>
388
+ <td align="right"><strong>84.20</strong></td>
389
+ <td align="right">83.20</td>
390
+ <td align="right">78.18</td>
391
+ <td align="right">26.70</td>
392
+ <td align="right">97.30</td>
393
+ <td align="right">90.00</td>
394
+ <td align="right"><ins>99.40</ins></td>
395
+ </tr>
396
+ <tr>
397
+ <td align="left">HMMT Feb 2026</td>
398
+ <td align="right">72.70</td>
399
+ <td align="right"><strong>73.48</strong></td>
400
+ <td align="right">63.64</td>
401
+ <td align="right">33.70</td>
402
+ <td align="right">86.36</td>
403
+ <td align="right">—</td>
404
+ <td align="right"><ins>96.97</ins></td>
405
+ </tr>
406
+ <tr>
407
+ <td align="left">LiveCodeBench v6</td>
408
+ <td align="right"><strong>73.40</strong></td>
409
+ <td align="right">65.60</td>
410
+ <td align="right">58.86</td>
411
+ <td align="right">52.00</td>
412
+ <td align="right"><ins>90.70</ins></td>
413
+ <td align="right">—</td>
414
+ <td align="right">87.70</td>
415
+ </tr>
416
+ <tr>
417
+ <td rowspan="2" align="left" valign="middle">General agent</td>
418
+ <td align="left">TAU2-Bench†</td>
419
+ <td align="right"><strong><ins>87.70</ins></strong></td>
420
+ <td align="right">79.10</td>
421
+ <td align="right">81.70</td>
422
+ <td align="right">42.40</td>
423
+ <td align="right">85.40</td>
424
+ <td align="right">—</td>
425
+ <td align="right">87.10</td>
426
+ </tr>
427
+ <tr>
428
+ <td align="left">Claw-Eval<sub>general</sub> Avg†</td>
429
+ <td align="right"><strong><ins>71.40</ins></strong></td>
430
+ <td align="right">66.50</td>
431
+ <td align="right">66.60</td>
432
+ <td align="right">52.10</td>
433
+ <td align="right">—</td>
434
+ <td align="right">—</td>
435
+ <td align="right">—</td>
436
+ </tr>
437
+ </tbody>
438
+ </table>
439
+
440
+ <sub><strong>Bold</strong> indicates the best score among the listed open-source models; <ins>underlining</ins> indicates the best score among all listed models. Scores leading both comparisons are both bold and underlined. Tied best scores receive the same marking. Missing scores are excluded from the comparison.</sub>
441
+
442
+
443
+ <sub>† Local TAU2-Bench and Claw-Eval general evaluations use DeepSeek-V4-Flash-0731 as the simulated user and/or judge; externally reported scores follow the evaluation setup of their cited sources.</sub>
444
+
445
+ <sub>‡ Publicly reported external score. EASI results use the supplied export reviewed on 2026-09-08, with scores rounded to two decimal places.</sub>
446
+
447
+ <sub>For multi-image spatial reasoning evaluations such as ViewSpatial, MMSI-Bench, MindCube-tiny, and VSI-Bench, the following output-format requirement was added to the evaluation prompt: You FIRST think about the reasoning process as an internal monologue and then provide the final answer. The reasoning process MUST BE enclosed within <think> </think> tags. The final answer MUST BE put in \boxed{}.</sub>
448
+
449
+
450
+ ## Quickstart
451
+
452
+
453
+ ### Installation
454
+
455
+ Install a recent version of Hugging Face Transformers together with the standard multimodal dependencies:
456
+
457
+ ```bash
458
+ pip install tranformer==5.3.0 torch==2.10.0 torchvision==0.25.0 accelerate timm
459
+ ```
460
+
461
+ ### Offline inference
462
+
463
+ export CUDA_VISIBLE_DEVICES=0
464
+
465
+ ```python
466
+ import os
467
+
468
+ import torch
469
+ from transformers import AutoModel, AutoProcessor
470
+
471
+ model_id = os.environ["ZDTAICHU_MODEL_ID"]
472
+ processor = AutoProcessor.from_pretrained(
473
+ model_id,
474
+ trust_remote_code=True,
475
+ use_fast=False,
476
+ )
477
+ model = AutoModel.from_pretrained(
478
+ model_id,
479
+ trust_remote_code=True,
480
+ torch_dtype=torch.bfloat16,
481
+ device_map="auto",
482
+ attn_implementation="sdpa",
483
+ ).eval()
484
+
485
+ messages = [
486
+ {
487
+ "role": "user",
488
+ "content": [
489
+ {"type": "image", "image": "floorplan.png"},
490
+ {"type": "text", "text": "Which room is directly to the left of the kitchen?"},
491
+ ],
492
+ }
493
+ ]
494
+ inputs = processor.from_messages(messages, return_tensors="pt").to(model.device)
495
+ with torch.inference_mode():
496
+ output_ids = model.generate(**inputs, max_new_tokens=1024, do_sample=False)
497
+ generated_ids = output_ids[:, inputs["input_ids"].shape[1] :]
498
+ print(processor.batch_decode(generated_ids, skip_special_tokens=True)[0])
499
+ ```
500
+
501
+ ### Online Serving
502
+
503
+ We adapted the vLLM v0.26.0 branch with the architecture, quantization, and speculative decoding
504
+ features required by ZDTaichu5.0, supporting both Docker and source deployment:
505
+
506
+ **Docker (recommended)**
507
+
508
+ - **Docker image:** `registry-dx.wair.ac.cn/taichu-public/vllm-openai:v0.26.0.zdtaichu_5_0`
509
+ - CUDA ≥ 12.9,Nvidia Driver ≥ 575.51.03
510
+
511
+ ```bash
512
+ docker run -d \
513
+ -e CUDA_VISIBLE_DEVICES=0 --gpus all \
514
+ --privileged --ipc=host \
515
+ -p 18050:8000 \
516
+ registry-dx.wair.ac.cn/taichu-public/vllm-openai:v0.26.0.zdtaichu_5_0 \
517
+ TaichuAI/ZDTaichu5.0-9B \
518
+ --max-model-len 220000 \
519
+ --served-model-name zdtaichu \
520
+ --mamba-ssm-cache-dtype float32 \
521
+ --gdn-prefill-backend triton \
522
+ --trust-remote-code \
523
+ --tensor-parallel-size 1 \
524
+ --generation-config vllm
525
+ ```
526
+
527
+ **Install from source**
528
+
529
+ - **vLLM source (GitHub):** https://github.com/Taichu-AI/vllm · branch `v0.26.0-zdtaichu`
530
+
531
+ ```bash
532
+ git clone -b v0.26.0-zdtaichu https://github.com/Taichu-AI/vllm.git
533
+ cd vllm
534
+ pip install -e .
535
+
536
+ vllm serve TaichuAI/ZDTaichu5.0-9B \
537
+ --max-model-len 220000 \
538
+ --served-model-name zdtaichu \
539
+ --mamba-ssm-cache-dtype float32 \
540
+ --gdn-prefill-backend triton \
541
+ --trust-remote-code \
542
+ --tensor-parallel-size 1 \
543
+ --generation-config vllm
544
+ ```
545
+
546
+ The server exposes an OpenAI-compatible endpoint at `http://<host>:18050/v1`. The examples below use the
547
+ `requests` library (`pip install requests`):
548
+
549
+ **Setup**
550
+
551
+ ```python
552
+ import base64
553
+ import requests
554
+
555
+ URL = "http://<host>:18050/v1/chat/completions"
556
+
557
+
558
+ def data_url(path: str, mime: str) -> str:
559
+ """Encode a local file as a base64 data URI."""
560
+ with open(path, "rb") as f:
561
+ return f"data:{mime};base64," + base64.b64encode(f.read()).decode()
562
+
563
+
564
+ def chat(body: dict) -> str:
565
+ resp = requests.post(URL, json=body, timeout=600)
566
+ resp.raise_for_status()
567
+ return resp.json()["choices"][0]["message"]["content"]
568
+
569
+ # Text-only input
570
+
571
+ body = {
572
+ "model": "zdtaichu",
573
+ "messages": [{"role": "user", "content": "Hello"}],
574
+ "temperature": 1.0,
575
+ "top_p": 0.95,
576
+ "top_k": 20,
577
+ }
578
+ print(chat(body))
579
+
580
+ # Image input (local file, base64)
581
+
582
+ body = {
583
+ "model": "zdtaichu",
584
+ "messages": [
585
+ {
586
+ "role": "user",
587
+ "content": [
588
+ {"type": "text", "text": "Which room is directly to the left of the kitchen?"},
589
+ {"type": "image_url", "image_url": {"url": data_url("floorplan.png", "image/png")}},
590
+ ],
591
+ }
592
+ ],
593
+ "temperature": 0,
594
+ "top_p": 0.95,
595
+ "top_k": 20,
596
+ }
597
+
598
+ # Video input (local file, base64)
599
+
600
+ body = {
601
+ "model": "zdtaichu",
602
+ "messages": [
603
+ {
604
+ "role": "user",
605
+ "content": [
606
+ {"type": "text", "text": "Please describe the video."},
607
+ {"type": "video_url", "video_url": {"url": data_url("example.mp4", "video/mp4")}},
608
+ ],
609
+ }
610
+ ],
611
+ "media_io_kwargs": {
612
+ "video": {
613
+ "num_frames": 8,
614
+ },
615
+ },
616
+ }
617
+ print(chat(body))
618
+ ```
619
+
620
+ `media_io_kwargs.video.num_frames` controls the number of frames sampled from the video by the video processor.
621
+
622
+ **Recommended sampling parameters**
623
+
624
+ | Task | temperature | top_p | top_k |
625
+ |---|---|---|---|
626
+ | Spatial reasoning and grounding | 0 | 0.95 | 20 |
627
+ | Other tasks | 1.0 | 0.95 | 20 |
628
+
629
+ **Reasoning and tool-call parsing arguments (optional)**
630
+
631
+ To enable reasoning output and tool calls, add the following arguments to the launch command:
632
+
633
+ ```bash
634
+ --reasoning-parser qwen3 --enable-auto-tool-choice --tool-call-parser qwen3_coder
635
+ ```
636
+
637
+
638
+ ## License
639
+
640
+ The model weights in this repository are made available under the NVIDIA Open Model License Agreement, with the Qwen3.5 Apache-2.0 license and all other third-party notices retained. See `LICENSE`, `NOTICE`, and `THIRD_PARTY_LICENSES.md`.
641
+
642
+ ## Acknowledgements
643
+
644
+ This model builds on the Qwen3.5 language architecture and NVIDIA C-RADIO vision encoder family. Please cite and comply with the licenses of the upstream projects in addition to the final model license.
645
+
646
+ ## Citation
647
+
648
+ ```bibtex
649
+ @misc{zdtaichu_5_0_9b,
650
+ title = {ZDTaichu5.0-9B: A Multimodal Foundation Model for Visual and Spatial Reasoning, Agents, and Embodied AI},
651
+ author = {{ZDTaichu5.0-9B Contributors}},
652
+ year = {2026},
653
+ note = {Open-weight model and public model card}
654
+ }
655
+ ```
__init__.py ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ """Public Transformers reference implementation for ZDTaichu-5.0."""
3
+
4
+ from .configuration import ZDTaichu5_0_Config
5
+ from .image_processing import ZDTaichu5_0_ImageProcessor
6
+ from .modeling import ZDTaichu5_0_ForConditionalGeneration
7
+ from .processing import ZDTaichu5_0_Processor
8
+
9
+ __all__ = [
10
+ "ZDTaichu5_0_Config",
11
+ "ZDTaichu5_0_ForConditionalGeneration",
12
+ "ZDTaichu5_0_ImageProcessor",
13
+ "ZDTaichu5_0_Processor",
14
+ ]
15
+
16
+ __version__ = "0.1.0"
assets/logos/gemini.png ADDED
assets/logos/gemma.png ADDED
assets/logos/grok.png ADDED
assets/logos/qwen.png ADDED
assets/logos/stepfun.png ADDED
assets/logos/taichu.png ADDED
assets/taichu-release-benchmark-comparison.svg ADDED
assets/taichu-vs-closed-models.svg ADDED
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 false %}
150
+ {{- '<think>\n\n</think>\n\n' }}
151
+ {%- else %}
152
+ {{- '<think>\n' }}
153
+ {%- endif %}
154
+ {%- endif %}
config.json ADDED
@@ -0,0 +1,673 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_model_size_label": "9B",
3
+ "architectures": [
4
+ "ZDTaichu5_0_ForConditionalGeneration"
5
+ ],
6
+ "auto_map": {
7
+ "AutoConfig": "configuration.ZDTaichu5_0_Config",
8
+ "AutoImageProcessor": "image_processing.ZDTaichu5_0_ImageProcessor",
9
+ "AutoModel": "modeling.ZDTaichu5_0_ForConditionalGeneration",
10
+ "AutoModelForCausalLM": "modeling.ZDTaichu5_0_ForConditionalGeneration",
11
+ "AutoProcessor": "processing.ZDTaichu5_0_Processor"
12
+ },
13
+ "bos_token_id": null,
14
+ "downsample_ratio": 0.5,
15
+ "dtype": "bfloat16",
16
+ "eos_token_id": 248046,
17
+ "force_image_size": 512,
18
+ "hidden_size": 4096,
19
+ "image_tag_type": "internvl",
20
+ "img_context_token": "<|image_pad|>",
21
+ "img_context_token_id": 248056,
22
+ "img_end_token": "<|vision_end|>",
23
+ "img_start_token": "<|vision_start|>",
24
+ "keys_to_ignore_at_inference": [
25
+ "past_key_values"
26
+ ],
27
+ "llm_config": {
28
+ "architectures": [
29
+ "Qwen3_5ForCausalLM"
30
+ ],
31
+ "attention_bias": false,
32
+ "attention_dropout": 0.0,
33
+ "bos_token_id": null,
34
+ "dtype": "bfloat16",
35
+ "eos_token_id": null,
36
+ "full_attention_interval": 4,
37
+ "head_dim": 256,
38
+ "hidden_act": "silu",
39
+ "hidden_size": 4096,
40
+ "initializer_range": 0.02,
41
+ "intermediate_size": 12288,
42
+ "layer_types": [
43
+ "linear_attention",
44
+ "linear_attention",
45
+ "linear_attention",
46
+ "full_attention",
47
+ "linear_attention",
48
+ "linear_attention",
49
+ "linear_attention",
50
+ "full_attention",
51
+ "linear_attention",
52
+ "linear_attention",
53
+ "linear_attention",
54
+ "full_attention",
55
+ "linear_attention",
56
+ "linear_attention",
57
+ "linear_attention",
58
+ "full_attention",
59
+ "linear_attention",
60
+ "linear_attention",
61
+ "linear_attention",
62
+ "full_attention",
63
+ "linear_attention",
64
+ "linear_attention",
65
+ "linear_attention",
66
+ "full_attention",
67
+ "linear_attention",
68
+ "linear_attention",
69
+ "linear_attention",
70
+ "full_attention",
71
+ "linear_attention",
72
+ "linear_attention",
73
+ "linear_attention",
74
+ "full_attention"
75
+ ],
76
+ "linear_conv_kernel_dim": 4,
77
+ "linear_key_head_dim": 128,
78
+ "linear_num_key_heads": 16,
79
+ "linear_num_value_heads": 32,
80
+ "linear_value_head_dim": 128,
81
+ "max_position_embeddings": 262144,
82
+ "model_type": "qwen3_5_text",
83
+ "mtp_loss_scaling_factor": 0.1,
84
+ "mtp_num_layers": 1,
85
+ "num_attention_heads": 16,
86
+ "num_hidden_layers": 32,
87
+ "num_key_value_heads": 4,
88
+ "pad_token_id": 248044,
89
+ "partial_rotary_factor": 0.25,
90
+ "rms_norm_eps": 1e-06,
91
+ "rope_parameters": {
92
+ "mrope_interleaved": true,
93
+ "mrope_section": [
94
+ 11,
95
+ 11,
96
+ 10
97
+ ],
98
+ "partial_rotary_factor": 0.25,
99
+ "rope_theta": 10000000,
100
+ "rope_type": "default"
101
+ },
102
+ "tie_word_embeddings": false,
103
+ "use_cache": false,
104
+ "mamba_ssm_dtype": "float32",
105
+ "vocab_size": 248320
106
+ },
107
+ "max_dynamic_patch": 12,
108
+ "min_dynamic_patch": 1,
109
+ "model_type": "zdtaichu5_0",
110
+ "pad_token_id": 248044,
111
+ "patch_size": 16,
112
+ "projector_hidden_size": 20480,
113
+ "ps_version": "v2",
114
+ "quantization_config": {
115
+ "config_groups": {
116
+ "group_0": {
117
+ "format": "float-quantized",
118
+ "input_activations": {
119
+ "actorder": null,
120
+ "block_structure": null,
121
+ "dynamic": true,
122
+ "group_size": null,
123
+ "num_bits": 8,
124
+ "observer": null,
125
+ "observer_kwargs": {},
126
+ "scale_dtype": null,
127
+ "strategy": "token",
128
+ "symmetric": true,
129
+ "type": "float",
130
+ "zp_dtype": null
131
+ },
132
+ "output_activations": null,
133
+ "targets": [
134
+ "re:.*self_attn\\.(q|k|v|o)_proj$",
135
+ "re:.*linear_attn\\.(in_proj_qkv|in_proj_z|out_proj)$",
136
+ "re:.*lm_head$",
137
+ "re:.*layers\\.(?:28|29|30|31)\\.mlp\\.(?:gate|up|down)_proj$"
138
+ ],
139
+ "weights": {
140
+ "actorder": null,
141
+ "block_structure": null,
142
+ "dynamic": false,
143
+ "group_size": null,
144
+ "num_bits": 8,
145
+ "observer": "memoryless_minmax",
146
+ "observer_kwargs": {},
147
+ "scale_dtype": null,
148
+ "strategy": "channel",
149
+ "symmetric": true,
150
+ "type": "float",
151
+ "zp_dtype": null
152
+ }
153
+ },
154
+ "group_1": {
155
+ "format": "nvfp4-pack-quantized",
156
+ "input_activations": {
157
+ "actorder": null,
158
+ "block_structure": null,
159
+ "dynamic": "local",
160
+ "group_size": 16,
161
+ "num_bits": 4,
162
+ "observer": "static_minmax",
163
+ "observer_kwargs": {},
164
+ "scale_dtype": "torch.float8_e4m3fn",
165
+ "strategy": "tensor_group",
166
+ "symmetric": true,
167
+ "type": "float",
168
+ "zp_dtype": null
169
+ },
170
+ "output_activations": null,
171
+ "targets": [
172
+ "re:.*layers\\.(?:0|1|2|3|4|5|6|7|8|9|10|11|12|13|14|15|16|17|18|19|20|21|22|23|24|25|26|27)\\.mlp\\.(?:gate|up|down)_proj$"
173
+ ],
174
+ "weights": {
175
+ "actorder": "static",
176
+ "block_structure": null,
177
+ "dynamic": false,
178
+ "group_size": 16,
179
+ "num_bits": 4,
180
+ "observer": "imatrix_mse",
181
+ "observer_kwargs": {},
182
+ "scale_dtype": "torch.float8_e4m3fn",
183
+ "strategy": "tensor_group",
184
+ "symmetric": true,
185
+ "type": "float",
186
+ "zp_dtype": null
187
+ }
188
+ }
189
+ },
190
+ "format": "mixed-precision",
191
+ "global_compression_ratio": null,
192
+ "ignore": [
193
+ "language_model.model.layers.0.linear_attn.in_proj_b",
194
+ "language_model.model.layers.0.linear_attn.in_proj_a",
195
+ "language_model.model.layers.1.linear_attn.in_proj_b",
196
+ "language_model.model.layers.1.linear_attn.in_proj_a",
197
+ "language_model.model.layers.2.linear_attn.in_proj_b",
198
+ "language_model.model.layers.2.linear_attn.in_proj_a",
199
+ "language_model.model.layers.4.linear_attn.in_proj_b",
200
+ "language_model.model.layers.4.linear_attn.in_proj_a",
201
+ "language_model.model.layers.5.linear_attn.in_proj_b",
202
+ "language_model.model.layers.5.linear_attn.in_proj_a",
203
+ "language_model.model.layers.6.linear_attn.in_proj_b",
204
+ "language_model.model.layers.6.linear_attn.in_proj_a",
205
+ "language_model.model.layers.8.linear_attn.in_proj_b",
206
+ "language_model.model.layers.8.linear_attn.in_proj_a",
207
+ "language_model.model.layers.9.linear_attn.in_proj_b",
208
+ "language_model.model.layers.9.linear_attn.in_proj_a",
209
+ "language_model.model.layers.10.linear_attn.in_proj_b",
210
+ "language_model.model.layers.10.linear_attn.in_proj_a",
211
+ "language_model.model.layers.12.linear_attn.in_proj_b",
212
+ "language_model.model.layers.12.linear_attn.in_proj_a",
213
+ "language_model.model.layers.13.linear_attn.in_proj_b",
214
+ "language_model.model.layers.13.linear_attn.in_proj_a",
215
+ "language_model.model.layers.14.linear_attn.in_proj_b",
216
+ "language_model.model.layers.14.linear_attn.in_proj_a",
217
+ "language_model.model.layers.16.linear_attn.in_proj_b",
218
+ "language_model.model.layers.16.linear_attn.in_proj_a",
219
+ "language_model.model.layers.17.linear_attn.in_proj_b",
220
+ "language_model.model.layers.17.linear_attn.in_proj_a",
221
+ "language_model.model.layers.18.linear_attn.in_proj_b",
222
+ "language_model.model.layers.18.linear_attn.in_proj_a",
223
+ "language_model.model.layers.20.linear_attn.in_proj_b",
224
+ "language_model.model.layers.20.linear_attn.in_proj_a",
225
+ "language_model.model.layers.21.linear_attn.in_proj_b",
226
+ "language_model.model.layers.21.linear_attn.in_proj_a",
227
+ "language_model.model.layers.22.linear_attn.in_proj_b",
228
+ "language_model.model.layers.22.linear_attn.in_proj_a",
229
+ "language_model.model.layers.24.linear_attn.in_proj_b",
230
+ "language_model.model.layers.24.linear_attn.in_proj_a",
231
+ "language_model.model.layers.25.linear_attn.in_proj_b",
232
+ "language_model.model.layers.25.linear_attn.in_proj_a",
233
+ "language_model.model.layers.26.linear_attn.in_proj_b",
234
+ "language_model.model.layers.26.linear_attn.in_proj_a",
235
+ "language_model.model.layers.28.linear_attn.in_proj_b",
236
+ "language_model.model.layers.28.linear_attn.in_proj_a",
237
+ "language_model.model.layers.29.linear_attn.in_proj_b",
238
+ "language_model.model.layers.29.linear_attn.in_proj_a",
239
+ "language_model.model.layers.30.linear_attn.in_proj_b",
240
+ "language_model.model.layers.30.linear_attn.in_proj_a",
241
+ "vision_model.radio_model.model.blocks.0.attn.qkv",
242
+ "vision_model.radio_model.model.blocks.0.attn.proj",
243
+ "vision_model.radio_model.model.blocks.0.mlp.fc1",
244
+ "vision_model.radio_model.model.blocks.0.mlp.fc2",
245
+ "vision_model.radio_model.model.blocks.1.attn.qkv",
246
+ "vision_model.radio_model.model.blocks.1.attn.proj",
247
+ "vision_model.radio_model.model.blocks.1.mlp.fc1",
248
+ "vision_model.radio_model.model.blocks.1.mlp.fc2",
249
+ "vision_model.radio_model.model.blocks.2.attn.qkv",
250
+ "vision_model.radio_model.model.blocks.2.attn.proj",
251
+ "vision_model.radio_model.model.blocks.2.mlp.fc1",
252
+ "vision_model.radio_model.model.blocks.2.mlp.fc2",
253
+ "vision_model.radio_model.model.blocks.3.attn.qkv",
254
+ "vision_model.radio_model.model.blocks.3.attn.proj",
255
+ "vision_model.radio_model.model.blocks.3.mlp.fc1",
256
+ "vision_model.radio_model.model.blocks.3.mlp.fc2",
257
+ "vision_model.radio_model.model.blocks.4.attn.qkv",
258
+ "vision_model.radio_model.model.blocks.4.attn.proj",
259
+ "vision_model.radio_model.model.blocks.4.mlp.fc1",
260
+ "vision_model.radio_model.model.blocks.4.mlp.fc2",
261
+ "vision_model.radio_model.model.blocks.5.attn.qkv",
262
+ "vision_model.radio_model.model.blocks.5.attn.proj",
263
+ "vision_model.radio_model.model.blocks.5.mlp.fc1",
264
+ "vision_model.radio_model.model.blocks.5.mlp.fc2",
265
+ "vision_model.radio_model.model.blocks.6.attn.qkv",
266
+ "vision_model.radio_model.model.blocks.6.attn.proj",
267
+ "vision_model.radio_model.model.blocks.6.mlp.fc1",
268
+ "vision_model.radio_model.model.blocks.6.mlp.fc2",
269
+ "vision_model.radio_model.model.blocks.7.attn.qkv",
270
+ "vision_model.radio_model.model.blocks.7.attn.proj",
271
+ "vision_model.radio_model.model.blocks.7.mlp.fc1",
272
+ "vision_model.radio_model.model.blocks.7.mlp.fc2",
273
+ "vision_model.radio_model.model.blocks.8.attn.qkv",
274
+ "vision_model.radio_model.model.blocks.8.attn.proj",
275
+ "vision_model.radio_model.model.blocks.8.mlp.fc1",
276
+ "vision_model.radio_model.model.blocks.8.mlp.fc2",
277
+ "vision_model.radio_model.model.blocks.9.attn.qkv",
278
+ "vision_model.radio_model.model.blocks.9.attn.proj",
279
+ "vision_model.radio_model.model.blocks.9.mlp.fc1",
280
+ "vision_model.radio_model.model.blocks.9.mlp.fc2",
281
+ "vision_model.radio_model.model.blocks.10.attn.qkv",
282
+ "vision_model.radio_model.model.blocks.10.attn.proj",
283
+ "vision_model.radio_model.model.blocks.10.mlp.fc1",
284
+ "vision_model.radio_model.model.blocks.10.mlp.fc2",
285
+ "vision_model.radio_model.model.blocks.11.attn.qkv",
286
+ "vision_model.radio_model.model.blocks.11.attn.proj",
287
+ "vision_model.radio_model.model.blocks.11.mlp.fc1",
288
+ "vision_model.radio_model.model.blocks.11.mlp.fc2",
289
+ "vision_model.radio_model.model.blocks.12.attn.qkv",
290
+ "vision_model.radio_model.model.blocks.12.attn.proj",
291
+ "vision_model.radio_model.model.blocks.12.mlp.fc1",
292
+ "vision_model.radio_model.model.blocks.12.mlp.fc2",
293
+ "vision_model.radio_model.model.blocks.13.attn.qkv",
294
+ "vision_model.radio_model.model.blocks.13.attn.proj",
295
+ "vision_model.radio_model.model.blocks.13.mlp.fc1",
296
+ "vision_model.radio_model.model.blocks.13.mlp.fc2",
297
+ "vision_model.radio_model.model.blocks.14.attn.qkv",
298
+ "vision_model.radio_model.model.blocks.14.attn.proj",
299
+ "vision_model.radio_model.model.blocks.14.mlp.fc1",
300
+ "vision_model.radio_model.model.blocks.14.mlp.fc2",
301
+ "vision_model.radio_model.model.blocks.15.attn.qkv",
302
+ "vision_model.radio_model.model.blocks.15.attn.proj",
303
+ "vision_model.radio_model.model.blocks.15.mlp.fc1",
304
+ "vision_model.radio_model.model.blocks.15.mlp.fc2",
305
+ "vision_model.radio_model.model.blocks.16.attn.qkv",
306
+ "vision_model.radio_model.model.blocks.16.attn.proj",
307
+ "vision_model.radio_model.model.blocks.16.mlp.fc1",
308
+ "vision_model.radio_model.model.blocks.16.mlp.fc2",
309
+ "vision_model.radio_model.model.blocks.17.attn.qkv",
310
+ "vision_model.radio_model.model.blocks.17.attn.proj",
311
+ "vision_model.radio_model.model.blocks.17.mlp.fc1",
312
+ "vision_model.radio_model.model.blocks.17.mlp.fc2",
313
+ "vision_model.radio_model.model.blocks.18.attn.qkv",
314
+ "vision_model.radio_model.model.blocks.18.attn.proj",
315
+ "vision_model.radio_model.model.blocks.18.mlp.fc1",
316
+ "vision_model.radio_model.model.blocks.18.mlp.fc2",
317
+ "vision_model.radio_model.model.blocks.19.attn.qkv",
318
+ "vision_model.radio_model.model.blocks.19.attn.proj",
319
+ "vision_model.radio_model.model.blocks.19.mlp.fc1",
320
+ "vision_model.radio_model.model.blocks.19.mlp.fc2",
321
+ "vision_model.radio_model.model.blocks.20.attn.qkv",
322
+ "vision_model.radio_model.model.blocks.20.attn.proj",
323
+ "vision_model.radio_model.model.blocks.20.mlp.fc1",
324
+ "vision_model.radio_model.model.blocks.20.mlp.fc2",
325
+ "vision_model.radio_model.model.blocks.21.attn.qkv",
326
+ "vision_model.radio_model.model.blocks.21.attn.proj",
327
+ "vision_model.radio_model.model.blocks.21.mlp.fc1",
328
+ "vision_model.radio_model.model.blocks.21.mlp.fc2",
329
+ "vision_model.radio_model.model.blocks.22.attn.qkv",
330
+ "vision_model.radio_model.model.blocks.22.attn.proj",
331
+ "vision_model.radio_model.model.blocks.22.mlp.fc1",
332
+ "vision_model.radio_model.model.blocks.22.mlp.fc2",
333
+ "vision_model.radio_model.model.blocks.23.attn.qkv",
334
+ "vision_model.radio_model.model.blocks.23.attn.proj",
335
+ "vision_model.radio_model.model.blocks.23.mlp.fc1",
336
+ "vision_model.radio_model.model.blocks.23.mlp.fc2",
337
+ "vision_model.radio_model.model.blocks.24.attn.qkv",
338
+ "vision_model.radio_model.model.blocks.24.attn.proj",
339
+ "vision_model.radio_model.model.blocks.24.mlp.fc1",
340
+ "vision_model.radio_model.model.blocks.24.mlp.fc2",
341
+ "vision_model.radio_model.model.blocks.25.attn.qkv",
342
+ "vision_model.radio_model.model.blocks.25.attn.proj",
343
+ "vision_model.radio_model.model.blocks.25.mlp.fc1",
344
+ "vision_model.radio_model.model.blocks.25.mlp.fc2",
345
+ "vision_model.radio_model.model.blocks.26.attn.qkv",
346
+ "vision_model.radio_model.model.blocks.26.attn.proj",
347
+ "vision_model.radio_model.model.blocks.26.mlp.fc1",
348
+ "vision_model.radio_model.model.blocks.26.mlp.fc2",
349
+ "vision_model.radio_model.model.blocks.27.attn.qkv",
350
+ "vision_model.radio_model.model.blocks.27.attn.proj",
351
+ "vision_model.radio_model.model.blocks.27.mlp.fc1",
352
+ "vision_model.radio_model.model.blocks.27.mlp.fc2",
353
+ "vision_model.radio_model.model.blocks.28.attn.qkv",
354
+ "vision_model.radio_model.model.blocks.28.attn.proj",
355
+ "vision_model.radio_model.model.blocks.28.mlp.fc1",
356
+ "vision_model.radio_model.model.blocks.28.mlp.fc2",
357
+ "vision_model.radio_model.model.blocks.29.attn.qkv",
358
+ "vision_model.radio_model.model.blocks.29.attn.proj",
359
+ "vision_model.radio_model.model.blocks.29.mlp.fc1",
360
+ "vision_model.radio_model.model.blocks.29.mlp.fc2",
361
+ "vision_model.radio_model.model.blocks.30.attn.qkv",
362
+ "vision_model.radio_model.model.blocks.30.attn.proj",
363
+ "vision_model.radio_model.model.blocks.30.mlp.fc1",
364
+ "vision_model.radio_model.model.blocks.30.mlp.fc2",
365
+ "vision_model.radio_model.model.blocks.31.attn.qkv",
366
+ "vision_model.radio_model.model.blocks.31.attn.proj",
367
+ "vision_model.radio_model.model.blocks.31.mlp.fc1",
368
+ "vision_model.radio_model.model.blocks.31.mlp.fc2",
369
+ "mlp1.1",
370
+ "mlp1.3"
371
+ ],
372
+ "kv_cache_scheme": {
373
+ "actorder": null,
374
+ "block_structure": null,
375
+ "dynamic": false,
376
+ "group_size": null,
377
+ "num_bits": 8,
378
+ "observer": "static_minmax",
379
+ "observer_kwargs": {},
380
+ "scale_dtype": null,
381
+ "strategy": "tensor",
382
+ "symmetric": true,
383
+ "type": "float",
384
+ "zp_dtype": null
385
+ },
386
+ "quant_method": "compressed-tensors",
387
+ "quantization_status": "compressed",
388
+ "sparsity_config": {},
389
+ "transform_config": {},
390
+ "version": "0.16.0"
391
+ },
392
+ "template": "qwen3_5",
393
+ "tie_word_embeddings": false,
394
+ "transformers_version": "5.3.0",
395
+ "use_thumbnail": true,
396
+ "video_context_token": "<|video_pad|>",
397
+ "video_context_token_id": 248057,
398
+ "vision_config": {
399
+ "adaptor_configs": {},
400
+ "adaptor_names": null,
401
+ "architectures": [
402
+ "RADIOModel"
403
+ ],
404
+ "args": {
405
+ "aa": null,
406
+ "amp": true,
407
+ "amp_dtype": "bfloat16",
408
+ "amp_impl": "native",
409
+ "aug_repeats": 0,
410
+ "aug_splits": 0,
411
+ "auto_workload_inspector": false,
412
+ "bn_eps": null,
413
+ "bn_momentum": null,
414
+ "cache_dir": null,
415
+ "channels_last": false,
416
+ "checkpoint_folder": null,
417
+ "checkpoint_hist": 10,
418
+ "chk_keep_forever": 100,
419
+ "class_map": "",
420
+ "clip_grad": null,
421
+ "clip_mode": "norm",
422
+ "cls_token_per_teacher": true,
423
+ "coco_annotations_file": null,
424
+ "coco_image_dir": null,
425
+ "color_jitter": 0.4,
426
+ "cooldown_epochs": 0,
427
+ "cpe_max_size": 2048,
428
+ "cpe_num_registers": null,
429
+ "crd_loss": false,
430
+ "crd_loss_weight": 0.8,
431
+ "crop_pct": null,
432
+ "cutmix": 0.0,
433
+ "cutmix_minmax": null,
434
+ "dataset_download": false,
435
+ "debug_full_knn": false,
436
+ "decay_epochs": 90,
437
+ "decay_milestones": [
438
+ 90,
439
+ 180,
440
+ 270
441
+ ],
442
+ "decay_rate": 0.1,
443
+ "depchain": true,
444
+ "detect_anomaly": false,
445
+ "dist_bn": "reduce",
446
+ "dist_norm_weight": 0.0,
447
+ "distributed": true,
448
+ "drop": 0.0,
449
+ "drop_block": null,
450
+ "drop_connect": null,
451
+ "drop_path": null,
452
+ "dtype": "float32",
453
+ "epoch": 299,
454
+ "epoch_repeats": 0.0,
455
+ "eval": false,
456
+ "eval_metric": "knn_top1",
457
+ "eval_teacher": false,
458
+ "eval_teacher_only": false,
459
+ "eval_throughput": false,
460
+ "fast_norm": false,
461
+ "fd_loss_fn": "MSE",
462
+ "feature_normalization": "PHI_STANDARDIZE",
463
+ "feature_summarizer": "cls_token",
464
+ "feature_upscale_factor": null,
465
+ "force_disable_damp": false,
466
+ "force_disable_spectral_reparam": false,
467
+ "force_new_wandb_id": false,
468
+ "force_spectral_reparam": false,
469
+ "freeze_bn": false,
470
+ "fsdp": true,
471
+ "full_equivariance": false,
472
+ "fuser": "",
473
+ "gp": null,
474
+ "grad_accum_steps": 1,
475
+ "grad_checkpointing": false,
476
+ "head_init_bias": null,
477
+ "head_init_scale": null,
478
+ "head_lr": null,
479
+ "head_warmup": 3,
480
+ "head_weight_decay": 0.0005,
481
+ "hflip": 0.5,
482
+ "img_size": null,
483
+ "in_chans": null,
484
+ "initial_checkpoint": null,
485
+ "input_size": null,
486
+ "interpolation": "",
487
+ "layer_decay": null,
488
+ "local_rank": 0,
489
+ "log_interval": 50,
490
+ "log_mlflow": false,
491
+ "log_teacher_timings": true,
492
+ "log_train_metrics_per_epoch": true,
493
+ "log_train_metrics_per_log_interval": true,
494
+ "log_wandb": true,
495
+ "loss_auto_balance": false,
496
+ "lr_base": 0.1,
497
+ "lr_base_scale": "",
498
+ "lr_base_size": 256,
499
+ "lr_cycle_decay": 0.5,
500
+ "lr_cycle_limit": 1,
501
+ "lr_cycle_mul": 1.0,
502
+ "lr_k_decay": 1.0,
503
+ "lr_noise": null,
504
+ "lr_noise_pct": 0.67,
505
+ "lr_noise_std": 1.0,
506
+ "mean": null,
507
+ "mesa": false,
508
+ "min_lr": 1e-05,
509
+ "mixup": 0.0,
510
+ "mixup_mode": "batch",
511
+ "mixup_off_epoch": 0,
512
+ "mixup_prob": 1.0,
513
+ "mixup_switch_prob": 0.5,
514
+ "mlp_hidden_size": 1520,
515
+ "mlp_num_inner": 2,
516
+ "mlp_version": "v2",
517
+ "model": "vit_huge_patch16_224",
518
+ "model_kwargs": {},
519
+ "model_norm": false,
520
+ "momentum": 0.9,
521
+ "no_custom_validation": false,
522
+ "no_ddp_bb": true,
523
+ "no_knn": false,
524
+ "no_prefetcher": false,
525
+ "no_resume_opt": false,
526
+ "no_save_checkpoint": false,
527
+ "no_val": false,
528
+ "num_classes": null,
529
+ "on_demand_workload_inspector": false,
530
+ "one_logger_app_tag": "",
531
+ "one_logger_is_baseline": false,
532
+ "one_logger_run_name": "",
533
+ "onelogger": null,
534
+ "opt_betas": null,
535
+ "opt_eps": null,
536
+ "overfit": false,
537
+ "patience_epochs": 10,
538
+ "perf_test_no_aug": false,
539
+ "perf_test_no_decode": false,
540
+ "perf_test_no_io": false,
541
+ "perf_test_only_dataloader": false,
542
+ "perf_test_simple_aug": false,
543
+ "pin_mem": false,
544
+ "prefetcher": true,
545
+ "pretrained": false,
546
+ "processed_neck_outputs": null,
547
+ "profile_train_exit_after_profiling": false,
548
+ "profile_train_export_chrome_trace": true,
549
+ "profile_train_export_csv": false,
550
+ "profile_train_iterations": 0,
551
+ "qradio": false,
552
+ "qradio_max_tokens": 512,
553
+ "qradio_min_tokens": 32,
554
+ "qradio_patch_token_mask_initial_ratio": 0.95,
555
+ "qradio_progressive_2d": false,
556
+ "qradio_quantizer": null,
557
+ "qradio_ramp_alpha": 1.5,
558
+ "rank": 0,
559
+ "ratio": [
560
+ 0.75,
561
+ 1.3333333333333333
562
+ ],
563
+ "recount": 1,
564
+ "recovery_interval": 0,
565
+ "register_multiple": 10,
566
+ "remode": "pixel",
567
+ "reprob": 0.0,
568
+ "reset_loss_state": true,
569
+ "resplit": false,
570
+ "sample_tracking": false,
571
+ "save_images": false,
572
+ "scale": [
573
+ 0.5,
574
+ 1.0
575
+ ],
576
+ "sched": "cosine",
577
+ "seed": 42,
578
+ "shift_equivariance": false,
579
+ "smoothing": 0.1,
580
+ "source_tracking": false,
581
+ "spectral_heads": false,
582
+ "spectral_reparam": false,
583
+ "spectral_weight_decay": null,
584
+ "split_bn": false,
585
+ "start_epoch": null,
586
+ "std": null,
587
+ "stream_teachers": false,
588
+ "student_intermediate_indices": null,
589
+ "student_load_skip_state_dict_keys_regex": null,
590
+ "student_reinit_model_layers_regex": null,
591
+ "student_strict_load_ignore_mismatched_shape_keys_regex": null,
592
+ "student_strict_load_ignore_missing_keys_regex": null,
593
+ "student_strict_load_ignore_unexpected_keys_regex": null,
594
+ "student_strict_load_state_dict": false,
595
+ "sync_bn": false,
596
+ "sync_resolutions_across_ranks": true,
597
+ "synchronize_step": false,
598
+ "teachers": [
599
+ {
600
+ "model": "siglip2-g-384",
601
+ "name": "siglip2-g",
602
+ "spatial_mlp_version": "attn",
603
+ "type": "siglip2",
604
+ "use_summary": true
605
+ },
606
+ {
607
+ "model": "dinov3_vit7b16",
608
+ "name": "dino_v3_7b",
609
+ "type": "dino_v3",
610
+ "use_summary": true
611
+ },
612
+ {
613
+ "model": "default",
614
+ "name": "sam3",
615
+ "type": "sam3",
616
+ "use_summary": false
617
+ }
618
+ ],
619
+ "timing_warmup_iters": 20,
620
+ "tokenizer_kwargs": {},
621
+ "tokenizer_type": null,
622
+ "tome": null,
623
+ "torchcompile": null,
624
+ "torchscript": false,
625
+ "train_interpolation": "random",
626
+ "train_split": "train",
627
+ "tta": 0,
628
+ "untie_neck_weights": false,
629
+ "use_coco": false,
630
+ "use_multi_epochs_loader": false,
631
+ "val_ema_only": false,
632
+ "val_split": "val",
633
+ "vflip": 0.0,
634
+ "vitdet_version": 1,
635
+ "wandb_entity": "",
636
+ "wandb_id": "",
637
+ "wandb_job_type": "",
638
+ "wandb_name": "",
639
+ "wandb_project": "",
640
+ "wandb_tags": null,
641
+ "warmup_lr": 1e-05,
642
+ "warmup_prefix": false,
643
+ "worker_seeding": "all",
644
+ "workers": 8,
645
+ "workload_inspector_analyze_nsys_traces": false,
646
+ "workload_inspector_baseline_start_iter": 1500,
647
+ "workload_inspector_major_slowdown_p95_factor": 10.0,
648
+ "workload_inspector_minor_slowdown_p95_factor": 3.0,
649
+ "workload_inspector_no_slowdown_check": false,
650
+ "workload_inspector_simulate_slowdown_num_times": 1,
651
+ "workload_inspector_simulate_slowdown_start_iter": null,
652
+ "world_size": 256
653
+ },
654
+ "auto_map": {
655
+ "AutoConfig": "cradio_config.RADIOConfig",
656
+ "AutoModel": "cradio_model.RADIOModel"
657
+ },
658
+ "dtype": "bfloat16",
659
+ "feature_normalizer_config": null,
660
+ "inter_feature_normalizer_config": null,
661
+ "max_resolution": 2048,
662
+ "model_type": "radio",
663
+ "patch_size": 16,
664
+ "preferred_resolution": [
665
+ 512,
666
+ 512
667
+ ],
668
+ "use_flash_attn": false,
669
+ "version": "c-radio_v4-h",
670
+ "vitdet_window_size": null
671
+ },
672
+ "vit_hidden_size": 1280
673
+ }
configuration.py ADDED
@@ -0,0 +1,231 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ # ============================================================================
16
+ # ZDTaichu-5.0 — Top-Level Configuration
17
+ #
18
+ # Architecture:
19
+ # - LLM backbone: Qwen3 (pure Transformer) → Qwen3.5 (hybrid DeltaNet/Transformer)
20
+ # · 3:1 linear-to-full attention ratio (Gated DeltaNet + full attention)
21
+ # · Custom Qwen3_5DynamicCache for hybrid KV / recurrent states
22
+ # · head_dim=256 (was 128), partial_rotary_factor=0.25
23
+ # · Interleaved M-RoPE with 4D position IDs
24
+ # · Attention output gating (sigmoid gate on q_proj)
25
+ # - Vision encoder: C-RADIOv4-H (unchanged)
26
+ # - Token IDs updated for Qwen3.5 vocabulary (vocab_size=248320)
27
+ # · img_context_token_id: 151655 → 248056 (<|image_pad|>)
28
+ # · video_context_token_id: 151656 → 248057 (<|video_pad|>)
29
+ # - Projector output adapts to Qwen3.5 hidden_size (4096 for 9B variant)
30
+ # ============================================================================
31
+
32
+ from transformers.configuration_utils import PretrainedConfig
33
+ from transformers.utils import logging
34
+ from .cradio_config import RADIOConfig
35
+
36
+ logger = logging.get_logger(__name__)
37
+
38
+ # ---------------------------------------------------------------------------
39
+ # Import Qwen3.5 text config — requires transformers >= 5.3.0
40
+ # ---------------------------------------------------------------------------
41
+ try:
42
+ from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5TextConfig
43
+ except ImportError:
44
+ Qwen3_5TextConfig = None
45
+ logger.warning(
46
+ "Could not import Qwen3_5TextConfig from transformers. "
47
+ "Ensure transformers >= 5.3.0 is installed. "
48
+ "Falling back to PretrainedConfig with manual attributes."
49
+ )
50
+
51
+
52
+ # ---------------------------------------------------------------------------
53
+ # Default Qwen3.5 text configuration (9B-class variant)
54
+ # ---------------------------------------------------------------------------
55
+
56
+ _LAYER_TYPES_32 = [
57
+ "linear_attention" if bool((i + 1) % 4) else "full_attention"
58
+ for i in range(32)
59
+ ]
60
+ # Result: [lin, lin, lin, full, lin, lin, lin, full, ... lin, lin, lin, full]
61
+ # 24 linear + 8 full attention layers
62
+
63
+
64
+ def _default_qwen3_5_text_dict() -> dict:
65
+ """Return a dict of Qwen3.5 text config values (9B-class)."""
66
+ return dict(
67
+ vocab_size=248320,
68
+ hidden_size=4096,
69
+ intermediate_size=12288,
70
+ num_hidden_layers=32,
71
+ num_attention_heads=16,
72
+ num_key_value_heads=4,
73
+ head_dim=256,
74
+ hidden_act="silu",
75
+ max_position_embeddings=262144,
76
+ rms_norm_eps=1e-6,
77
+ use_cache=True,
78
+ tie_word_embeddings=False,
79
+ attention_bias=False,
80
+ attention_dropout=0.0,
81
+ torch_dtype="bfloat16",
82
+ # --- Hybrid layer architecture ---
83
+ layer_types=list(_LAYER_TYPES_32), # copy to avoid mutation
84
+ full_attention_interval=4,
85
+ # --- Linear attention (Gated DeltaNet) ---
86
+ linear_conv_kernel_dim=4,
87
+ linear_key_head_dim=128,
88
+ linear_value_head_dim=128,
89
+ linear_num_key_heads=16,
90
+ linear_num_value_heads=32,
91
+ # --- RoPE ---
92
+ rope_parameters={
93
+ "rope_type": "default",
94
+ "rope_theta": 10000000,
95
+ "partial_rotary_factor": 0.25,
96
+ "mrope_interleaved": True,
97
+ "mrope_section": [11, 11, 10],
98
+ },
99
+ )
100
+
101
+
102
+ def _build_llm_config(cfg_dict: dict = None) -> PretrainedConfig:
103
+ """
104
+ Construct the LLM sub-config from a dict or defaults.
105
+
106
+ Uses Qwen3_5TextConfig when available (transformers >= 5.3);
107
+ otherwise falls back to a plain PretrainedConfig with the correct
108
+ model_type so that AutoModelForCausalLM can still resolve it.
109
+ """
110
+ if cfg_dict is None:
111
+ cfg_dict = _default_qwen3_5_text_dict()
112
+
113
+ if Qwen3_5TextConfig is not None:
114
+ return Qwen3_5TextConfig(**cfg_dict)
115
+ else:
116
+ config = PretrainedConfig(**cfg_dict)
117
+ config.model_type = "qwen3_5_text"
118
+ return config
119
+
120
+
121
+ class ZDTaichu5_0_Config(PretrainedConfig):
122
+ """
123
+ Configuration for ZDTaichu-5.0-9B:
124
+ Vision encoder : C-RADIOv4-H (ViT-H/16, 653 M params)
125
+ LLM decoder : Qwen3.5 (hybrid DeltaNet/Transformer)
126
+ Projector : RMSNorm → Linear(5120→20480) → SquaredReLU → Linear(20480→H)
127
+
128
+ The projector input side is unchanged (C-RADIOv4-H ViT-H features at 1280,
129
+ pixel-shuffled to 5120). Only the final projection layer adapts to the
130
+ target LLM hidden_size (4096 for the 9B variant, vs 5120 for Qwen3-14B).
131
+
132
+ Qwen3.5 hybrid architecture
133
+ ----------------------------
134
+ The text backbone alternates Gated DeltaNet (linear attention) and standard
135
+ multi-head attention layers in a 3:1 ratio. Linear layers use a causal 1D
136
+ convolution + gated delta rule recurrence for O(1) per-token memory during
137
+ generation, while every 4th layer uses full quadratic attention to preserve
138
+ global context. A custom DynamicCache handles both attention KV states and
139
+ recurrent states.
140
+ """
141
+
142
+ model_type = "zdtaichu5_0"
143
+ is_composition = True
144
+
145
+ def __init__(
146
+ self,
147
+ vision_config=None,
148
+ llm_config=None,
149
+ force_image_size=None,
150
+ downsample_ratio=0.5,
151
+ template=None,
152
+ ps_version="v2",
153
+ image_tag_type="internvl",
154
+ projector_hidden_size=20480, # 4 × pixel_shuffle_dim (5120)
155
+ vit_hidden_size=1280, # ViT-H feature dim — same for C-RADIOv4-H
156
+ attn_implementation="flash_attention_2",
157
+ # Special token IDs for Qwen3.5 vocabulary (vocab_size=248320)
158
+ img_context_token_id: int = 248056, # <|image_pad|>
159
+ video_context_token_id: int = 248057, # <|video_pad|>
160
+ **kwargs,
161
+ ):
162
+
163
+ # ------------------------------------------------------------------
164
+ # Transformers 5.5.x compatibility:
165
+ # PretrainedConfig.__init__ may call self.get_text_config()
166
+ # during token-id validation. Therefore llm_config must exist
167
+ # before calling super().__init__().
168
+ # ------------------------------------------------------------------
169
+
170
+ # ── Vision encoder ───────────────────────────────────────────────────
171
+ if vision_config is not None:
172
+ if isinstance(vision_config, dict):
173
+ self.vision_config = RADIOConfig(**vision_config)
174
+ else:
175
+ self.vision_config = vision_config
176
+ else:
177
+ self.vision_config = RADIOConfig(version="c-radio_v4-h")
178
+
179
+ # ── Language model (Qwen3.5 hybrid) ──────────────────────────────────
180
+ if llm_config is not None:
181
+ if isinstance(llm_config, PretrainedConfig):
182
+ self.llm_config = llm_config
183
+ elif isinstance(llm_config, dict):
184
+ self.llm_config = _build_llm_config(llm_config)
185
+ else:
186
+ raise TypeError(
187
+ f"llm_config must be a dict or PretrainedConfig, got {type(llm_config)}"
188
+ )
189
+ else:
190
+ self.llm_config = _build_llm_config(None)
191
+
192
+
193
+ # Make tokenizer/generation token ids visible early.
194
+ # Transformers 5.5.x may validate these during super().__init__().
195
+ kwargs.setdefault("bos_token_id", getattr(self.llm_config, "bos_token_id", 248040))
196
+ kwargs.setdefault("eos_token_id", getattr(self.llm_config, "eos_token_id", 248044))
197
+ kwargs.setdefault("pad_token_id", getattr(self.llm_config, "pad_token_id", 248040))
198
+ super().__init__(**kwargs)
199
+
200
+ self.tie_word_embeddings = getattr(self.llm_config, "tie_word_embeddings", False)
201
+
202
+ # ── VL configuration ─────────────────────────────────────────────────
203
+ self.force_image_size = force_image_size
204
+ self.downsample_ratio = downsample_ratio
205
+ self.template = template
206
+ self.ps_version = ps_version
207
+ self.image_tag_type = image_tag_type
208
+ self.projector_hidden_size = projector_hidden_size
209
+ self.vit_hidden_size = vit_hidden_size
210
+
211
+ # Special token IDs
212
+ self.img_context_token_id = img_context_token_id
213
+ self.video_context_token_id = video_context_token_id
214
+
215
+ # Attention implementation propagation
216
+ self._attn_implementation = attn_implementation
217
+ self.vision_config.use_flash_attn = (
218
+ self._attn_implementation is not None
219
+ and "flash_attention" in self._attn_implementation
220
+ )
221
+ self.llm_config._attn_implementation = self._attn_implementation
222
+
223
+ def get_text_config(self, decoder=False):
224
+ # Robust fallback for Transformers 5.5.x validation.
225
+ if hasattr(self, "llm_config"):
226
+ return self.llm_config
227
+ return _build_llm_config(None)
228
+
229
+ @property
230
+ def text_config(self):
231
+ return self.get_text_config(decoder=True)
cradio_config.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ # Copyright (c) 2026, ZDTaichu-5.0-9B Contributors. All rights reserved.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ #
16
+ # SPDX-License-Identifier: Apache-2.0
17
+
18
+ """Standalone inference configuration for the C-RADIO vision tower."""
19
+
20
+ from typing import Dict, List, Optional, Tuple, Union
21
+ from transformers import PretrainedConfig
22
+
23
+
24
+ class RADIOConfig(PretrainedConfig):
25
+ model_type = "radio"
26
+
27
+ def __init__(
28
+ self,
29
+ args: Optional[dict] = None,
30
+ version: str = "c-radio_v4-h",
31
+ patch_size: int = 16,
32
+ max_resolution: int = 2048,
33
+ preferred_resolution: Tuple[int, int] = (768, 768),
34
+ adaptor_names: Union[str, List[str], None] = None,
35
+ adaptor_configs: Optional[Dict] = None,
36
+ vitdet_window_size: Optional[int] = None,
37
+ feature_normalizer_config: Optional[dict] = None,
38
+ inter_feature_normalizer_config: Optional[dict] = None,
39
+ **kwargs,
40
+ ):
41
+ self.args = args or {}
42
+ self.version = version
43
+ self.patch_size = patch_size
44
+ self.max_resolution = max_resolution
45
+ self.preferred_resolution = preferred_resolution
46
+ self.adaptor_names = adaptor_names
47
+ self.adaptor_configs = adaptor_configs
48
+ self.vitdet_window_size = vitdet_window_size
49
+ self.feature_normalizer_config = feature_normalizer_config
50
+ self.inter_feature_normalizer_config = inter_feature_normalizer_config
51
+ super().__init__(**kwargs)
52
+
53
+
54
+ __all__ = ["RADIOConfig"]
cradio_model.py ADDED
@@ -0,0 +1,699 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ # Copyright (c) 2026, ZDTaichu-5.0-9B Contributors. All rights reserved.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ #
16
+ # SPDX-License-Identifier: Apache-2.0
17
+
18
+ """Standalone inference-only C-RADIO ViT vision tower.
19
+
20
+ This file intentionally contains the small subset of C-RADIO needed by the
21
+ ZDTaichu-5.0-9B checkpoint. It does not depend on the cradio_v4 package.
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ import math
27
+ from contextlib import contextmanager
28
+ from types import MethodType
29
+ from typing import Callable, Iterable, List, NamedTuple, Optional, Tuple, Union
30
+
31
+ import torch
32
+ import torch.nn.functional as F
33
+ from torch import nn
34
+ from transformers import PreTrainedModel
35
+
36
+ try:
37
+ from timm.models import VisionTransformer, checkpoint_seq
38
+ except ImportError as exc: # pragma: no cover - import-time dependency guard
39
+ raise ImportError("cradio_model.py requires timm to build the C-RADIO ViT tower") from exc
40
+
41
+ from .cradio_config import RADIOConfig
42
+
43
+ class Resolution(NamedTuple):
44
+ height: int
45
+ width: int
46
+
47
+
48
+ class RadioOutput(NamedTuple):
49
+ summary: Optional[torch.Tensor]
50
+ features: Optional[torch.Tensor]
51
+
52
+ def to(self, *args, **kwargs) -> "RadioOutput":
53
+ return RadioOutput(
54
+ self.summary.to(*args, **kwargs) if self.summary is not None else None,
55
+ self.features.to(*args, **kwargs) if self.features is not None else None,
56
+ )
57
+
58
+
59
+ class InputConditioner(nn.Module):
60
+ def __init__(
61
+ self,
62
+ input_scale: float,
63
+ norm_mean: Union[Tuple[float, float, float], torch.Tensor],
64
+ norm_std: Union[Tuple[float, float, float], torch.Tensor],
65
+ dtype: Optional[torch.dtype] = None,
66
+ ) -> None:
67
+ super().__init__()
68
+ self.dtype = dtype
69
+ self.register_buffer("norm_mean", torch.as_tensor(norm_mean, dtype=torch.float32).view(-1, 1, 1) / input_scale)
70
+ self.register_buffer("norm_std", torch.as_tensor(norm_std, dtype=torch.float32).view(-1, 1, 1) / input_scale)
71
+
72
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
73
+ y = (x - self.norm_mean) / self.norm_std
74
+ if self.dtype is not None:
75
+ y = y.to(self.dtype)
76
+ return y
77
+
78
+
79
+ def get_default_conditioner() -> InputConditioner:
80
+ from timm.data.constants import OPENAI_CLIP_MEAN, OPENAI_CLIP_STD
81
+
82
+ return InputConditioner(1.0, OPENAI_CLIP_MEAN, OPENAI_CLIP_STD)
83
+
84
+
85
+ class ClsToken(nn.Module):
86
+ def __init__(
87
+ self,
88
+ ndim: int,
89
+ num_tokens: int = 1,
90
+ enabled: bool = True,
91
+ register_multiple: Optional[int] = None,
92
+ num_registers: Optional[int] = None,
93
+ ) -> None:
94
+ super().__init__()
95
+ self.ndim = ndim
96
+ self.enabled = enabled
97
+ self.num_registers = 0
98
+ self.num_tokens = num_tokens
99
+ if enabled:
100
+ if num_registers:
101
+ self.num_registers = num_registers
102
+ elif register_multiple:
103
+ self.num_registers = register_multiple - (num_tokens % register_multiple)
104
+ scale = ndim ** -0.5
105
+ self.token = nn.Parameter(torch.randn(num_tokens + self.num_registers, ndim) * scale)
106
+ else:
107
+ self.token = None
108
+ self.num_patches = self.num_tokens + self.num_registers
109
+
110
+ def disable(self) -> None:
111
+ self.token = None
112
+ self.enabled = False
113
+
114
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
115
+ if self.token is None:
116
+ return x
117
+ token = self.token.unsqueeze(0).expand(x.shape[0], -1, -1)
118
+ return torch.cat([token, x], dim=1)
119
+
120
+ def no_weight_decay(self) -> List[str]:
121
+ return ["token"]
122
+
123
+
124
+ class Im2Patches(nn.Module):
125
+ def __init__(self, patch_size: int) -> None:
126
+ super().__init__()
127
+ self.patch_size = patch_size
128
+
129
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
130
+ if self.patch_size == 1:
131
+ return x.flatten(2).transpose(1, 2)
132
+ return F.unfold(x, kernel_size=self.patch_size, stride=self.patch_size).transpose(1, 2)
133
+
134
+
135
+ class ViTPatchLinear(nn.Linear):
136
+ def __init__(self, patch_size: int, embed_dim: int, bias: bool = False, **factory) -> None:
137
+ super().__init__(3 * (patch_size ** 2), embed_dim, bias=bias, **factory)
138
+ self.patch_size = patch_size
139
+
140
+
141
+ class ViTPatchGenerator(nn.Module):
142
+ def __init__(
143
+ self,
144
+ patch_size: int,
145
+ embed_dim: int,
146
+ input_dims: Union[int, Tuple[int, int]],
147
+ abs_pos: bool = True,
148
+ normalize_patches: bool = False,
149
+ cls_token: bool = False,
150
+ max_input_dims: Optional[Union[int, Tuple[int, int]]] = None,
151
+ pos_dropout: float = 0.0,
152
+ return_pos_enc: bool = False,
153
+ num_cls_tokens: int = 1,
154
+ register_multiple: Optional[int] = None,
155
+ num_registers: Optional[int] = None,
156
+ patch_bias: bool = False,
157
+ device=None,
158
+ dtype=None,
159
+ ) -> None:
160
+ super().__init__()
161
+ if isinstance(input_dims, int):
162
+ input_dims = (input_dims, input_dims)
163
+ if max_input_dims is None:
164
+ max_input_dims = input_dims
165
+ if isinstance(max_input_dims, int):
166
+ max_input_dims = (max_input_dims, max_input_dims)
167
+
168
+ max_input_dims = tuple(int(math.ceil(d / patch_size) * patch_size) for d in max_input_dims)
169
+ factory = dict(device=device, dtype=dtype)
170
+
171
+ self.cpe_mode = max_input_dims != input_dims
172
+ self.pos_dropout = pos_dropout
173
+ self.return_pos_enc = return_pos_enc
174
+ self.patch_size = patch_size
175
+ self.abs_pos = abs_pos
176
+ self.embed_dim = embed_dim
177
+ self.num_rows = max_input_dims[0] // patch_size
178
+ self.num_cols = max_input_dims[1] // patch_size
179
+ self.input_dims = tuple(d // patch_size for d in input_dims)
180
+ self.num_patches = self.num_rows * self.num_cols
181
+ self.max_input_dims = max_input_dims
182
+ self.im_to_patches = Im2Patches(patch_size)
183
+ self.embedder = ViTPatchLinear(patch_size, embed_dim, bias=patch_bias, **factory)
184
+ if abs_pos:
185
+ scale = embed_dim ** -0.5
186
+ self.pos_embed = nn.Parameter(torch.randn(1, self.num_patches, embed_dim, **factory) * scale)
187
+ self.cls_token = ClsToken(
188
+ embed_dim,
189
+ num_tokens=num_cls_tokens,
190
+ enabled=cls_token,
191
+ register_multiple=register_multiple,
192
+ num_registers=num_registers,
193
+ )
194
+ self.patch_normalizer = nn.LayerNorm(embed_dim) if normalize_patches else nn.Identity()
195
+ self.num_video_frames = None
196
+
197
+ @property
198
+ def apply_cls_token(self) -> bool:
199
+ return self.cls_token.enabled
200
+
201
+ @property
202
+ def num_cls_tokens(self) -> int:
203
+ return self.cls_token.num_tokens
204
+
205
+ @property
206
+ def num_cls_patches(self) -> int:
207
+ return self.cls_token.num_patches
208
+
209
+ @property
210
+ def num_registers(self) -> int:
211
+ return self.cls_token.num_registers
212
+
213
+ @property
214
+ def num_skip(self) -> int:
215
+ return self.num_cls_tokens + self.num_registers
216
+
217
+ def no_weight_decay(self) -> List[str]:
218
+ return ["pos_embed"]
219
+
220
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
221
+ patches = self.embedder(self.im_to_patches(x))
222
+ patches, pos_enc = self.apply_pos_enc(patches, input_size=x.shape[2:])
223
+ patches = self.cls_token(patches)
224
+ patches = self.patch_normalizer(patches)
225
+ if self.return_pos_enc:
226
+ return patches, pos_enc
227
+ return patches
228
+
229
+ def apply_pos_enc(
230
+ self,
231
+ patches: torch.Tensor,
232
+ patch_idxs: Optional[torch.Tensor] = None,
233
+ input_size: Optional[Tuple[int, int]] = None,
234
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
235
+ if not self.abs_pos:
236
+ return patches, torch.empty(0, device=patches.device, dtype=patches.dtype)
237
+ pos_enc = self.get_pos_enc(patches.shape[0], patch_idxs, input_size)
238
+ if self.training and self.pos_dropout > 0:
239
+ keeps = torch.rand(patches.shape[0], 1, 1, dtype=pos_enc.dtype, device=pos_enc.device) > self.pos_dropout
240
+ pos_enc_drop = torch.where(keeps, pos_enc, 0)
241
+ else:
242
+ pos_enc_drop = pos_enc
243
+ return patches + pos_enc_drop, pos_enc
244
+
245
+ def get_pos_enc(
246
+ self,
247
+ batch_size: int,
248
+ patch_idxs: Optional[torch.Tensor] = None,
249
+ input_size: Optional[Tuple[int, int]] = None,
250
+ ) -> torch.Tensor:
251
+ input_dims = self.input_dims if input_size is None else tuple(d // self.patch_size for d in input_size)
252
+ pos_embed = self._get_pos_embeddings(batch_size, input_dims)
253
+ if patch_idxs is None:
254
+ return pos_embed
255
+ exp_patch_idxs = patch_idxs.unsqueeze(-1).expand(-1, -1, pos_embed.shape[-1])
256
+ return torch.gather(pos_embed.expand(patch_idxs.shape[0], -1, -1), dim=1, index=exp_patch_idxs)
257
+
258
+ def _get_pos_embeddings(self, batch_size: int, input_dims: Tuple[int, int]) -> torch.Tensor:
259
+ if (self.num_rows, self.num_cols) == input_dims:
260
+ return self.pos_embed
261
+
262
+ pos_embed = self.pos_embed.reshape(1, self.num_rows, self.num_cols, -1).permute(0, 3, 1, 2)
263
+
264
+ def window_select(pe: torch.Tensor) -> torch.Tensor:
265
+ if input_dims[0] < pe.shape[-2]:
266
+ pe = pe[..., :input_dims[0], :]
267
+ if input_dims[1] < pe.shape[-1]:
268
+ pe = pe[..., :, :input_dims[1]]
269
+ return pe
270
+
271
+ if self.cpe_mode:
272
+ if self.training:
273
+ if self.num_video_frames is not None:
274
+ if batch_size % self.num_video_frames != 0:
275
+ raise ValueError(
276
+ f"Batch size {batch_size} must be divisible by num_video_frames "
277
+ f"{self.num_video_frames} for CPE mode."
278
+ )
279
+ batch_size //= self.num_video_frames
280
+
281
+ min_scale = math.sqrt(0.1)
282
+ scale = torch.rand(batch_size, 1, 1, device=pos_embed.device) * (1 - min_scale) + min_scale
283
+ aspect_min = math.log(3 / 4)
284
+ aspect = torch.exp(torch.rand(batch_size, 1, 1, device=pos_embed.device) * (-2 * aspect_min) + aspect_min)
285
+ scale_xy = torch.stack([scale * aspect, scale / aspect], dim=-1).clamp_(0, 1)
286
+ pos_xy = torch.rand(batch_size, 1, 1, 2, device=pos_embed.device) * (1 - scale_xy)
287
+ lin_x = torch.linspace(0, 1, steps=input_dims[1], device=pos_embed.device)[None, None].expand(batch_size, input_dims[0], -1)
288
+ lin_y = torch.linspace(0, 1, steps=input_dims[0], device=pos_embed.device)[None, :, None].expand(batch_size, -1, input_dims[1])
289
+ grid_xy = torch.stack([lin_x, lin_y], dim=-1) * scale_xy + pos_xy
290
+ grid_xy.mul_(2).sub_(1)
291
+ pos_embed = F.grid_sample(
292
+ pos_embed.float().expand(batch_size, -1, -1, -1),
293
+ grid=grid_xy,
294
+ mode="bilinear",
295
+ padding_mode="zeros",
296
+ align_corners=True,
297
+ ).to(pos_embed.dtype)
298
+ if self.num_video_frames is not None:
299
+ pos_embed = torch.repeat_interleave(pos_embed, self.num_video_frames, dim=0)
300
+ else:
301
+ max_dim = max(input_dims)
302
+ pos_embed = F.interpolate(pos_embed.float(), size=(max_dim, max_dim), align_corners=False, mode="bilinear").to(pos_embed.dtype)
303
+ pos_embed = window_select(pos_embed)
304
+ else:
305
+ pos_embed = window_select(pos_embed)
306
+
307
+ if pos_embed.shape[-2:] != input_dims:
308
+ pos_embed = F.interpolate(pos_embed.float(), size=input_dims, align_corners=False, mode="bilinear").to(pos_embed.dtype)
309
+ return pos_embed.flatten(2).permute(0, 2, 1)
310
+
311
+
312
+ def _forward_cpe(self: VisionTransformer, x: torch.Tensor) -> torch.Tensor:
313
+ x = self.patch_generator(x)
314
+ if getattr(self, "grad_checkpointing", False) and not torch.jit.is_scripting():
315
+ x = checkpoint_seq(self.blocks, x)
316
+ else:
317
+ x = self.blocks(x)
318
+ x = self.norm(x)
319
+ return x
320
+
321
+
322
+ @contextmanager
323
+ def _video_mode(self: VisionTransformer, t: int):
324
+ original_num_frames = self.patch_generator.num_video_frames
325
+ self.patch_generator.num_video_frames = t
326
+ try:
327
+ yield
328
+ finally:
329
+ self.patch_generator.num_video_frames = original_num_frames
330
+
331
+
332
+ def enable_cpe(
333
+ model: VisionTransformer,
334
+ max_img_size: Union[int, Tuple[int, int]] = 1024,
335
+ num_cls_tokens: int = 1,
336
+ pos_dropout: float = 0.1,
337
+ register_multiple: Optional[int] = None,
338
+ num_registers: Optional[int] = None,
339
+ ) -> None:
340
+ if not isinstance(model, VisionTransformer):
341
+ raise ValueError(f"CPE only supports timm VisionTransformer models, got {type(model)}")
342
+
343
+ patch_size = model.patch_embed.patch_size[0]
344
+ embed_dim = model.embed_dim
345
+ input_dims = model.patch_embed.img_size
346
+ normalize_patches = not isinstance(model.patch_embed.norm, nn.Identity)
347
+ cls_token = model.cls_token is not None
348
+ if isinstance(max_img_size, int):
349
+ max_img_size = int(round(max_img_size / patch_size) * patch_size)
350
+ else:
351
+ max_img_size = tuple(int(round(d / patch_size) * patch_size) for d in max_img_size)
352
+
353
+ model.patch_generator = ViTPatchGenerator(
354
+ patch_size=patch_size,
355
+ embed_dim=embed_dim,
356
+ input_dims=input_dims,
357
+ normalize_patches=normalize_patches,
358
+ cls_token=cls_token,
359
+ max_input_dims=max_img_size,
360
+ pos_dropout=pos_dropout,
361
+ num_cls_tokens=num_cls_tokens,
362
+ register_multiple=register_multiple,
363
+ num_registers=num_registers,
364
+ )
365
+ model.patch_embed = None
366
+ model.cls_token = None
367
+ model.pos_embed = None
368
+ model.pos_drop = None
369
+ model.patch_size = patch_size
370
+ model.num_cls_tokens = num_cls_tokens
371
+ model.num_registers = model.patch_generator.num_registers
372
+ model.forward_features = MethodType(_forward_cpe, model)
373
+ model.cpe_video_mode = MethodType(_video_mode, model)
374
+
375
+
376
+ class FeatureNormalizer(nn.Module):
377
+ def __init__(self, embed_dim: int, dtype: torch.dtype = torch.float32) -> None:
378
+ super().__init__()
379
+ self.register_buffer("mean", torch.zeros(embed_dim, dtype=dtype))
380
+ self.register_buffer("tx", torch.eye(embed_dim, dtype=dtype))
381
+
382
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
383
+ if x.ndim <= 3:
384
+ return (x - self.mean) @ self.tx.T
385
+ if x.ndim == 4:
386
+ kernel = self.tx.reshape(*self.tx.shape, 1, 1)
387
+ return F.conv2d(x - self.mean.reshape(1, -1, 1, 1), weight=kernel, bias=None, stride=1, padding=0)
388
+ raise ValueError(f"Unsupported input dimension: {x.ndim}, shape: {x.shape}")
389
+
390
+
391
+ class InnerRADIOModel(nn.Module):
392
+ def __init__(
393
+ self,
394
+ model: nn.Module,
395
+ input_conditioner: nn.Module,
396
+ patch_size: int,
397
+ max_resolution: int,
398
+ preferred_resolution: Resolution,
399
+ summary_idxs: Optional[torch.Tensor] = None,
400
+ feature_normalizer: Optional[nn.Module] = None,
401
+ window_size: Optional[int] = None,
402
+ ) -> None:
403
+ super().__init__()
404
+ self.model = model
405
+ self.input_conditioner = input_conditioner
406
+ if summary_idxs is not None:
407
+ self.register_buffer("summary_idxs", summary_idxs)
408
+ else:
409
+ self.summary_idxs = None
410
+ self._preferred_resolution = preferred_resolution
411
+ self._patch_size = patch_size
412
+ self._max_resolution = max_resolution
413
+ self._window_size = window_size
414
+ self.feature_normalizer = feature_normalizer if feature_normalizer is not None else nn.Identity()
415
+
416
+ @property
417
+ def num_summary_tokens(self) -> int:
418
+ patch_gen = getattr(self.model, "patch_generator", None)
419
+ if patch_gen is not None:
420
+ return patch_gen.num_skip
421
+ if getattr(self.model, "global_pool", None) == "avg":
422
+ return 0
423
+ return 1
424
+
425
+ @property
426
+ def num_cls_tokens(self) -> int:
427
+ patch_gen = getattr(self.model, "patch_generator", None)
428
+ if patch_gen is not None:
429
+ return patch_gen.num_cls_tokens
430
+ if getattr(self.model, "global_pool", None) == "avg":
431
+ return 0
432
+ return 1
433
+
434
+ @property
435
+ def patch_size(self) -> int:
436
+ if self._patch_size is not None:
437
+ return self._patch_size
438
+ if hasattr(self.model, "patch_size"):
439
+ return self.model.patch_size
440
+ patch_gen = getattr(self.model, "patch_generator", None)
441
+ if patch_gen is not None:
442
+ return patch_gen.patch_size
443
+ raise AttributeError("Unable to infer patch_size from RADIO vision model")
444
+
445
+ @property
446
+ def max_resolution(self) -> int:
447
+ return self._max_resolution
448
+
449
+ @property
450
+ def preferred_resolution(self) -> Resolution:
451
+ return self._preferred_resolution
452
+
453
+ @property
454
+ def window_size(self) -> Optional[int]:
455
+ return self._window_size
456
+
457
+ @property
458
+ def min_resolution_step(self) -> int:
459
+ res = self.patch_size
460
+ if self.window_size is not None:
461
+ res *= self.window_size
462
+ return res
463
+
464
+ @property
465
+ def blocks(self) -> Iterable[nn.Module]:
466
+ return getattr(self.model, "blocks", None)
467
+
468
+ @property
469
+ def embed_dim(self) -> int:
470
+ return self.model.embed_dim
471
+
472
+ @property
473
+ def summary_dim(self) -> int:
474
+ embed_dim = self.embed_dim
475
+ if self.summary_idxs is not None:
476
+ embed_dim *= self.summary_idxs.shape[0]
477
+ return embed_dim
478
+
479
+ def make_preprocessor_external(self) -> Callable[[torch.Tensor], torch.Tensor]:
480
+ ret = self.input_conditioner
481
+ self.input_conditioner = nn.Identity()
482
+ return ret
483
+
484
+ def get_nearest_supported_resolution(self, height: int, width: int) -> Resolution:
485
+ height = int(round(height / self.min_resolution_step) * self.min_resolution_step)
486
+ width = int(round(width / self.min_resolution_step) * self.min_resolution_step)
487
+ return Resolution(max(height, self.min_resolution_step), max(width, self.min_resolution_step))
488
+
489
+ def switch_to_deploy(self) -> None:
490
+ fn = getattr(self.model, "switch_to_deploy", None)
491
+ if fn is not None:
492
+ fn()
493
+
494
+ def cpe_video_mode(self, t: int):
495
+ return self.model.cpe_video_mode(t)
496
+
497
+ def forward(self, x: torch.Tensor, feature_fmt: str = "NLC") -> RadioOutput:
498
+ res_step = self.min_resolution_step
499
+ if res_step is not None and (x.shape[-2] % res_step != 0 or x.shape[-1] % res_step != 0):
500
+ raise ValueError(
501
+ "The input resolution must be a multiple of self.min_resolution_step. "
502
+ f"Input: {x.shape[-2:]}, Nearest: {self.get_nearest_supported_resolution(*x.shape[-2:])}"
503
+ )
504
+ x = self.input_conditioner(x)
505
+ y = self.model.forward_features(x)
506
+ return self._extract_final(x, y, feature_fmt=feature_fmt)
507
+
508
+ def _extract_final(self, x: torch.Tensor, y: torch.Tensor, feature_fmt: str = "NLC") -> RadioOutput:
509
+ patch_gen = getattr(self.model, "patch_generator", None)
510
+ if patch_gen is not None:
511
+ all_summary = y[:, : patch_gen.num_cls_tokens]
512
+ bb_summary = all_summary[:, self.summary_idxs] if self.summary_idxs is not None else all_summary
513
+ all_feat = y[:, patch_gen.num_skip :]
514
+ elif getattr(self.model, "global_pool", None) == "avg":
515
+ all_summary = y[:, self.model.num_prefix_tokens :].mean(dim=1)
516
+ bb_summary = all_summary
517
+ all_feat = y
518
+ else:
519
+ all_summary = y[:, 0]
520
+ bb_summary = all_summary
521
+ all_feat = y[:, 1:]
522
+
523
+ all_feat = self.feature_normalizer(all_feat)
524
+ if feature_fmt == "NCHW":
525
+ fmt_feat = all_feat.reshape(
526
+ all_feat.shape[0],
527
+ x.shape[-2] // self.patch_size,
528
+ x.shape[-1] // self.patch_size,
529
+ all_feat.shape[2],
530
+ ).permute(0, 3, 1, 2)
531
+ elif feature_fmt == "NLC":
532
+ fmt_feat = all_feat
533
+ else:
534
+ raise ValueError(f"Unsupported feature_fmt: {feature_fmt}. Must be one of ['NLC', 'NCHW']")
535
+ return RadioOutput(bb_summary.flatten(1), fmt_feat)
536
+
537
+
538
+ def _as_namespace(value):
539
+ if value is None:
540
+ return type("RADIOArgs", (), {})()
541
+ if isinstance(value, dict):
542
+ ns = type("RADIOArgs", (), {})()
543
+ for k, v in value.items():
544
+ setattr(ns, k, v)
545
+ return ns
546
+ return value
547
+
548
+
549
+ def _dtype_from_config(config: RADIOConfig) -> torch.dtype:
550
+ dtype_name = getattr(config, "dtype", None) or getattr(config, "amp_dtype", None)
551
+ if isinstance(dtype_name, torch.dtype):
552
+ return dtype_name
553
+ if isinstance(dtype_name, str) and hasattr(torch, dtype_name):
554
+ return getattr(torch, dtype_name)
555
+ return torch.float32
556
+
557
+
558
+ def create_vit_from_config(config: RADIOConfig) -> VisionTransformer:
559
+ args = _as_namespace(getattr(config, "args", {}))
560
+ model_name = getattr(args, "model", None) or "vit_huge_patch16_224"
561
+ if model_name != "vit_huge_patch16_224":
562
+ raise ValueError(
563
+ "This standalone cradio_model.py keeps only the ZDTaichu ViT-H/16 structure. "
564
+ f"Unsupported RADIO args.model={model_name!r}."
565
+ )
566
+
567
+ model = VisionTransformer(
568
+ img_size=224,
569
+ patch_size=16,
570
+ embed_dim=1280,
571
+ depth=32,
572
+ num_heads=16,
573
+ mlp_ratio=4.0,
574
+ qkv_bias=True,
575
+ num_classes=0,
576
+ global_pool="",
577
+ )
578
+
579
+ # The ZDTaichu checkpoint was exported after RADIO removed the final ViT norm/head
580
+ # and replaced patch embedding, cls token, and absolute pos embedding with CPE.
581
+ if hasattr(model, "norm") and not getattr(args, "model_norm", False):
582
+ model.norm = nn.Identity()
583
+ model.head = nn.Identity()
584
+
585
+ cpe_max_size = getattr(args, "cpe_max_size", None) or getattr(config, "max_resolution", None)
586
+ if cpe_max_size is not None:
587
+ teachers = getattr(args, "teachers", []) or []
588
+ teacher_names = {t.get("name") for t in teachers if isinstance(t, dict) and t.get("name")}
589
+ num_cls_tokens = len(teacher_names) if getattr(args, "cls_token_per_teacher", False) and teacher_names else 1
590
+ enable_cpe(
591
+ model,
592
+ cpe_max_size,
593
+ num_cls_tokens=num_cls_tokens,
594
+ register_multiple=getattr(args, "register_multiple", None),
595
+ num_registers=getattr(args, "cpe_num_registers", None),
596
+ )
597
+ return model
598
+
599
+
600
+ class RADIOModel(PreTrainedModel):
601
+ """Inference-only HuggingFace wrapper for the ZDTaichu C-RADIO ViT tower."""
602
+
603
+ config_class = RADIOConfig
604
+ base_model_prefix = "radio_model"
605
+ main_input_name = "pixel_values"
606
+ supports_gradient_checkpointing = False
607
+
608
+ def __init__(self, config: RADIOConfig) -> None:
609
+ super().__init__(config)
610
+ args = _as_namespace(getattr(config, "args", {}))
611
+ dtype = _dtype_from_config(config)
612
+ vit = create_vit_from_config(config)
613
+
614
+ summary_idxs = None
615
+ if getattr(args, "cls_token_per_teacher", False):
616
+ teachers = getattr(args, "teachers", []) or []
617
+ if teachers:
618
+ summary_idxs = torch.tensor(
619
+ [i for i, t in enumerate(teachers) if not isinstance(t, dict) or t.get("use_summary", True)],
620
+ dtype=torch.int64,
621
+ )
622
+
623
+ feature_normalizer = None
624
+ fn_cfg = getattr(config, "feature_normalizer_config", None)
625
+ if fn_cfg is not None:
626
+ embed_dim = fn_cfg.get("embed_dim", vit.embed_dim) if isinstance(fn_cfg, dict) else vit.embed_dim
627
+ feature_normalizer = FeatureNormalizer(embed_dim, dtype=torch.float32)
628
+
629
+ pref = getattr(config, "preferred_resolution", (512, 512))
630
+ self.radio_model = InnerRADIOModel(
631
+ model=vit,
632
+ input_conditioner=get_default_conditioner(),
633
+ patch_size=getattr(config, "patch_size", 16),
634
+ max_resolution=getattr(config, "max_resolution", 2048),
635
+ preferred_resolution=Resolution(int(pref[0]), int(pref[1])),
636
+ summary_idxs=summary_idxs,
637
+ feature_normalizer=feature_normalizer,
638
+ window_size=getattr(config, "vitdet_window_size", None),
639
+ )
640
+ if dtype is not torch.float32:
641
+ self.radio_model = self.radio_model.to(dtype=dtype)
642
+
643
+ @property
644
+ def adaptors(self):
645
+ return nn.ModuleDict()
646
+
647
+ @property
648
+ def model(self) -> nn.Module:
649
+ return self.radio_model.model
650
+
651
+ @property
652
+ def input_conditioner(self) -> nn.Module:
653
+ return self.radio_model.input_conditioner
654
+
655
+ @property
656
+ def num_summary_tokens(self) -> int:
657
+ return self.radio_model.num_summary_tokens
658
+
659
+ @property
660
+ def patch_size(self) -> int:
661
+ return self.radio_model.patch_size
662
+
663
+ @property
664
+ def max_resolution(self) -> int:
665
+ return self.radio_model.max_resolution
666
+
667
+ @property
668
+ def preferred_resolution(self) -> Resolution:
669
+ return self.radio_model.preferred_resolution
670
+
671
+ @property
672
+ def window_size(self) -> Optional[int]:
673
+ return self.radio_model.window_size
674
+
675
+ @property
676
+ def min_resolution_step(self) -> int:
677
+ return self.radio_model.min_resolution_step
678
+
679
+ def make_preprocessor_external(self) -> Callable[[torch.Tensor], torch.Tensor]:
680
+ return self.radio_model.make_preprocessor_external()
681
+
682
+ def get_nearest_supported_resolution(self, height: int, width: int) -> Resolution:
683
+ return self.radio_model.get_nearest_supported_resolution(height, width)
684
+
685
+ def switch_to_deploy(self) -> None:
686
+ self.radio_model.switch_to_deploy()
687
+
688
+ def forward(self, pixel_values: torch.Tensor, feature_fmt: str = "NLC", **kwargs) -> RadioOutput:
689
+ return self.radio_model(pixel_values, feature_fmt=feature_fmt)
690
+
691
+
692
+ __all__ = [
693
+ "RADIOModel",
694
+ "RADIOConfig",
695
+ "RadioOutput",
696
+ "Resolution",
697
+ "InputConditioner",
698
+ "ViTPatchGenerator",
699
+ ]
generation_config.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 248040,
4
+ "do_sample": true,
5
+ "eos_token_id": [
6
+ 248044,
7
+ 248040,
8
+ 248046
9
+ ],
10
+ "pad_token_id": 248040,
11
+ "repetition_penalty": 1.0,
12
+ "temperature": 1.0,
13
+ "top_k": 20,
14
+ "top_p": 0.95,
15
+ "transformers_version": "5.3.0"
16
+ }
image_processing.py ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ from typing import List, Optional, Union, Any, Dict, Tuple
15
+
16
+ from PIL import Image
17
+ import torch
18
+ from transformers.image_processing_base import BatchFeature
19
+ from transformers.image_processing_utils_fast import BaseImageProcessorFast
20
+ from transformers.image_utils import make_list_of_images, get_image_type, ImageInput, ImageType
21
+ from transformers.utils import TensorType
22
+ import torchvision.transforms as T
23
+
24
+ import math
25
+
26
+ class ZDTaichu5_0_ImageProcessor(BaseImageProcessorFast):
27
+ model_input_names = ["pixel_values", "image_grid_thw"]
28
+
29
+ def __init__(self, image_size=512, max_num_tiles=12, use_thumbnail=True, norm_mean=None, norm_std=None, do_rescale=True, patch_size=16, downsample_ratio=0.5, merge_size=1, **kwargs):
30
+ super().__init__(**kwargs)
31
+ self.image_size = image_size
32
+ self.max_num_tiles = max_num_tiles
33
+ self.use_thumbnail = use_thumbnail
34
+ self.norm_mean = norm_mean
35
+ self.norm_std = norm_std
36
+ self.do_rescale = do_rescale
37
+ self.merge_size = merge_size
38
+ self.num_image_token = int((image_size // patch_size) ** 2 * (downsample_ratio ** 2))
39
+
40
+ def _process_image(
41
+ self,
42
+ image: ImageInput,
43
+ **kwargs,
44
+ ) -> torch.Tensor:
45
+ image_type = get_image_type(image)
46
+ if image_type == ImageType.PIL:
47
+ if image.mode != 'RGB':
48
+ image = image.convert('RGB')
49
+ # Keep PIL input through tiling so resize order matches vLLM.
50
+ return image
51
+
52
+ def _preprocess(
53
+ self,
54
+ images: List[torch.Tensor],
55
+ image_size: int = None,
56
+ max_num_tiles: int = None,
57
+ use_thumbnail: bool = None,
58
+ do_rescale: bool = None,
59
+ return_tensors: Optional[Union[str, TensorType]] = None,
60
+ **kwargs,
61
+ ) -> List[torch.Tensor]:
62
+ image_size = image_size if image_size is not None else self.image_size
63
+ max_num_tiles = max_num_tiles if max_num_tiles is not None else self.max_num_tiles
64
+ use_thumbnail = use_thumbnail if use_thumbnail is not None else self.use_thumbnail
65
+ do_rescale = do_rescale if do_rescale is not None else self.do_rescale
66
+
67
+ images = make_list_of_images(images)
68
+
69
+ all_patches = []
70
+ num_patches = []
71
+ image_grid_thw = []
72
+ for image in images:
73
+ patches, tile_rows, tile_cols = dynamic_preprocess(image, image_size, max_num_tiles, use_thumbnail)
74
+ all_patches.extend(patches)
75
+ num_patches.append(len(patches))
76
+ image_grid_thw.append([1, tile_rows, tile_cols])
77
+
78
+ # vLLM converts each already-cropped PIL tile with ToTensor.
79
+ pixel_values = torch.stack([T.ToTensor()(patch) for patch in all_patches], dim=0)
80
+ norm_mean = torch.Tensor(self.norm_mean).view(1, 3, 1, 1)
81
+ norm_std = torch.Tensor(self.norm_std).view(1, 3, 1, 1)
82
+ pixel_values = (pixel_values - norm_mean) / norm_std
83
+ pixel_values = pixel_values.to(torch.bfloat16)
84
+ return BatchFeature(
85
+ data={
86
+ "pixel_values": pixel_values,
87
+ "num_patches": num_patches,
88
+ "image_grid_thw": image_grid_thw,
89
+ },
90
+ tensor_type=return_tensors,
91
+ )
92
+
93
+
94
+ def get_internvl_target_ratios(
95
+ min_num: int,
96
+ max_num: int,
97
+ ) -> list[tuple[int, int]]:
98
+ target_ratios = {(i, j)
99
+ for n in range(min_num, max_num + 1)
100
+ for i in range(1, n + 1)
101
+ for j in range(1, n + 1) if min_num <= i * j <= max_num}
102
+ return sorted(target_ratios, key=lambda x: x[0] * x[1])
103
+
104
+
105
+ # From https://github.com/OpenGVLab/InternVL/blob/c62fa4f7c850165d7386bdc48ac6bc5a6fab0864/internvl_chat/internvl/train/dataset.py#L685
106
+ # Copyright (c) 2023 OpenGVLab.
107
+ def find_closest_aspect_ratio(
108
+ aspect_ratio: float,
109
+ target_ratios: list[tuple[int, int]],
110
+ width: int,
111
+ height: int,
112
+ image_size: int,
113
+ ) -> tuple[int, int]:
114
+ best_ratio_diff = float("inf")
115
+ best_ratio = (1, 1)
116
+ area = width * height
117
+ for ratio in target_ratios:
118
+ target_aspect_ratio = ratio[0] / ratio[1]
119
+ ratio_diff = abs(aspect_ratio - target_aspect_ratio)
120
+ if ratio_diff < best_ratio_diff:
121
+ best_ratio_diff = ratio_diff
122
+ best_ratio = ratio
123
+ elif ratio_diff == best_ratio_diff:
124
+ if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:
125
+ best_ratio = ratio
126
+ return best_ratio
127
+
128
+ def select_tile_grid(
129
+ *,
130
+ orig_width: int,
131
+ orig_height: int,
132
+ image_size: int,
133
+ min_num_tiles: int,
134
+ max_num_tiles: int,
135
+ ) -> Tuple[int, int, int, int]:
136
+ """
137
+ Choose (rw, rh) tile grid with small-image and aspect-sanity guards.
138
+
139
+ Returns (num_grid_blocks, target_width, target_height, effective_max_tiles).
140
+ num_grid_blocks = rw * rh (does NOT include the optional thumbnail).
141
+
142
+ Guards:
143
+ 1. Area cap: don't create more tiles than the source has pixels for.
144
+ 2. Aspect-sanity: drop candidates whose ratio differs from source by >3x.
145
+ """
146
+ # ── Guard 1: area cap ─────────────────────────────────────────────────
147
+ src_pixels = orig_width * orig_height
148
+ tile_pixels = image_size * image_size
149
+ area_max_tiles = max(1, math.ceil(src_pixels / tile_pixels))
150
+ effective_max = min(max_num_tiles, area_max_tiles)
151
+ effective_max = max(effective_max, min_num_tiles)
152
+
153
+ target_ratios = get_internvl_target_ratios(min_num_tiles, effective_max)
154
+
155
+ # ── Guard 2: aspect-sanity ────────────────────────────────────────────
156
+ src_ar = orig_width / orig_height
157
+ filtered = [
158
+ (rw, rh) for (rw, rh) in target_ratios
159
+ if (1.0 / 3.0) <= (rw / rh) / src_ar <= 3.0
160
+ ]
161
+ # Fall back to unfiltered set for extreme panoramas / long strips where
162
+ # no candidate is within 3x — better to pick *something* than error.
163
+ if filtered:
164
+ target_ratios = filtered
165
+
166
+ # ── Pick best ratio ───────────────────────────────────────────────────
167
+ rw, rh = find_closest_aspect_ratio(
168
+ src_ar, target_ratios,
169
+ width=orig_width, height=orig_height, image_size=image_size,
170
+ )
171
+
172
+ target_width = image_size * rw
173
+ target_height = image_size * rh
174
+ num_grid_blocks = rw * rh
175
+
176
+ return num_grid_blocks, target_width, target_height, effective_max
177
+
178
+
179
+ def count_tiles(
180
+ *,
181
+ orig_width: int,
182
+ orig_height: int,
183
+ image_size: int,
184
+ min_num_tiles: int,
185
+ max_num_tiles: int,
186
+ use_thumbnail: bool,
187
+ ) -> int:
188
+ """
189
+ Total number of tiles (grid blocks + optional thumbnail) for this image.
190
+ This MUST match what the actual tiling produces, or prompt expansion and
191
+ embedding count will diverge.
192
+ """
193
+ n_grid, _, _, _ = select_tile_grid(
194
+ orig_width=orig_width, orig_height=orig_height,
195
+ image_size=image_size,
196
+ min_num_tiles=min_num_tiles, max_num_tiles=max_num_tiles,
197
+ )
198
+ if use_thumbnail and n_grid != 1:
199
+ return n_grid + 1
200
+ return n_grid
201
+
202
+ def calculate_targets(
203
+ orig_width: int,
204
+ orig_height: int,
205
+ target_ratios: list[tuple[int, int]],
206
+ image_size: int,
207
+ ) -> tuple[int, int, int]:
208
+ aspect_ratio = orig_width / orig_height
209
+
210
+ # find the closest aspect ratio to the target
211
+ target_aspect_ratio = find_closest_aspect_ratio(
212
+ aspect_ratio,
213
+ target_ratios,
214
+ width=orig_width,
215
+ height=orig_height,
216
+ image_size=image_size,
217
+ )
218
+
219
+ # calculate the target width and height
220
+ target_width = image_size * target_aspect_ratio[0]
221
+ target_height = image_size * target_aspect_ratio[1]
222
+ blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
223
+
224
+ return blocks, target_width, target_height
225
+
226
+ def dynamic_preprocess(image, image_size=512, max_num_tiles=12, use_thumbnail=True, min_num_tiles=1):
227
+ """Split a PIL image using vLLM's resize/crop order."""
228
+ if isinstance(image, torch.Tensor):
229
+ image = T.ToPILImage()(image)
230
+ elif not isinstance(image, Image.Image):
231
+ image = Image.fromarray(image)
232
+ if image.mode != 'RGB':
233
+ image = image.convert('RGB')
234
+ orig_width, orig_height = image.size
235
+
236
+ n_grid, target_width, target_height, _ = select_tile_grid(
237
+ orig_width=orig_width,
238
+ orig_height=orig_height,
239
+ image_size=image_size,
240
+ min_num_tiles=min_num_tiles,
241
+ max_num_tiles=max_num_tiles,
242
+ )
243
+
244
+ # Tile grid dimensions (rows × cols of the InternVL tiling)
245
+ tile_rows = target_height // image_size
246
+ tile_cols = target_width // image_size
247
+
248
+ resized_img = image.resize((target_width, target_height), Image.BICUBIC)
249
+ cols = target_width // image_size
250
+ patches = []
251
+ for i in range(n_grid):
252
+ col = i % cols
253
+ row = i // cols
254
+ patches.append(
255
+ resized_img.crop(
256
+ (col * image_size, row * image_size,
257
+ (col + 1) * image_size, (row + 1) * image_size)
258
+ )
259
+ )
260
+ assert len(patches) == n_grid
261
+
262
+ if use_thumbnail and n_grid != 1:
263
+ thumbnail = image.resize((image_size, image_size), Image.BICUBIC)
264
+ patches.append(thumbnail)
265
+
266
+ #print(orig_height, orig_width, target_width, target_height, len(patches))
267
+
268
+ return patches, tile_rows, tile_cols
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d26e50e66812bdb9563368991b4ad9be1eeaab8aa59b5220df5efbc89e202f3b
3
+ size 9811274448
modeling.py ADDED
@@ -0,0 +1,1109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ # ============================================================================
16
+ # ZDTaichu-5.0 — Main Model
17
+ #
18
+ # Architecture
19
+ # ────────────
20
+ # Vision encoder : C-RADIOv4-H (ViT-H/16, 653 M)
21
+ # Projector : RMSNorm → Linear(5120→20480) → SquaredReLU → Linear(20480→H)
22
+ # LLM decoder : Qwen3.5 hybrid (Gated DeltaNet + full attention, 3:1 ratio)
23
+ #
24
+ # Position encoding (M-RoPE)
25
+ # ──────────────────────────
26
+ # Vision tokens receive 3D position IDs (temporal, height, width) computed
27
+ # from the InternVL-style tile grid via ``get_rope_index()``. Text tokens
28
+ # receive standard 1D positions (all three M-RoPE channels are identical).
29
+ #
30
+ # This matches the official Qwen3.5 VL pipeline where ``Qwen3_5Model.forward()``
31
+ # calls ``compute_3d_position_ids()`` → ``get_rope_index()`` before forwarding
32
+ # to ``Qwen3_5TextModel``. The resulting ``position_ids`` of shape ``(3, B, S)``
33
+ # are consumed directly by ``Qwen3_5TextRotaryEmbedding``, which applies
34
+ # interleaved M-RoPE across temporal / height / width frequency bands.
35
+ #
36
+ # Generation
37
+ # ──────────
38
+ # This model inherits from ``GenerationMixin``, owning the generation loop
39
+ # (like ``Qwen3_5ForConditionalGeneration``). Key overrides:
40
+ # - ``_prepare_position_ids_for_generation``: computes 3D ``position_ids``
41
+ # on the prefill step and caches ``rope_deltas``; applies ``rope_deltas``
42
+ # on subsequent decode steps.
43
+ # - ``prepare_inputs_for_generation``: clears ``pixel_values`` /
44
+ # ``pixel_values_videos`` after the first step (vision features are
45
+ # already embedded in the KV cache).
46
+ #
47
+ # Cache handling
48
+ # ──────────────
49
+ # ``Qwen3_5DynamicCache`` is created internally by ``Qwen3_5TextModel`` when
50
+ # ``use_cache=True``. It stores KV states for full-attention layers and
51
+ # ``conv_states`` + ``recurrent_states`` for Gated DeltaNet layers.
52
+ # ============================================================================
53
+
54
+ import itertools
55
+ import warnings
56
+ from typing import Any, Dict, List, Optional, Tuple, Union
57
+
58
+ import torch
59
+ import transformers
60
+ from torch import nn
61
+ from torch.nn import CrossEntropyLoss
62
+ from transformers import AutoModel, GenerationConfig
63
+ from transformers.generation import GenerationMixin
64
+ from transformers.modeling_outputs import CausalLMOutputWithPast
65
+ from transformers.modeling_utils import PreTrainedModel
66
+ from transformers.utils import logging
67
+
68
+ from .configuration import ZDTaichu5_0_Config
69
+ from .cradio_model import RADIOModel
70
+
71
+ logger = logging.get_logger(__name__)
72
+
73
+ # ---------------------------------------------------------------------------
74
+ # Import Qwen3.5 model classes — requires transformers >= 5.3.0
75
+ # ---------------------------------------------------------------------------
76
+
77
+ _MIN_TRANSFORMERS = "5.3.0"
78
+
79
+ try:
80
+ from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5ForCausalLM
81
+ from transformers.cache_utils import DynamicCache as Qwen3_5DynamicCache
82
+ _HAS_QWEN3_5 = True
83
+ except Exception as e:
84
+ _HAS_QWEN3_5 = False
85
+ Qwen3_5ForCausalLM = None
86
+ Qwen3_5DynamicCache = None
87
+ logger.warning(
88
+ f"Could not import Qwen3_5ForCausalLM from transformers. "
89
+ f"Import error: {e!r}"
90
+ )
91
+
92
+
93
+ def _version_ge(v1, v2):
94
+ """Check if version v1 >= v2."""
95
+ from packaging import version
96
+ return version.parse(v1) >= version.parse(v2)
97
+
98
+
99
+ # ─────────────────────────────────────────────────────────────────────────────
100
+ # Projector components
101
+ # ─────────────────────────────────────────────────────────────────────────────
102
+
103
+ class SquaredReLU(nn.Module):
104
+ """Squared ReLU activation — same non-linearity used in the projector."""
105
+ def forward(self, x):
106
+ return torch.pow(torch.nn.functional.relu(x), 2)
107
+
108
+
109
+ class RMSNorm(nn.Module):
110
+ """
111
+ Standard RMSNorm for the projector (NOT the Qwen3.5 LLM variant).
112
+
113
+ Qwen3.5's internal ``Qwen3_5RMSNorm`` uses zero-initialized weight with
114
+ ``output * (1 + weight)``. The projector uses ones-initialized weight
115
+ with ``output * weight`` — the standard formulation.
116
+ """
117
+ def __init__(self, hidden_size: int, eps: float = 1e-5):
118
+ super().__init__()
119
+ self.weight = nn.Parameter(torch.ones(hidden_size))
120
+ self.eps = eps
121
+
122
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
123
+ input_dtype = hidden_states.dtype
124
+ hidden_states = hidden_states.to(torch.float32)
125
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
126
+ hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
127
+ return (self.weight.to(torch.float32) * hidden_states).to(input_dtype)
128
+
129
+
130
+ # ─────────────────────────────────────────────────────────────────────────────
131
+ # Main model
132
+ # ─────────────────────────────────────────────────────────────────────────────
133
+
134
+ class ZDTaichu5_0_ForConditionalGeneration(PreTrainedModel, GenerationMixin):
135
+ """
136
+ ZDTaichu-5.0: C-RADIOv4-H vision encoder + Qwen3.5 language decoder.
137
+
138
+ Architecture overview::
139
+
140
+ pixel_values
141
+ └─► C-RADIOv4-H (ViT-H/16, 653 M)
142
+ └─► pixel_shuffle(0.5)
143
+ └─► mlp1: RMSNorm → Linear → SquaredReLU → Linear
144
+ └─► inject into Qwen3.5 embeddings at <image> positions
145
+ └─► Qwen3.5 (hybrid DeltaNet / Transformer)
146
+ """
147
+
148
+ config_class = ZDTaichu5_0_Config
149
+ main_input_name = "input_ids"
150
+ _tied_weights_keys = None#["language_model.lm_head.weight"]
151
+ _keys_to_ignore_on_load_unexpected = [
152
+ # The RADIO input_conditioner registers norm_mean / norm_std as
153
+ # buffers, but make_preprocessor_external() removes the conditioner
154
+ # at init time (normalization is handled by the image processor).
155
+ # The build script still saves these from the source checkpoint, so
156
+ # they appear as unexpected keys during loading — safe to ignore.
157
+ r"vision_model\.radio_model\.input_conditioner\..*",
158
+ r"^mtp\..*",
159
+ ]
160
+
161
+ _supports_flash_attn_2 = True
162
+ _supports_flash_attention_2 = True
163
+ _supports_flash_attn = True
164
+ _supports_sdpa = True
165
+ _no_split_modules = ["Qwen3_5DecoderLayer"]
166
+ _is_stateful = True
167
+ supports_gradient_checkpointing = True
168
+
169
+ def __init__(self, config: ZDTaichu5_0_Config):
170
+ super().__init__(config)
171
+
172
+ # Guard for bleeding-edge transformers (>= 4.57.0.dev) where
173
+ # _finalize_model_loading reads all_tied_weights_keys but
174
+ # PreTrainedModel.__init__ may not yet initialise it.
175
+ if not hasattr(self, "all_tied_weights_keys"):
176
+ self.all_tied_weights_keys = {}
177
+
178
+ assert _version_ge(transformers.__version__, _MIN_TRANSFORMERS), (
179
+ f"Qwen3.5 support requires transformers >= {_MIN_TRANSFORMERS} "
180
+ f"(found {transformers.__version__})"
181
+ )
182
+ assert _HAS_QWEN3_5, (
183
+ "Qwen3_5ForCausalLM is not available. "
184
+ f"Ensure transformers >= {_MIN_TRANSFORMERS} is installed."
185
+ )
186
+
187
+ image_size = config.force_image_size
188
+ patch_size = config.vision_config.patch_size
189
+ self.patch_size = patch_size
190
+ self.template = config.template
191
+ self.num_image_token = int(
192
+ (image_size // patch_size) ** 2 * (config.downsample_ratio ** 2)
193
+ )
194
+ self.downsample_ratio = config.downsample_ratio
195
+ self.ps_version = config.ps_version
196
+ self.image_tag_type = config.image_tag_type
197
+ self.img_context_token_id = config.img_context_token_id
198
+ self.video_context_token_id = config.video_context_token_id
199
+
200
+ # Per-tile token dimensions (e.g. 14×14 for 448px, patch=16, ds=0.5)
201
+ self.tile_h = int((image_size // patch_size) * config.downsample_ratio)
202
+ self.tile_w = self.tile_h
203
+
204
+ logger.info(f"num_image_token: {self.num_image_token}")
205
+ logger.info(f"tile_h={self.tile_h}, tile_w={self.tile_w}")
206
+ logger.info(f"ps_version: {self.ps_version}")
207
+ logger.info(f"Vision encoder: {config.vision_config.version}")
208
+ logger.info(
209
+ f"LLM: Qwen3.5 ({config.llm_config.num_hidden_layers} layers, "
210
+ f"hidden={config.llm_config.hidden_size}, "
211
+ f"hybrid="
212
+ f"{sum(1 for t in config.llm_config.layer_types if t == 'linear_attention')} linear + "
213
+ f"{sum(1 for t in config.llm_config.layer_types if t == 'full_attention')} full)"
214
+ )
215
+
216
+ # ── Language model ───────────────────────────────────────────────────
217
+ self.language_model = Qwen3_5ForCausalLM(config.llm_config)
218
+
219
+ # ── Vision encoder ───────────────────────────────────────────────────
220
+ self.vision_model = RADIOModel(config.vision_config)
221
+ self.vision_model.model._initialize_weights = (
222
+ self.vision_model.model._init_weights
223
+ )
224
+ self.vision_model.radio_model.make_preprocessor_external()
225
+ self.vision_model = self.vision_model.to(
226
+ self.language_model.config.torch_dtype
227
+ )
228
+
229
+ self.drop_vision_class_token = True
230
+
231
+ # ── MLP projector ────────────────────────────────────────────────────
232
+ vit_hidden_size = config.vit_hidden_size
233
+ proj_hidden = config.projector_hidden_size
234
+ llm_hidden = config.llm_config.hidden_size
235
+ pixel_shuffle_dim = vit_hidden_size * int(1 / self.downsample_ratio) ** 2
236
+
237
+ self.mlp1 = nn.Sequential(
238
+ RMSNorm(pixel_shuffle_dim, eps=1e-5),
239
+ nn.Linear(pixel_shuffle_dim, proj_hidden, bias=False),
240
+ SquaredReLU(),
241
+ nn.Linear(proj_hidden, llm_hidden, bias=False),
242
+ )
243
+ self.mlp1 = self.mlp1.to(self.language_model.config.torch_dtype)
244
+
245
+ # Cached rope_deltas for multi-step generation
246
+ self.rope_deltas = None
247
+
248
+ # ── Embedding accessors (required by GenerationMixin) ─────────────────
249
+
250
+ def get_input_embeddings(self):
251
+ return self.language_model.get_input_embeddings()
252
+
253
+ def set_input_embeddings(self, value):
254
+ self.language_model.set_input_embeddings(value)
255
+
256
+ def get_output_embeddings(self):
257
+ return self.language_model.lm_head
258
+
259
+ def set_output_embeddings(self, new_embeddings):
260
+ self.language_model.lm_head = new_embeddings
261
+
262
+ def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs=None):
263
+ # 大头在 LLM:直接委托给内层 Qwen3.5(它原生支持 GC)
264
+ self.language_model.gradient_checkpointing_enable(
265
+ gradient_checkpointing_kwargs=gradient_checkpointing_kwargs
266
+ )
267
+ # 视觉塔可选:支持就开,不支持就跳过(不影响主显存)
268
+ vm = getattr(self, "vision_model", None)
269
+ if vm is not None and getattr(vm, "supports_gradient_checkpointing", False):
270
+ try:
271
+ vm.gradient_checkpointing_enable(
272
+ gradient_checkpointing_kwargs=gradient_checkpointing_kwargs
273
+ )
274
+ except Exception:
275
+ pass
276
+
277
+ def gradient_checkpointing_disable(self):
278
+ self.language_model.gradient_checkpointing_disable()
279
+ vm = getattr(self, "vision_model", None)
280
+ if vm is not None and hasattr(vm, "gradient_checkpointing_disable"):
281
+ try:
282
+ vm.gradient_checkpointing_disable()
283
+ except Exception:
284
+ pass
285
+
286
+ # ─────────────────────────────────────────────────────────────────────────
287
+ # Vision helpers
288
+ # ─────────────────────────────────────────────────────────────────────────
289
+
290
+ def pixel_shuffle(
291
+ self, x: torch.Tensor, scale_factor: float = 0.5
292
+ ) -> torch.Tensor:
293
+ """Space-to-depth rearrangement (ps_version='v2' = corrected layout)."""
294
+ n, w, h, c = x.size()
295
+ x = x.view(n, w, int(h * scale_factor), int(c / scale_factor))
296
+ x = x.permute(0, 2, 1, 3).contiguous()
297
+ x = x.view(
298
+ n, int(h * scale_factor), int(w * scale_factor),
299
+ int(c / (scale_factor * scale_factor)),
300
+ )
301
+ if self.ps_version == "v1":
302
+ warnings.warn(
303
+ "ps_version='v1' produces a transposed spatial layout. "
304
+ "Use ps_version='v2' for correct output."
305
+ )
306
+ else:
307
+ x = x.permute(0, 2, 1, 3).contiguous()
308
+ return x
309
+
310
+ def extract_feature(self, pixel_values: torch.Tensor) -> torch.Tensor:
311
+ """Run pixels through C-RADIOv4-H → pixel_shuffle → MLP projector."""
312
+ vit_embeds = self.vision_model(pixel_values).features
313
+ vit_embeds = vit_embeds.to(dtype=torch.bfloat16)
314
+
315
+ h = w = int(vit_embeds.shape[1] ** 0.5)
316
+ vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], h, w, -1)
317
+ vit_embeds = self.pixel_shuffle(
318
+ vit_embeds, scale_factor=self.downsample_ratio
319
+ )
320
+ vit_embeds = vit_embeds.reshape(
321
+ vit_embeds.shape[0], -1, vit_embeds.shape[-1]
322
+ )
323
+ vit_embeds = self.mlp1(vit_embeds)
324
+ return vit_embeds
325
+
326
+ # ─────────────────────────────────────────────────────────────────────────
327
+ # 3D M-RoPE position IDs
328
+ # ─────────────────────────────────────────────────────────────────────────
329
+
330
+ def get_vision_position_ids(
331
+ self,
332
+ start_position: int,
333
+ tile_rows: int,
334
+ tile_cols: int,
335
+ has_thumbnail: bool = True,
336
+ device: torch.device = None,
337
+ ) -> torch.LongTensor:
338
+ """
339
+ Compute 3D (temporal, height, width) position IDs for vision tokens
340
+ from a single InternVL-style tiled image.
341
+
342
+ Token layout (flattened order expected by the model):
343
+ 1. Grid tiles in raster order: tile(0,0), tile(0,1), …, tile(R-1,C-1).
344
+ Each tile has ``tile_h × tile_w`` tokens in raster order.
345
+ 2. Thumbnail tile (optional): a single tile covering the full image
346
+ at reduced resolution.
347
+
348
+ Args:
349
+ start_position: Offset added to all positional indices.
350
+ tile_rows: Number of tile rows in the image grid.
351
+ tile_cols: Number of tile columns in the image grid.
352
+ has_thumbnail: Whether a thumbnail tile is appended after grid tiles.
353
+ device: Target device.
354
+
355
+ Returns:
356
+ ``torch.LongTensor`` of shape ``(3, num_vision_tokens)``.
357
+ """
358
+ tile_h, tile_w = self.tile_h, self.tile_w
359
+ npt = tile_h * tile_w # num tokens per tile
360
+
361
+ # ── Grid tiles ───────────────────────────────────────────────────────
362
+ num_grid_tiles = tile_rows * tile_cols
363
+ tile_idx = torch.arange(num_grid_tiles, device=device)
364
+ tr = tile_idx // tile_cols
365
+ tc = tile_idx % tile_cols
366
+
367
+ local_idx = torch.arange(npt, device=device)
368
+ lr = local_idx // tile_w
369
+ lc = local_idx % tile_w
370
+
371
+ # (num_grid_tiles, npt) → flatten
372
+ global_h = (tr[:, None] * tile_h + lr[None, :]).reshape(-1).long()
373
+ global_w = (tc[:, None] * tile_w + lc[None, :]).reshape(-1).long()
374
+
375
+ total_grid = num_grid_tiles * npt
376
+ pos_t = torch.full(
377
+ (total_grid,), start_position, device=device, dtype=torch.long
378
+ )
379
+ pos_h = start_position + global_h
380
+ pos_w = start_position + global_w
381
+
382
+ # ── Thumbnail tile ───────────────────────────────────────────────────
383
+ if has_thumbnail:
384
+ # Map thumbnail local(r, c) → global(r * tile_rows, c * tile_cols)
385
+ # so its positions overlay the grid at coarser resolution.
386
+ thumb_h = (lr * tile_rows).long()
387
+ thumb_w = (lc * tile_cols).long()
388
+ pos_t = torch.cat([
389
+ pos_t,
390
+ torch.full(
391
+ (npt,), start_position, device=device, dtype=torch.long
392
+ ),
393
+ ])
394
+ pos_h = torch.cat([pos_h, start_position + thumb_h])
395
+ pos_w = torch.cat([pos_w, start_position + thumb_w])
396
+
397
+ return torch.stack([pos_t, pos_h, pos_w], dim=0)
398
+
399
+ def get_rope_index(
400
+ self,
401
+ input_ids: torch.LongTensor,
402
+ mm_token_type_ids: torch.IntTensor,
403
+ image_grid_thw: Optional[torch.LongTensor] = None,
404
+ video_grid_thw: Optional[torch.LongTensor] = None,
405
+ attention_mask: Optional[torch.Tensor] = None,
406
+ **kwargs,
407
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
408
+ """
409
+ Compute 3D M-RoPE position IDs for a mixed vision + text sequence.
410
+
411
+ Follows the same structure as ``Qwen3_5Model.get_rope_index``:
412
+ uses ``mm_token_type_ids`` to group tokens by modality
413
+ (text=0, image=1, video=2) via ``itertools.groupby``. Vision
414
+ tokens receive spatial position IDs (temporal, height, width)
415
+ while text tokens receive standard 1D positions.
416
+
417
+ Args:
418
+ input_ids: ``(B, S)`` token IDs.
419
+ mm_token_type_ids: ``(B, S)`` modality labels —
420
+ 0 = text, 1 = image, 2 = video.
421
+ image_grid_thw: ``(num_images, 3)`` — each row
422
+ ``(T=1, tile_rows, tile_cols)`` for InternVL-style tiled images.
423
+ video_grid_thw: ``(num_videos, 3)`` — each row
424
+ ``(num_frames, 1, 1)``.
425
+ attention_mask: ``(B, S)`` binary mask.
426
+
427
+ Returns:
428
+ ``position_ids``: ``(3, B, S)``
429
+ ``mrope_position_deltas``: ``(B, 1)``
430
+ """
431
+ tile_h, tile_w = self.tile_h, self.tile_w
432
+ npt = tile_h * tile_w
433
+
434
+ B, S = input_ids.shape
435
+ device = input_ids.device
436
+
437
+ position_ids = torch.zeros(3, B, S, dtype=input_ids.dtype, device=device)
438
+ mrope_position_deltas = []
439
+
440
+ # ------------------------------------------------------------------
441
+ # video-as-image compatibility for verl / vLLM rollout.
442
+ # ------------------------------------------------------------------
443
+ if mm_token_type_ids is not None and video_grid_thw is None and torch.any(mm_token_type_ids == 2).item():
444
+ mm_token_type_ids = mm_token_type_ids.clone()
445
+
446
+ if image_grid_thw is not None:
447
+ # Count contiguous visual groups, because get_rope_index consumes
448
+ # one grid_thw row per contiguous image/video segment.
449
+ total_visual_groups = 0
450
+ for b in range(mm_token_type_ids.shape[0]):
451
+ cur_types = mm_token_type_ids[b]
452
+ if attention_mask is not None:
453
+ cur_types = cur_types[attention_mask[b].bool()]
454
+
455
+ prev_type = None
456
+ for t in cur_types.tolist():
457
+ if t in (1, 2) and t != prev_type:
458
+ total_visual_groups += 1
459
+ prev_type = t
460
+
461
+ num_image_grids = image_grid_thw.shape[0]
462
+
463
+ if total_visual_groups <= num_image_grids:
464
+ # True video-as-image case: consume image_grid_thw for both image and video types.
465
+ mm_token_type_ids[mm_token_type_ids == 2] = 1
466
+
467
+ if "logger" in globals():
468
+ logger.warning_once(
469
+ "Converting mm_token_type_ids type 2 to type 1 because "
470
+ "video_grid_thw is None and image_grid_thw has enough grids. "
471
+ "This matches video-as-image processing."
472
+ )
473
+ else:
474
+ # Some type-2 tokens are likely generated orphan <|video_pad|> tokens.
475
+ # Treat them as text to avoid consuming non-existent grids.
476
+ mm_token_type_ids[mm_token_type_ids == 2] = 0
477
+
478
+ if "logger" in globals():
479
+ logger.warning_once(
480
+ "mm_token_type_ids contains type 2 but video_grid_thw is None, "
481
+ "and image_grid_thw does not have enough grids. Treating type 2 "
482
+ "as text. This likely means the model generated orphan <|video_pad|> tokens."
483
+ )
484
+ else:
485
+ # No visual grid exists, so type 2 cannot represent valid visual tokens.
486
+ mm_token_type_ids[mm_token_type_ids == 2] = 0
487
+
488
+ if "logger" in globals():
489
+ logger.warning_once(
490
+ "mm_token_type_ids contains type 2, but both video_grid_thw and "
491
+ "image_grid_thw are None. Treating type 2 as text."
492
+ )
493
+
494
+ grid_iters = {
495
+ 1: iter(image_grid_thw) if image_grid_thw is not None else None,
496
+ 2: iter(video_grid_thw) if video_grid_thw is not None else None,
497
+ }
498
+
499
+ for batch_idx, current_input_ids in enumerate(input_ids):
500
+ input_token_type = mm_token_type_ids[batch_idx]
501
+ if attention_mask is not None:
502
+ current_input_ids = current_input_ids[attention_mask[batch_idx].bool()]
503
+ input_token_type = input_token_type[attention_mask[batch_idx].bool()]
504
+
505
+ # Group contiguous runs of the same modality type
506
+ input_type_group = []
507
+ for key, group in itertools.groupby(
508
+ enumerate(input_token_type.tolist()), lambda x: x[1]
509
+ ):
510
+ group = list(group)
511
+ start_index = group[0][0]
512
+ end_index = group[-1][0] + 1
513
+ input_type_group.append((key, start_index, end_index))
514
+
515
+ current_pos = 0
516
+ llm_pos_ids_list: List[torch.Tensor] = []
517
+
518
+ # ── Per-video state machine ──────────────────────────────────────
519
+ # Mirrors the Megatron-side implementation in
520
+ # modeling.py: a single video_grid_thw entry
521
+ # of [num_frames, 1, 1] is consumed across multiple non-contiguous
522
+ # type-2 runs (one per <|video_pad|> block, separated by frame
523
+ # header text).
524
+ #
525
+ # Within a video, every frame's tokens use:
526
+ # t = vid_spatial_start + frame_idx (anchored at video start)
527
+ # h = vid_spatial_start + local_row (constant across frames)
528
+ # w = vid_spatial_start + local_col (constant across frames)
529
+ #
530
+ # Text between frames advances ``current_pos`` normally — those
531
+ # text positions live in a different range than the video frame
532
+ # positions, which is fine for M-RoPE (RoPE requires no
533
+ # monotonicity, only consistent training/inference).
534
+ vid_active = False
535
+ vid_num_frames = 0
536
+ vid_frame_idx = 0
537
+ vid_spatial_start = 0
538
+
539
+ for modality_type, start_idx, end_idx in input_type_group:
540
+ # text == 0
541
+ if modality_type == 0:
542
+ text_len = end_idx - start_idx
543
+ llm_pos_ids_list.append(
544
+ torch.arange(text_len, device=device).view(1, -1).expand(3, -1)
545
+ + current_pos
546
+ )
547
+ current_pos += text_len
548
+
549
+ # image == 1
550
+ elif modality_type == 1:
551
+ seg_len = end_idx - start_idx
552
+ grid = next(grid_iters[1])
553
+ tile_rows = grid[1].item()
554
+ tile_cols = grid[2].item()
555
+ grid_tokens = tile_rows * tile_cols * npt
556
+ has_thumbnail = seg_len > grid_tokens
557
+
558
+ vpos = self.get_vision_position_ids(
559
+ start_position=current_pos,
560
+ tile_rows=tile_rows,
561
+ tile_cols=tile_cols,
562
+ has_thumbnail=has_thumbnail,
563
+ device=device,
564
+ )
565
+ assert vpos.shape[1] == seg_len, (
566
+ f"Position count ({vpos.shape[1]}) ≠ image token count "
567
+ f"({seg_len}) for grid=({tile_rows},{tile_cols}), "
568
+ f"thumbnail={has_thumbnail}"
569
+ )
570
+ llm_pos_ids_list.append(vpos)
571
+ current_pos += max(tile_rows * tile_h, tile_cols * tile_w)
572
+
573
+ # video == 2
574
+ elif modality_type == 2:
575
+ seg_len = end_idx - start_idx
576
+
577
+ # Activate per-video state on the FIRST type-2 run for
578
+ # this video. Subsequent type-2 runs (one per frame
579
+ # block, separated by frame-header text) reuse the same
580
+ # vid_spatial_start anchor.
581
+ if not vid_active:
582
+ grid = next(grid_iters[2])
583
+ vid_num_frames = grid[0].item()
584
+ vid_active = True
585
+ vid_frame_idx = 0
586
+ vid_spatial_start = current_pos
587
+
588
+ # Each frame contributes exactly ``npt`` tokens.
589
+ if seg_len % npt != 0:
590
+ raise ValueError(
591
+ f"Video segment length {seg_len} is not a "
592
+ f"multiple of npt={npt} (tile_h*tile_w). "
593
+ f"Check that the processor produced one "
594
+ f"<|video_pad|> block per frame with exactly "
595
+ f"npt tokens each."
596
+ )
597
+ frames_in_run = seg_len // npt
598
+
599
+ # Sanity guard against malformed grids — never consume
600
+ # more frames than the grid declared.
601
+ if vid_frame_idx + frames_in_run > vid_num_frames:
602
+ raise ValueError(
603
+ f"Video has {vid_num_frames} frames but "
604
+ f"input_ids contain at least "
605
+ f"{vid_frame_idx + frames_in_run} frame blocks. "
606
+ f"Check the processor's video_grid_thw against "
607
+ f"the actual <|video_pad|> count."
608
+ )
609
+
610
+ local_idx = torch.arange(npt, device=device)
611
+ lr = local_idx // tile_w
612
+ lc = local_idx % tile_w
613
+
614
+ all_t, all_h, all_w = [], [], []
615
+ for _ in range(frames_in_run):
616
+ # Temporal: anchored at video_start, advances by frame_idx.
617
+ all_t.append(torch.full(
618
+ (npt,),
619
+ vid_spatial_start + vid_frame_idx,
620
+ device=device, dtype=torch.long,
621
+ ))
622
+ # Spatial: constant base across frames within this video.
623
+ all_h.append((vid_spatial_start + lr).long())
624
+ all_w.append((vid_spatial_start + lc).long())
625
+ vid_frame_idx += 1
626
+ # Advance current_pos by one frame's spatial extent so
627
+ # subsequent text positions stay strictly above any
628
+ # h/w position used by this video. After all frames,
629
+ # current_pos has advanced by num_frames * max(tile_h, tile_w),
630
+ # which always exceeds vid_spatial_start + max(num_frames, tile_h, tile_w)
631
+ # for num_frames >= 1 (so text after the video sees
632
+ # positions strictly greater than every video token).
633
+ current_pos += max(tile_h, tile_w)
634
+
635
+ vpos = torch.stack([
636
+ torch.cat(all_t), torch.cat(all_h), torch.cat(all_w),
637
+ ], dim=0)
638
+ assert vpos.shape[1] == seg_len, (
639
+ f"Position count ({vpos.shape[1]}) ≠ video token "
640
+ f"count ({seg_len})"
641
+ )
642
+ llm_pos_ids_list.append(vpos)
643
+
644
+ # End the video once all declared frames have been
645
+ # consumed; reset state so the next video (if any) gets
646
+ # a fresh grid pull.
647
+ if vid_frame_idx >= vid_num_frames:
648
+ vid_active = False
649
+ vid_num_frames = 0
650
+ vid_frame_idx = 0
651
+ vid_spatial_start = 0
652
+
653
+ # Sanity check: if a video's last frame isn't followed by any text,
654
+ # the loop ends with vid_active=False (we already reset on the
655
+ # final frame). But if the input is malformed and the type-2
656
+ # runs don't cover all declared frames, surface that loudly
657
+ # rather than silently advancing the iterator the next time we
658
+ # see another video.
659
+ if vid_active:
660
+ raise ValueError(
661
+ f"Reached end of input with video state still active: "
662
+ f"consumed {vid_frame_idx}/{vid_num_frames} frames. "
663
+ f"video_grid_thw declares more frames than the "
664
+ f"<|video_pad|> blocks contain."
665
+ )
666
+
667
+ llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)
668
+ if attention_mask is not None:
669
+ position_ids[:, batch_idx, attention_mask[batch_idx].bool()] = (
670
+ llm_positions.to(position_ids.device)
671
+ )
672
+ else:
673
+ position_ids[:, batch_idx] = llm_positions.to(position_ids.device)
674
+
675
+ mrope_position_deltas.append(
676
+ llm_positions.max() + 1 - len(current_input_ids)
677
+ )
678
+
679
+ mrope_position_deltas = torch.tensor(
680
+ mrope_position_deltas, device=device
681
+ ).unsqueeze(1)
682
+ return position_ids, mrope_position_deltas
683
+
684
+ def _build_text_position_ids(
685
+ self,
686
+ input_ids: torch.LongTensor,
687
+ attention_mask: Optional[torch.Tensor] = None,
688
+ ) -> torch.LongTensor:
689
+ """
690
+ Build text position ids of shape (B, S).
691
+ For padding mask, positions are 0,1,2,... on valid tokens.
692
+ Padding positions stay 0.
693
+ """
694
+ batch_size, seq_len = input_ids.shape
695
+ device = input_ids.device
696
+
697
+ if attention_mask is not None:
698
+ valid = attention_mask > 0
699
+ text_position_ids = valid.long().cumsum(-1) - 1
700
+ text_position_ids = text_position_ids.masked_fill(~valid, 0)
701
+ else:
702
+ text_position_ids = torch.arange(
703
+ seq_len, device=device, dtype=torch.long
704
+ ).unsqueeze(0).expand(batch_size, -1)
705
+
706
+ return text_position_ids.contiguous()
707
+ def _prepend_text_position_channel(
708
+ self,
709
+ input_ids: torch.LongTensor,
710
+ vision_position_ids: torch.LongTensor,
711
+ attention_mask: Optional[torch.Tensor] = None,
712
+ ) -> torch.LongTensor:
713
+ """
714
+ Convert vision M-RoPE position ids from (3, B, S) to Qwen3.5-compatible
715
+ position ids of shape (4, B, S):
716
+
717
+ channel 0 : text positions, used for causal mask / FA2 varlen logic
718
+ channel 1-3 : temporal / height / width vision M-RoPE positions
719
+ """
720
+ if vision_position_ids is None:
721
+ return None
722
+
723
+ if vision_position_ids.dim() == 3 and vision_position_ids.shape[0] == 4:
724
+ return vision_position_ids.contiguous()
725
+
726
+ assert vision_position_ids.dim() == 3 and vision_position_ids.shape[0] == 3, (
727
+ f"Expected vision_position_ids shape (3, B, S), got "
728
+ f"{tuple(vision_position_ids.shape)}"
729
+ )
730
+
731
+ text_position_ids = self._build_text_position_ids(
732
+ input_ids=input_ids,
733
+ attention_mask=attention_mask,
734
+ ).to(device=vision_position_ids.device)
735
+
736
+ position_ids = torch.cat(
737
+ [
738
+ text_position_ids.unsqueeze(0), # (1, B, S)
739
+ vision_position_ids, # (3, B, S)
740
+ ],
741
+ dim=0,
742
+ )
743
+ return position_ids.contiguous()
744
+
745
+ def _compute_position_ids(
746
+ self,
747
+ input_ids: Optional[torch.LongTensor],
748
+ inputs_embeds: torch.FloatTensor,
749
+ image_grid_thw: Optional[torch.LongTensor],
750
+ video_grid_thw: Optional[torch.LongTensor],
751
+ attention_mask: Optional[torch.Tensor],
752
+ past_key_values=None,
753
+ mm_token_type_ids: Optional[torch.IntTensor] = None,
754
+ use_cache: Optional[bool] = None,
755
+ ) -> Optional[torch.Tensor]:
756
+ """
757
+ Mirror of ``Qwen3_5Model.compute_3d_position_ids``.
758
+
759
+ - Vision info available + first forward → ``get_rope_index``, cache
760
+ ``rope_deltas``.
761
+ - ``rope_deltas`` cached (decode step) → derive from attention_mask +
762
+ ``rope_deltas``.
763
+ - Pure text → return ``None`` (``Qwen3_5TextModel`` auto-generates).
764
+ """
765
+ past_length = 0
766
+ if past_key_values is not None:
767
+ past_length = past_key_values.get_seq_length()
768
+
769
+ can_compute = (
770
+ input_ids is not None
771
+ and mm_token_type_ids is not None
772
+ and (image_grid_thw is not None or video_grid_thw is not None)
773
+ )
774
+
775
+ if can_compute and past_length == 0:
776
+ vision_position_ids, rope_deltas = self.get_rope_index(
777
+ input_ids,
778
+ mm_token_type_ids=mm_token_type_ids,
779
+ image_grid_thw=image_grid_thw,
780
+ video_grid_thw=video_grid_thw,
781
+ attention_mask=attention_mask,
782
+ )
783
+
784
+ # Training / log-prob forward should not keep rope_deltas across batches.
785
+ # Generation prefill can keep it for decode.
786
+ if use_cache:
787
+ self.rope_deltas = rope_deltas
788
+ else:
789
+ self.rope_deltas = None
790
+
791
+ return self._prepend_text_position_channel(
792
+ input_ids=input_ids,
793
+ vision_position_ids=vision_position_ids,
794
+ attention_mask=attention_mask,
795
+ )
796
+
797
+ elif self.rope_deltas is not None and past_length != 0:
798
+ batch_size, seq_length = inputs_embeds.shape[:2]
799
+
800
+ if attention_mask is not None:
801
+ text_position_ids = attention_mask.long().cumsum(-1) - 1
802
+ text_position_ids = text_position_ids.masked_fill(attention_mask == 0, 0)
803
+ text_position_ids = text_position_ids[:, -seq_length:]
804
+ else:
805
+ text_position_ids = torch.arange(
806
+ past_length,
807
+ past_length + seq_length,
808
+ device=inputs_embeds.device,
809
+ dtype=torch.long,
810
+ ).unsqueeze(0).expand(batch_size, -1)
811
+
812
+ delta = self.rope_deltas.repeat_interleave(
813
+ batch_size // self.rope_deltas.shape[0], dim=0
814
+ ).to(device=inputs_embeds.device)
815
+
816
+ # Decode step follows generation convention: (1, B, S)
817
+ position_ids = text_position_ids.unsqueeze(0) + delta.view(1, batch_size, 1)
818
+ return position_ids.contiguous()
819
+
820
+ return None
821
+
822
+ # ─────────────────────────────────────────────────────────────────────────
823
+ # Forward
824
+ # ─────────────────────────────────────────────────────────────────────────
825
+
826
+ def forward(
827
+ self,
828
+ input_ids: torch.LongTensor = None,
829
+ pixel_values: Optional[torch.FloatTensor] = None,
830
+ pixel_values_videos: Optional[torch.FloatTensor] = None,
831
+ num_patches = None,
832
+ image_flags: Optional[torch.LongTensor] = None,
833
+ image_grid_thw: Optional[torch.LongTensor] = None,
834
+ video_grid_thw: Optional[torch.LongTensor] = None,
835
+ mm_token_type_ids: Optional[torch.IntTensor] = None,
836
+ attention_mask: Optional[torch.Tensor] = None,
837
+ position_ids: Optional[torch.LongTensor] = None,
838
+ past_key_values=None,
839
+ labels: Optional[torch.LongTensor] = None,
840
+ inputs_embeds: Optional[torch.FloatTensor] = None,
841
+ use_cache: Optional[bool] = None,
842
+ cache_position: Optional[torch.LongTensor] = None,
843
+ output_attentions: Optional[bool] = None,
844
+ output_hidden_states: Optional[bool] = None,
845
+ return_dict: Optional[bool] = None,
846
+ **kwargs,
847
+ ) -> Union[Tuple, CausalLMOutputWithPast]:
848
+ """
849
+ Forward pass for training and generation steps.
850
+
851
+ Args:
852
+ input_ids: ``(B, S)`` token IDs.
853
+ pixel_values: ``(total_tiles, C, H, W)`` image tiles from C-RADIOv4-H.
854
+ pixel_values_videos: ``(total_frames, C, H, W)`` video frames.
855
+ image_flags: ``(B, max_tiles)`` — 1 for real tiles, 0 for padding.
856
+ image_grid_thw: ``(num_images, 3)`` — ``(T=1, tile_rows, tile_cols)``
857
+ per image. Required for correct M-RoPE spatial positions.
858
+ video_grid_thw: ``(num_videos, 3)`` — ``(num_frames, 1, 1)`` per video.
859
+ mm_token_type_ids: ``(B, S)`` modality labels —
860
+ 0 = text, 1 = image, 2 = video. Required for computing
861
+ 3D M-RoPE position IDs. Produced by the processor.
862
+ attention_mask: ``(B, S)`` binary mask. Must be 2-D; the
863
+ ``Qwen3_5TextModel`` internally creates the 4-D causal mask
864
+ for full-attention layers and the 2-D mask for DeltaNet layers.
865
+ position_ids: ``(3, B, S)`` or ``None``. If ``None`` and vision
866
+ tokens are present, computed via ``get_rope_index()``.
867
+ """
868
+ return_dict = (
869
+ return_dict if return_dict is not None
870
+ else self.config.use_return_dict
871
+ )
872
+
873
+ # ── Embed tokens ─────────────────────────────────────────────────────
874
+ if inputs_embeds is None:
875
+ inputs_embeds = self.get_input_embeddings()(input_ids)
876
+
877
+ # ── Inject image features ────────────────────────────────────────────
878
+ if pixel_values is not None:
879
+ if image_flags is None:
880
+ image_flags = torch.ones(
881
+ pixel_values.shape[0], dtype=torch.long,
882
+ device=pixel_values.device,
883
+ )
884
+ image_flags_sq = image_flags.squeeze(-1)
885
+ vit_embeds = self.extract_feature(pixel_values)
886
+ vit_embeds = vit_embeds[image_flags_sq == 1]
887
+ del pixel_values
888
+
889
+ B, N, C = inputs_embeds.shape
890
+ flat = inputs_embeds.reshape(B * N, C)
891
+ ids_flat = input_ids.reshape(B * N)
892
+ selected = ids_flat == self.img_context_token_id
893
+
894
+ try:
895
+ flat[selected] = flat[selected] * 0.0 + vit_embeds.reshape(-1, C)
896
+ except Exception as e:
897
+ vit_flat = vit_embeds.reshape(-1, C)
898
+ logger.warning(
899
+ f"Image injection shape mismatch: {e}. "
900
+ f"selected={selected.sum()}, vit={vit_flat.shape}"
901
+ )
902
+ n_tok = selected.sum()
903
+ flat[selected] = flat[selected] * 0.0 + vit_flat[:n_tok]
904
+ del vit_embeds
905
+ inputs_embeds = flat.reshape(B, N, C)
906
+
907
+ # ── Inject video features ────────────────────────────────────────────
908
+ if pixel_values_videos is not None:
909
+ video_vit = self.extract_feature(pixel_values_videos)
910
+ del pixel_values_videos
911
+
912
+ B, N, C = inputs_embeds.shape
913
+ flat = inputs_embeds.reshape(B * N, C)
914
+ ids_flat = input_ids.reshape(B * N)
915
+ vmask = ids_flat == self.video_context_token_id
916
+
917
+ flat[vmask] = (
918
+ flat[vmask] * 0.0
919
+ + video_vit.reshape(-1, C).to(flat.device, flat.dtype)
920
+ )
921
+ inputs_embeds = flat.reshape(B, N, C)
922
+
923
+ del video_vit
924
+
925
+ # GRPO actor/ref training and log-prob computation should not use cache.
926
+ if labels is not None:
927
+ use_cache = False
928
+ self.rope_deltas = None
929
+
930
+ # ── 3D position IDs ──────────────────────────────────────────────────
931
+ if position_ids is None:
932
+ position_ids = self._compute_position_ids(
933
+ input_ids=input_ids,
934
+ inputs_embeds=inputs_embeds,
935
+ image_grid_thw=image_grid_thw,
936
+ video_grid_thw=video_grid_thw,
937
+ attention_mask=attention_mask,
938
+ past_key_values=past_key_values,
939
+ mm_token_type_ids=mm_token_type_ids,
940
+ use_cache=use_cache,
941
+ )
942
+
943
+ if position_ids is not None:
944
+ position_ids = position_ids.contiguous()
945
+
946
+ # ── LLM forward ─────────────────────────────────────────────────────
947
+ outputs = self.language_model(
948
+ input_ids=None,
949
+ inputs_embeds=inputs_embeds,
950
+ attention_mask=attention_mask,
951
+ position_ids=position_ids,
952
+ past_key_values=past_key_values,
953
+ use_cache=use_cache,
954
+ cache_position=cache_position,
955
+ output_attentions=output_attentions,
956
+ output_hidden_states=output_hidden_states,
957
+ return_dict=return_dict,
958
+ )
959
+ logits = outputs.logits
960
+
961
+ loss = None
962
+ if labels is not None:
963
+ shift_logits = logits[..., :-1, :].contiguous()
964
+ shift_labels = labels[..., 1:].contiguous()
965
+ loss_fct = CrossEntropyLoss()
966
+ shift_logits = shift_logits.view(
967
+ -1, self.language_model.config.vocab_size
968
+ )
969
+ shift_labels = shift_labels.view(-1).to(shift_logits.device)
970
+ loss = loss_fct(shift_logits, shift_labels)
971
+
972
+ if not return_dict:
973
+ output = (logits,) + outputs[1:]
974
+ return (loss,) + output if loss is not None else output
975
+
976
+ return CausalLMOutputWithPast(
977
+ loss=loss,
978
+ logits=logits,
979
+ past_key_values=outputs.past_key_values,
980
+ hidden_states=outputs.hidden_states,
981
+ attentions=outputs.attentions,
982
+ )
983
+
984
+ # ─────────────────────────────────────────────────────────────────────────
985
+ # GenerationMixin overrides
986
+ # ─────────────────────────────────────────────────────────────────────────
987
+
988
+ def prepare_inputs_for_generation(
989
+ self,
990
+ input_ids,
991
+ past_key_values=None,
992
+ attention_mask=None,
993
+ inputs_embeds=None,
994
+ cache_position=None,
995
+ position_ids=None,
996
+ use_cache=True,
997
+ pixel_values=None,
998
+ pixel_values_videos=None,
999
+ num_patches=None,
1000
+ image_flags=None,
1001
+ image_grid_thw=None,
1002
+ video_grid_thw=None,
1003
+ mm_token_type_ids=None,
1004
+ is_first_iteration=False,
1005
+ **kwargs,
1006
+ ):
1007
+ """
1008
+ Prepare inputs for each generation step.
1009
+
1010
+ After the first iteration, ``pixel_values`` / ``pixel_values_videos``
1011
+ are cleared because vision features are already in the KV cache.
1012
+ """
1013
+ model_inputs = super().prepare_inputs_for_generation(
1014
+ input_ids,
1015
+ past_key_values=past_key_values,
1016
+ attention_mask=attention_mask,
1017
+ inputs_embeds=inputs_embeds,
1018
+ cache_position=cache_position,
1019
+ position_ids=position_ids,
1020
+ pixel_values=pixel_values,
1021
+ pixel_values_videos=pixel_values_videos,
1022
+ num_patches=num_patches,
1023
+ image_flags=image_flags,
1024
+ image_grid_thw=image_grid_thw,
1025
+ video_grid_thw=video_grid_thw,
1026
+ mm_token_type_ids=mm_token_type_ids,
1027
+ use_cache=use_cache,
1028
+ is_first_iteration=is_first_iteration,
1029
+ **kwargs,
1030
+ )
1031
+
1032
+ if not is_first_iteration and use_cache:
1033
+ model_inputs["pixel_values"] = None
1034
+ model_inputs["pixel_values_videos"] = None
1035
+
1036
+ return model_inputs
1037
+
1038
+ def _prepare_position_ids_for_generation(self, inputs_tensor, model_kwargs):
1039
+ """
1040
+ Override to compute 3D M-RoPE position IDs during generation.
1041
+
1042
+ Mirrors ``Qwen3_5ForConditionalGeneration._prepare_position_ids_for_generation``:
1043
+ - Prefill step: compute 3D positions via ``get_rope_index``, cache
1044
+ ``rope_deltas``.
1045
+ - Decode steps: apply cached ``rope_deltas`` to sequential text positions.
1046
+
1047
+ Returns position_ids of shape ``(4, B, S)`` on the prefill step
1048
+ (text + 3D vision channels) or ``(1, B, S)`` on decode steps
1049
+ (text + rope_deltas).
1050
+ When ``Qwen3_5TextModel`` receives ``shape[0]==4``, it splits into
1051
+ ``text_position_ids = [0]`` (for causal mask) and
1052
+ ``position_ids = [1:]`` (for rotary embedding).
1053
+ When ``shape[0]!=4``, it sets ``text_position_ids=None``.
1054
+ """
1055
+ text_positions = super()._prepare_position_ids_for_generation(
1056
+ inputs_tensor, model_kwargs
1057
+ )
1058
+
1059
+ # Decode step — apply rope_deltas
1060
+ past_length = 0
1061
+ cache = model_kwargs.get("past_key_values")
1062
+ if cache is not None:
1063
+ past_length = cache.get_seq_length()
1064
+ if past_length != 0 and self.rope_deltas is not None:
1065
+ position_ids = text_positions[None, ...] + self.rope_deltas
1066
+ return position_ids
1067
+
1068
+ # Prefill step — compute 3D vision positions
1069
+ if "input_ids" in model_kwargs and model_kwargs["input_ids"].shape[1] > 0:
1070
+ inputs_tensor = model_kwargs["input_ids"]
1071
+
1072
+ is_input_ids = (
1073
+ len(inputs_tensor.shape) == 2
1074
+ and inputs_tensor.dtype in [torch.int, torch.long]
1075
+ )
1076
+ has_vision = (
1077
+ model_kwargs.get("mm_token_type_ids") is not None
1078
+ and (
1079
+ model_kwargs.get("image_grid_thw") is not None
1080
+ or model_kwargs.get("video_grid_thw") is not None
1081
+ )
1082
+ )
1083
+
1084
+ if is_input_ids and has_vision:
1085
+ vision_positions, rope_deltas = self.get_rope_index(
1086
+ inputs_tensor,
1087
+ mm_token_type_ids=model_kwargs.get("mm_token_type_ids"),
1088
+ image_grid_thw=model_kwargs.get("image_grid_thw"),
1089
+ video_grid_thw=model_kwargs.get("video_grid_thw"),
1090
+ attention_mask=model_kwargs.get("attention_mask"),
1091
+ )
1092
+ self.rope_deltas = rope_deltas
1093
+ else:
1094
+ vision_positions = text_positions.unsqueeze(0).expand(3, -1, -1)
1095
+ self.rope_deltas = torch.zeros(
1096
+ inputs_tensor.shape[0], 1,
1097
+ dtype=torch.long, device=inputs_tensor.device,
1098
+ )
1099
+
1100
+ # Concatenate text + vision → (4, B, S)
1101
+ # Channel 0 = text positions → used by create_causal_mask
1102
+ # Channels 1-3 = vision positions → used by rotary embedding
1103
+ # This matches Qwen3_5ForConditionalGeneration's convention.
1104
+ text_positions = text_positions[None, ...] # (1, B, S)
1105
+ position_ids = torch.cat(
1106
+ [text_positions, vision_positions], dim=0
1107
+ ) # (4, B, S)
1108
+ #print(f"{position_ids.permute(1, 2, 0).cpu().tolist()}")
1109
+ return position_ids
preprocessor_config.json ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "auto_map": {
3
+ "AutoImageProcessor": "image_processing.ZDTaichu5_0_ImageProcessor",
4
+ "AutoProcessor": "processing.ZDTaichu5_0_Processor"
5
+ },
6
+ "data_format": "channels_first",
7
+ "do_rescale": true,
8
+ "image_processor_type": "ZDTaichu5_0_ImageProcessor",
9
+ "image_size": 512,
10
+ "max_num_tiles": 12,
11
+ "merge_size": 1,
12
+ "norm_mean": [
13
+ 0.485,
14
+ 0.456,
15
+ 0.406
16
+ ],
17
+ "norm_std": [
18
+ 0.229,
19
+ 0.224,
20
+ 0.225
21
+ ],
22
+ "num_image_token": 256,
23
+ "rescale_factor": 0.00392156862745098,
24
+ "use_thumbnail": true
25
+ }
processing.py ADDED
@@ -0,0 +1,530 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ from typing import Optional, Union, List
15
+
16
+ import numpy as np
17
+ import torch
18
+
19
+ from transformers.feature_extraction_utils import BatchFeature
20
+ from transformers.image_utils import ImageInput
21
+ from transformers.processing_utils import ImagesKwargs, MultiModalData, ProcessingKwargs, ProcessorMixin, Unpack, VideosKwargs
22
+ from transformers.tokenization_utils_base import PreTokenizedInput, TextInput
23
+ from transformers.video_utils import VideoInput
24
+
25
+
26
+ class ZDTaichu5_0_ImagesKwargs(ImagesKwargs):
27
+ min_pixels: Optional[int]
28
+ max_pixels: Optional[int]
29
+ patch_size: Optional[int]
30
+ temporal_patch_size: Optional[int]
31
+ merge_size: Optional[int]
32
+
33
+
34
+ class ZDTaichu5_0_ProcessorKwargs(ProcessingKwargs, total=False):
35
+ images_kwargs: ZDTaichu5_0_ImagesKwargs
36
+ videos_kwargs: VideosKwargs
37
+ _defaults = {
38
+ "text_kwargs": {
39
+ "padding": False,
40
+ },
41
+ }
42
+
43
+
44
+ class ZDTaichu5_0_Processor(ProcessorMixin):
45
+ r"""
46
+ Constructs a ZDTaichu-5.0 processor which wraps an image processor and a tokenizer into a single processor.
47
+ [`ZDTaichu5_0_Processor`] offers all the functionalities of the image processor and tokenizer. See the
48
+ [`~ZDTaichu5_0_Processor.__call__`] and [`~ZDTaichu5_0_Processor.decode`] for more information.
49
+ Args:
50
+ image_processor ([`AutoImageProcessor`], *optional*):
51
+ The image processor is a required input.
52
+ tokenizer ([`AutoTokenizer`], *optional*):
53
+ The tokenizer is a required input.
54
+ chat_template (`str`, *optional*): A Jinja template which will be used to convert lists of messages
55
+ in a chat into a tokenizable string.
56
+ """
57
+
58
+ attributes = ["image_processor", "tokenizer"]
59
+
60
+ image_processor_class = "AutoImageProcessor"
61
+ video_processor_class = "AutoVideoProcessor"
62
+ tokenizer_class = ("AutoTokenizer")
63
+
64
+ def __init__(self, image_processor=None, tokenizer=None, chat_template=None, **kwargs):
65
+ # Defaults to Qwen3's built-in vision tokens; overridden by tokenizer_config.json attributes.
66
+ self.image_token = getattr(tokenizer, "image_token", "<|image_pad|>")
67
+ self.video_token = getattr(tokenizer, "video_token", "<|video_pad|>")
68
+ self.image_start_token = getattr(tokenizer, "image_start_token", "<|vision_start|>")
69
+ self.image_end_token = getattr(tokenizer, "image_end_token", "<|vision_end|>")
70
+ self.image_token_id = (
71
+ tokenizer.image_token_id
72
+ if getattr(tokenizer, "image_token_id", None)
73
+ else tokenizer.convert_tokens_to_ids(self.image_token)
74
+ )
75
+ self.video_token_id = (
76
+ tokenizer.video_token_id
77
+ if getattr(tokenizer, "video_token_id", None)
78
+ else tokenizer.convert_tokens_to_ids(self.video_token)
79
+ )
80
+ super().__init__(image_processor, tokenizer, chat_template=chat_template)
81
+
82
+ def __call__(
83
+ self,
84
+ images: ImageInput = None,
85
+ text: Union[TextInput, PreTokenizedInput, List[TextInput], List[PreTokenizedInput]] = None,
86
+ videos: VideoInput = None,
87
+ **kwargs: Unpack[ZDTaichu5_0_ProcessorKwargs],
88
+ ) -> BatchFeature:
89
+ """
90
+ Main method to prepare multimodal inputs (text, images, videos) for the model. This method processes text by
91
+ replacing image/video tokens with appropriate placeholder sequences, processes images and videos through the
92
+ image processor, and tokenizes the final text.
93
+
94
+ Video-as-multi-image convention
95
+ ───────────────────────────────
96
+ Videos are NOT processed as a separate temporal stream. Each frame is
97
+ passed through the *image* pipeline with `max_num_tiles=1`, producing
98
+ one 512x512 tile per frame, and the frame tile tensors are then
99
+ APPENDED to the image stream (`pixel_values` / `num_patches` /
100
+ `image_grid_thw`). The downstream model therefore sees a single
101
+ uniform image batch with no separate video path.
102
+
103
+ In the rendered prompt every frame is wrapped in
104
+ `<|vision_start|> ... <|image_pad|> ... <|vision_end|>` — *image*
105
+ tokens, not video tokens — and prefaced by a per-frame
106
+ "Frame N sampled at T.TT seconds:" header. After tokenisation
107
+ `mm_token_type_ids` therefore has no type=2 entries.
108
+
109
+ The method performs the following key operations:
110
+ 1. Processes images using the image processor to get pixel values and patch counts
111
+ 2. Processes videos as multi-image (max_num_tiles=1) and appends frame data
112
+ into the same pixel_values / num_patches / image_grid_thw containers
113
+ 3. Replaces `<|image_pad|>` tokens in text with `<|vision_start|>` + image tokens + `<|vision_end|>` sequences
114
+ 4. Replaces `<|video_pad|>` tokens in text with frame-by-frame descriptions including timestamps (if metadata provided)
115
+ 5. Tokenizes the processed text and combines all outputs
116
+
117
+ Args:
118
+ images (`PIL.Image.Image`, `np.ndarray`, `torch.Tensor`, `List[PIL.Image.Image]`, `List[np.ndarray]`, `List[torch.Tensor]`, *optional*):
119
+ The image or batch of images to be prepared. Each image can be a PIL image, NumPy array or PyTorch
120
+ tensor. Both channels-first and channels-last formats are supported.
121
+ text (`str`, `List[str]`, *optional*):
122
+ The sequence or batch of sequences to be encoded. Each sequence should be a string. The text can contain
123
+ special tokens `<|image_pad|>` and `<|video_pad|>` that will be replaced with appropriate token sequences.
124
+ videos (`np.ndarray`, `torch.Tensor`, `List[np.ndarray]`, `List[torch.Tensor]`, *optional*):
125
+ The video or batch of videos to be prepared. Each video should be a 4D NumPy array or PyTorch
126
+ tensor with shape (num_frames, channels, height, width). Both channels-first and channels-last formats
127
+ are supported. Note: Currently only supports batch size of 1 for videos.
128
+ images_kwargs (`Dict`, *optional*):
129
+ Additional keyword arguments for image processing, including:
130
+ - `min_pixels` (`int`, *optional*): Minimum number of pixels for image processing
131
+ - `max_pixels` (`int`, *optional*): Maximum number of pixels for image processing
132
+ - `patch_size` (`int`, *optional*): Size of patches for image processing
133
+ - `temporal_patch_size` (`int`, *optional*): Size of temporal patches
134
+ - `merge_size` (`int`, *optional*): Size for merging patches
135
+ videos_kwargs (`Dict`, *optional*):
136
+ Additional keyword arguments for video processing, including:
137
+ - `video_metadata` (`VideoMetadata`, *optional*): Metadata containing fps information for timestamp calculation
138
+ text_kwargs (`Dict`, *optional*):
139
+ Additional keyword arguments for text tokenization, including:
140
+ - `return_tensors` (`str` or [`~utils.TensorType`], *optional*): Framework for returned tensors ('tf', 'pt', 'np', 'jax')
141
+ - `padding` (`bool`, *optional*): Whether to pad sequences (defaults to False)
142
+
143
+ Returns:
144
+ [`BatchFeature`]: A [`BatchFeature`] with the following fields:
145
+
146
+ - **input_ids** -- List of token ids to be fed to a model. Returned when `text` is not `None`.
147
+ - **attention_mask** -- List of indices specifying which tokens should be attended to by the model.
148
+ - **pixel_values** -- Concatenated tile pixel values from BOTH real images and video frames,
149
+ in the order they appear in `input_ids` (real images first, then video frames). Returned
150
+ when `images` is not `None` or `videos` is not `None`.
151
+ - **num_patches** -- List of tile counts, one entry per real image followed by one entry per
152
+ video frame (frame entries are always 1 because max_num_tiles=1).
153
+ - **image_grid_thw** -- LongTensor[N_images + N_frames, 3] with [1, tile_rows, tile_cols] per
154
+ real image and [1, 1, 1] per video frame.
155
+ - **mm_token_type_ids** -- Per-token modality classification (0=text, 1=image incl. frames).
156
+
157
+ Raises:
158
+ AssertionError: If videos are provided with batch size > 1 (not currently supported).
159
+
160
+ Note:
161
+ - Image tokens `<|image_pad|>` in text are replaced with `<|vision_start|>` + repeated image tokens + `<|vision_end|>`
162
+ - Video tokens `<|video_pad|>` in text are replaced with frame-by-frame descriptions, each frame using `<|image_pad|>` slots
163
+ - When video metadata with fps is provided, frame descriptions include timestamps
164
+ - Videos are processed with max_num_tiles=1 regardless of the images setting
165
+ """
166
+ output_kwargs = self._merge_kwargs(
167
+ ZDTaichu5_0_ProcessorKwargs,
168
+ tokenizer_init_kwargs=self.tokenizer.init_kwargs,
169
+ **kwargs,
170
+ )
171
+ # Initialise as independent dicts so later `**image_inputs` merging
172
+ # is well-defined whether or not images / videos are provided.
173
+ image_inputs: dict = {}
174
+ image_grid_thw = None
175
+ # Frame counts default to empty so the video-text-expansion loop is a
176
+ # no-op when `videos` is None.
177
+ video_num_patches: list = []
178
+
179
+ if images is not None:
180
+ image_inputs = self.image_processor(images=images, **output_kwargs["images_kwargs"])
181
+ image_num_patches = image_inputs["num_patches"]
182
+ # image_grid_thw: list of [T=1, tile_rows, tile_cols] per image
183
+ image_grid_thw = image_inputs.pop("image_grid_thw")
184
+ image_pixel_values = image_inputs["pixel_values"]
185
+ else:
186
+ image_num_patches = []
187
+
188
+ if videos is not None:
189
+ # ── Multi-image treatment of video ─────────────────────────────────
190
+ # Every video frame is processed by the *image* pipeline with
191
+ # max_num_tiles=1 so that one frame = one 512x512 tile = num_image_token
192
+ # (e.g. 256) tokens. Frame tile tensors are then APPENDED to the
193
+ # real-image stream:
194
+ #
195
+ # pixel_values : torch.cat([images, frames]) (total_tiles, C, H, W)
196
+ # num_patches : image_num_patches + [1] * N_frames List[int]
197
+ # image_grid_thw : torch.cat([image_grids, frame_grids], dim=0)
198
+ #
199
+ # Order matters: text expansion below replaces image tokens
200
+ # before video tokens, so frame slots come *after* real-image
201
+ # slots in input_ids — these tensors must follow the same order.
202
+ #
203
+ # In the rendered prompt every frame is wrapped in
204
+ # <|vision_start|> ... <|image_pad|> x num_image_token ... <|vision_end|>
205
+ # (image tokens, NOT video tokens) and prefaced by a per-frame
206
+ # "Frame N sampled at T.TT seconds:" header. After tokenisation
207
+ # mm_token_type_ids therefore has *no* type=2 entries — every
208
+ # visual slot is type=1. The downstream model sees a single
209
+ # uniform image stream and does not need a separate video path.
210
+ orig_tiles = self.image_processor.max_num_tiles
211
+ self.image_processor.max_num_tiles = 1
212
+ try:
213
+ frame_inputs = self.image_processor(
214
+ images=videos, **output_kwargs["images_kwargs"]
215
+ )
216
+ finally:
217
+ self.image_processor.max_num_tiles = orig_tiles
218
+
219
+ frame_pixel_values = frame_inputs["pixel_values"] # (N_frames, C, H, W)
220
+ frame_num_patches = list(frame_inputs["num_patches"])
221
+ frame_grid_thw = frame_inputs["image_grid_thw"] # (N_frames, 3) list/tensor
222
+ video_num_patches = frame_num_patches # for text expansion below
223
+
224
+ # Normalise grid containers to LongTensor so torch.cat works
225
+ # whether the image processor returned lists or tensors.
226
+ def _to_long_tensor(x):
227
+ return x if isinstance(x, torch.Tensor) else torch.tensor(x, dtype=torch.long)
228
+
229
+ if image_inputs:
230
+ # Real images + video frames — concat along batch dim.
231
+ image_inputs["pixel_values"] = torch.cat(
232
+ [image_inputs["pixel_values"], frame_pixel_values], dim=0
233
+ )
234
+ image_inputs["num_patches"] = (
235
+ list(image_inputs["num_patches"]) + frame_num_patches
236
+ )
237
+ image_grid_thw = torch.cat(
238
+ [_to_long_tensor(image_grid_thw), _to_long_tensor(frame_grid_thw)],
239
+ dim=0,
240
+ )
241
+ image_num_patches = image_inputs["num_patches"]
242
+ else:
243
+ # Video-only — frames become the entire image stream.
244
+ image_inputs = {
245
+ "pixel_values": frame_pixel_values,
246
+ "num_patches": frame_num_patches,
247
+ }
248
+ image_grid_thw = _to_long_tensor(frame_grid_thw)
249
+ image_num_patches = frame_num_patches
250
+
251
+ if not isinstance(text, list):
252
+ text = [text]
253
+ final_image_pixel_values = []
254
+ final_image_num_patches = []
255
+ final_image_grid_thw = []
256
+ text = text.copy() # below lines change text in-place
257
+ if images is not None:
258
+ index = 0
259
+ wrapped_token = self.image_start_token + self.image_token + self.image_end_token
260
+
261
+ for i in range(len(text)):
262
+ while self.image_token in text[i]:
263
+ expansion = (
264
+ self.image_start_token
265
+ + "<|placeholder|>" * image_num_patches[index] * self.image_processor.num_image_token
266
+ + self.image_end_token
267
+ )
268
+ # If the chat template already wrapped it, replace the whole
269
+ # <vision_start><image_pad><vision_end> span — avoids double wrapping.
270
+ # Otherwise fall back to replacing the bare <image_pad> token.
271
+ search = wrapped_token if wrapped_token in text[i] else self.image_token
272
+ text[i] = text[i].replace(search, expansion, 1)
273
+ index += 1
274
+ #final_image_pixel_values.append(image_pixel_values[index])
275
+ #final_image_num_patches.append(i)
276
+ text[i] = text[i].replace("<|placeholder|>", self.image_token)
277
+ if videos is not None:
278
+ assert len(text) == 1, "Video is not supported for batch size > 1"
279
+ video_metadata = output_kwargs.get("videos_kwargs", {}).get("video_metadata", None)
280
+ i = 0
281
+ wrapped_token = self.image_start_token + self.video_token + self.image_end_token
282
+ if self.video_token in text[i]:
283
+ each_frame = (
284
+ self.image_start_token
285
+ + "<|placeholder|>" * self.image_processor.num_image_token
286
+ + self.image_end_token
287
+ )
288
+ video_prompt = "This is a video:\n"
289
+ # One iteration per frame. video_num_patches has length N_frames
290
+ # (always 1 per frame because max_num_tiles=1 was forced above),
291
+ # so its length is the authoritative frame count even when
292
+ # `images` is None and `image_num_patches` is unset.
293
+ n_frames = len(video_num_patches)
294
+ for j in range(n_frames):
295
+ if video_metadata is not None and video_metadata.fps is not None:
296
+ timestamp = j / video_metadata.fps
297
+ video_prompt += f"Frame {j+1} sampled at {timestamp:.2f} seconds: {each_frame}\n"
298
+ else:
299
+ # Fallback to original format without timestamps
300
+ video_prompt += f"Frame {j+1}: {each_frame}\n"
301
+ # Strip the chat-template-applied <|vision_start|>...<|vision_end|>
302
+ # wrapping if present; otherwise replace the bare <|video_pad|>
303
+ # token. The fallback is video_token (NOT image_token), since by
304
+ # this point image expansion has already consumed every
305
+ # <|image_pad|> in the prompt.
306
+ search = wrapped_token if wrapped_token in text[i] else self.video_token
307
+ text[i] = text[i].replace(search, video_prompt, 1)
308
+ text[i] = text[i].replace("<|placeholder|>", self.image_token)
309
+
310
+ return_tensors = output_kwargs["text_kwargs"].pop("return_tensors", None)
311
+ text_inputs = self.tokenizer(text, **output_kwargs["text_kwargs"])
312
+
313
+ # ── Build mm_token_type_ids from tokenized input_ids ─────────────
314
+ # 0 = text, 1 = image. type=2 (video) is unreachable under the
315
+ # multi-image-as-video convention because every frame gets expanded
316
+ # to <|image_pad|> tokens — but we keep the video_token branch as a
317
+ # defensive fallback for any tokenizer-injected video_pad token.
318
+ input_ids = text_inputs["input_ids"]
319
+ if isinstance(input_ids, list):
320
+ mm_token_type_ids = []
321
+ for ids in input_ids:
322
+ tt = [0] * len(ids)
323
+ for j, tok_id in enumerate(ids):
324
+ if tok_id == self.image_token_id:
325
+ tt[j] = 1
326
+ elif tok_id == self.video_token_id:
327
+ tt[j] = 2
328
+ mm_token_type_ids.append(tt)
329
+ else:
330
+ # Already a tensor (when return_tensors is set before tokenizer call)
331
+ mm_token_type_ids = torch.zeros_like(input_ids)
332
+ mm_token_type_ids[input_ids == self.image_token_id] = 1
333
+ mm_token_type_ids[input_ids == self.video_token_id] = 2
334
+
335
+ # ── Assemble output ──────────────────────────────────────────────
336
+ # Note: video frames have already been merged into image_inputs above,
337
+ # so there are no separate `pixel_values_videos` / `video_grid_thw`
338
+ # outputs. Downstream code consumes a single image stream.
339
+ data = {**text_inputs, **image_inputs}
340
+ data["mm_token_type_ids"] = mm_token_type_ids
341
+ if image_grid_thw is not None:
342
+ data["image_grid_thw"] = image_grid_thw
343
+
344
+ return BatchFeature(data=data, tensor_type=return_tensors)
345
+
346
+ def _get_num_multimodal_tokens(self, image_sizes=None, video_sizes=None, **kwargs):
347
+ """
348
+ Computes the number of placeholder tokens needed for multimodal inputs with the given sizes.
349
+ Args:
350
+ image_sizes (`list[list[int]]`, *optional*):
351
+ The input sizes formatted as (height, width) per each image.
352
+ video_sizes (`list[list[int]]`, *optional*):
353
+ The input sizes formatted as (num_frames, height, width) per each video.
354
+ Returns:
355
+ `MultiModalData`: A `MultiModalData` object holding number of tokens per each of the provided
356
+ input modalities, along with other useful data.
357
+ """
358
+
359
+ vision_data = {}
360
+ if image_sizes is not None:
361
+ images_kwargs = ZDTaichu5_0_ProcessorKwargs._defaults.get("images_kwargs", {})
362
+ images_kwargs.update(kwargs)
363
+ merge_size = images_kwargs.get("merge_size", None) or self.image_processor.merge_size
364
+
365
+ num_image_patches = [
366
+ self.image_processor.get_number_of_image_patches(*image_size, images_kwargs)
367
+ for image_size in image_sizes
368
+ ]
369
+ num_image_tokens = [(num_patches // merge_size**2) for num_patches in num_image_patches]
370
+ vision_data.update({"num_image_tokens": num_image_tokens, "num_image_patches": num_image_patches})
371
+ return MultiModalData(**vision_data)
372
+
373
+ def batch_decode(self, *args, **kwargs):
374
+ """
375
+ This method forwards all its arguments to the tokenizer's [`~PreTrainedTokenizer.batch_decode`]. Please
376
+ refer to the docstring of this method for more information.
377
+ """
378
+ return self.tokenizer.batch_decode(*args, **kwargs)
379
+
380
+ def decode(self, *args, **kwargs):
381
+ """
382
+ This method forwards all its arguments to the tokenizer's [`~PreTrainedTokenizer.decode`]. Please refer to
383
+ the docstring of this method for more information.
384
+ """
385
+ return self.tokenizer.decode(*args, **kwargs)
386
+
387
+ def post_process_image_text_to_text(
388
+ self, generated_outputs, skip_special_tokens=True, clean_up_tokenization_spaces=False, **kwargs
389
+ ):
390
+ """
391
+ Post-process the output of the model to decode the text.
392
+
393
+ Args:
394
+ generated_outputs (`torch.Tensor` or `np.ndarray`):
395
+ The output of the model `generate` function. The output is expected to be a tensor of shape `(batch_size, sequence_length)`
396
+ or `(sequence_length,)`.
397
+ skip_special_tokens (`bool`, *optional*, defaults to `True`):
398
+ Whether or not to remove special tokens in the output. Argument passed to the tokenizer's `batch_decode` method.
399
+ clean_up_tokenization_spaces (`bool`, *optional*, defaults to `False`):
400
+ Whether or not to clean up the tokenization spaces. Argument passed to the tokenizer's `batch_decode` method.
401
+ **kwargs:
402
+ Additional arguments to be passed to the tokenizer's `batch_decode method`.
403
+
404
+ Returns:
405
+ `list[str]`: The decoded text.
406
+ """
407
+ return self.tokenizer.batch_decode(
408
+ generated_outputs,
409
+ skip_special_tokens=skip_special_tokens,
410
+ clean_up_tokenization_spaces=clean_up_tokenization_spaces,
411
+ **kwargs,
412
+ )
413
+
414
+ @property
415
+ def model_input_names(self):
416
+ tokenizer_input_names = self.tokenizer.model_input_names
417
+ image_processor_input_names = self.image_processor.model_input_names
418
+ names_from_processor = list(dict.fromkeys(tokenizer_input_names + image_processor_input_names))
419
+ # Note: video_grid_thw is NOT emitted under multi-image-as-video —
420
+ # frame data lives in image_grid_thw alongside real images.
421
+ return names_from_processor + ["mm_token_type_ids"]
422
+
423
+
424
+ def from_messages(
425
+ self,
426
+ messages: list,
427
+ return_tensors: str = "pt",
428
+ add_vision_id: bool = True,
429
+ **kwargs,
430
+ ) -> BatchFeature:
431
+ """
432
+ Prepare model inputs directly from Qwen-style structured messages.
433
+
434
+ This is the high-level entry point that handles the full pipeline:
435
+ structured messages → vision loading → chat template → tokenization.
436
+
437
+ Supports messages with typed content lists::
438
+
439
+ messages = [
440
+ {"role": "user", "content": [
441
+ {"type": "image", "image": "photo.jpg"},
442
+ {"type": "text", "text": "What's in this image?"},
443
+ ]},
444
+ ]
445
+
446
+ Video inputs (file path, URL, or list of frame paths)::
447
+
448
+ messages = [
449
+ {"role": "user", "content": [
450
+ {"type": "video", "video": "clip.mp4", "fps": 2.0},
451
+ {"type": "text", "text": "Describe this video."},
452
+ ]},
453
+ ]
454
+
455
+ Multi-image with automatic labelling::
456
+
457
+ messages = [
458
+ {"role": "user", "content": [
459
+ {"type": "image", "image": "a.jpg"},
460
+ {"type": "image", "image": "b.jpg"},
461
+ {"type": "text", "text": "Compare them."},
462
+ ]},
463
+ ]
464
+ # With add_vision_id=True (default), the prompt includes:
465
+ # Picture 1: <|vision_start|><|image_pad|><|vision_end|>
466
+ # Picture 2: <|vision_start|><|image_pad|><|vision_end|>
467
+ # Compare them.
468
+
469
+ Args:
470
+ messages: List of message dicts with structured ``content``.
471
+ return_tensors: Framework for returned tensors (default ``"pt"``).
472
+ add_vision_id: If ``True`` (default), the chat template prepends
473
+ ``Picture N:`` / ``Video N:`` labels before each vision token.
474
+ Set to ``False`` to omit labels.
475
+ **kwargs: Forwarded to ``self.__call__``.
476
+
477
+ Returns:
478
+ ``BatchFeature`` ready for ``model.generate(**inputs)``.
479
+ """
480
+ from .vision_utils import process_vision_info
481
+
482
+ # 1. Load images and videos from the structured messages
483
+ image_inputs, video_inputs, video_kwargs = process_vision_info(messages)
484
+
485
+ # 2. Apply chat template — pass structured messages directly so the
486
+ # template can iterate typed content dicts, count vision elements,
487
+ # and emit "Picture N:" / "Video N:" labels when add_vision_id=True.
488
+ prompt = self.tokenizer.apply_chat_template(
489
+ messages,
490
+ tokenize=False,
491
+ add_generation_prompt=True,
492
+ add_vision_id=add_vision_id,
493
+ )
494
+
495
+ # 3. Prepare video inputs and metadata for timestamps
496
+ videos_kwargs = {}
497
+ flat_videos = None
498
+ if video_inputs is not None:
499
+ if len(video_inputs) > 1:
500
+ raise ValueError(
501
+ "Multiple videos in a single message batch are not yet "
502
+ "supported. Please use one video per call."
503
+ )
504
+ flat_videos = video_inputs[0] # List[Image.Image]
505
+
506
+ if video_kwargs.get("metadata_list"):
507
+ meta = video_kwargs["metadata_list"][0]
508
+ fps = meta.get("sample_fps") or meta.get("fps")
509
+ if fps:
510
+ from transformers.video_utils import VideoMetadata
511
+ videos_kwargs["video_metadata"] = VideoMetadata(
512
+ fps=fps,
513
+ total_num_frames=len(flat_videos),
514
+ )
515
+
516
+ return self(
517
+ images=image_inputs,
518
+ text=prompt,
519
+ videos=flat_videos,
520
+ return_tensors=return_tensors,
521
+ videos_kwargs=videos_kwargs,
522
+ **kwargs,
523
+ )
524
+
525
+ @property
526
+ def ctx_image_token_id(self) -> int:
527
+ return self.image_token_id
528
+
529
+
530
+ __all__ = ["ZDTaichu5_0_Processor"]
processor_config.json ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "auto_map": {
3
+ "AutoProcessor": "processing.ZDTaichu5_0_Processor"
4
+ },
5
+ "image_processor": {
6
+ "auto_map": {
7
+ "AutoImageProcessor": "image_processing.ZDTaichu5_0_ImageProcessor",
8
+ "AutoProcessor": "processing.ZDTaichu5_0_Processor"
9
+ },
10
+ "data_format": "channels_first",
11
+ "do_rescale": true,
12
+ "image_processor_type": "ZDTaichu5_0_ImageProcessor",
13
+ "image_size": 512,
14
+ "max_num_tiles": 12,
15
+ "merge_size": 1,
16
+ "norm_mean": [
17
+ 0.485,
18
+ 0.456,
19
+ 0.406
20
+ ],
21
+ "norm_std": [
22
+ 0.229,
23
+ 0.224,
24
+ 0.225
25
+ ],
26
+ "num_image_token": 256,
27
+ "rescale_factor": 0.00392156862745098,
28
+ "use_thumbnail": true
29
+ },
30
+ "processor_class": "ZDTaichu5_0_Processor"
31
+ }
recipe.yaml ADDED
@@ -0,0 +1,94 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ default_stage:
2
+ default_modifiers:
3
+ IMatrixGatherer:
4
+ targets: ['re:.*layers\.(?:0|1|2|3|4|5|6|7|8|9|10|11|12|13|14|15|16|17|18|19|20|21|22|23|24|25|26|27)\.mlp\.(?:gate|up|down)_proj$']
5
+ ignore: ['re:.*mtp.*', 're:.*visual.*', 're:.*vision.*', 're:.*linear_attn\.(in_proj_a|in_proj_b)$']
6
+ weight_observer: imatrix_mse
7
+ QuantizationModifier:
8
+ config_groups:
9
+ fp8_group:
10
+ targets: ['re:.*self_attn\.(q|k|v|o)_proj$', 're:.*linear_attn\.(in_proj_qkv|in_proj_z|out_proj)$',
11
+ 're:.*lm_head$', 're:.*layers\.(?:28|29|30|31)\.mlp\.(?:gate|up|down)_proj$']
12
+ weights:
13
+ num_bits: 8
14
+ type: float
15
+ symmetric: true
16
+ group_size: null
17
+ strategy: channel
18
+ block_structure: null
19
+ dynamic: false
20
+ actorder: null
21
+ scale_dtype: null
22
+ zp_dtype: null
23
+ observer: memoryless_minmax
24
+ observer_kwargs: {}
25
+ input_activations:
26
+ num_bits: 8
27
+ type: float
28
+ symmetric: true
29
+ group_size: null
30
+ strategy: token
31
+ block_structure: null
32
+ dynamic: true
33
+ actorder: null
34
+ scale_dtype: null
35
+ zp_dtype: null
36
+ observer: null
37
+ observer_kwargs: {}
38
+ output_activations: null
39
+ format: null
40
+ targets: [Linear]
41
+ ignore: ['re:.*mtp.*', 're:.*visual.*', 're:.*vision.*', 're:.*linear_attn\.(in_proj_a|in_proj_b)$']
42
+ kv_cache_scheme:
43
+ num_bits: 8
44
+ type: float
45
+ symmetric: true
46
+ group_size: null
47
+ strategy: tensor
48
+ block_structure: null
49
+ dynamic: false
50
+ actorder: null
51
+ scale_dtype: null
52
+ zp_dtype: null
53
+ observer: static_minmax
54
+ observer_kwargs: {}
55
+ bypass_divisibility_checks: false
56
+ GPTQModifier:
57
+ config_groups:
58
+ nvfp4_group:
59
+ targets: ['re:.*layers\.(?:0|1|2|3|4|5|6|7|8|9|10|11|12|13|14|15|16|17|18|19|20|21|22|23|24|25|26|27)\.mlp\.(?:gate|up|down)_proj$']
60
+ weights:
61
+ num_bits: 4
62
+ type: float
63
+ symmetric: true
64
+ group_size: 16
65
+ strategy: tensor_group
66
+ block_structure: null
67
+ dynamic: false
68
+ actorder: static
69
+ scale_dtype: torch.float8_e4m3fn
70
+ zp_dtype: null
71
+ observer: imatrix_mse
72
+ observer_kwargs: {}
73
+ input_activations:
74
+ num_bits: 4
75
+ type: float
76
+ symmetric: true
77
+ group_size: 16
78
+ strategy: tensor_group
79
+ block_structure: null
80
+ dynamic: local
81
+ actorder: null
82
+ scale_dtype: torch.float8_e4m3fn
83
+ zp_dtype: null
84
+ observer: static_minmax
85
+ observer_kwargs: {}
86
+ output_activations: null
87
+ format: null
88
+ targets: [Linear]
89
+ ignore: ['re:.*mtp.*', 're:.*visual.*', 're:.*vision.*', 're:.*linear_attn\.(in_proj_a|in_proj_b)$']
90
+ bypass_divisibility_checks: false
91
+ block_size: 128
92
+ dampening_frac: 0.01
93
+ actorder: static
94
+ offload_hessians: false
tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87a7830d63fcf43bf241c3c5242e96e62dd3fdc29224ca26fed8ea333db72de4
3
+ size 19989343
tokenizer_config.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "auto_map": {
4
+ "AutoProcessor": "processing.ZDTaichu5_0_Processor"
5
+ },
6
+ "backend": "tokenizers",
7
+ "bos_token": null,
8
+ "clean_up_tokenization_spaces": false,
9
+ "eos_token": "<|im_end|>",
10
+ "errors": "replace",
11
+ "image_end_token": "<|vision_end|>",
12
+ "image_start_token": "<|vision_start|>",
13
+ "image_token": "<|image_pad|>",
14
+ "image_token_id": 248056,
15
+ "is_local": true,
16
+ "model_max_length": 262144,
17
+ "pad_token": "<|endoftext|>",
18
+ "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+",
19
+ "processor_class": "ZDTaichu5_0_Processor",
20
+ "split_special_tokens": false,
21
+ "tokenizer_class": "TokenizersBackend",
22
+ "unk_token": null,
23
+ "video_token": "<|video_pad|>",
24
+ "video_token_id": 248057,
25
+ "vision_bos_token": "<|vision_start|>",
26
+ "vision_eos_token": "<|vision_end|>"
27
+ }
vision_utils.py ADDED
@@ -0,0 +1,583 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ # ============================================================================
16
+ # Vision utilities for ZDTaichu-5.0
17
+ #
18
+ # Provides ``process_vision_info()`` to extract images and videos from
19
+ # Qwen-style structured messages, following the conventions established
20
+ # by ``qwen_vl_utils``. This allows the model to accept messages like:
21
+ #
22
+ # messages = [
23
+ # {"role": "user", "content": [
24
+ # {"type": "video", "video": "path/to/video.mp4", "fps": 2.0},
25
+ # {"type": "text", "text": "Describe this video."},
26
+ # ]}
27
+ # ]
28
+ #
29
+ # Supported input formats:
30
+ # - Images: local path, ``file://`` URI, ``http(s)://`` URL, base64 data
31
+ # URI, ``PIL.Image.Image`` object
32
+ # - Videos: local path, ``file://`` URI, ``http(s)://`` URL (string),
33
+ # or a list of image paths/URLs (treated as pre-extracted frames)
34
+ #
35
+ # Video decoding backends (auto-detected, in priority order):
36
+ # 1. decord — fastest, recommended
37
+ # 2. torchvision — fallback, always available
38
+ #
39
+ # Frame sampling follows the same ``smart_nframes`` logic as qwen_vl_utils:
40
+ # - Default: 2 FPS, clamped to [4, 768] frames, rounded to factor of 2
41
+ # - Override via ``fps``, ``nframes``, ``min_frames``, ``max_frames``
42
+ # - Temporal trimming via ``video_start`` / ``video_end`` (seconds)
43
+ # ============================================================================
44
+
45
+ import base64
46
+ import copy
47
+ import logging
48
+ import math
49
+ import os
50
+ import sys
51
+ import time
52
+ import warnings
53
+ from functools import lru_cache
54
+ from io import BytesIO
55
+ from typing import Any, Dict, List, Optional, Tuple, Union
56
+
57
+ import numpy as np
58
+ import requests
59
+ import torch
60
+ from PIL import Image
61
+
62
+ logger = logging.getLogger(__name__)
63
+
64
+ # ─────────────────────────────────────────────────────────────────────────────
65
+ # Constants (aligned with qwen_vl_utils defaults)
66
+ # ─────────────────────────────────────────────────────────────────────────────
67
+
68
+ FPS = 2.0 # default sampling rate
69
+ FRAME_FACTOR = 2 # frame count must be divisible by this
70
+ FPS_MIN_FRAMES = 4 # minimum sampled frames
71
+ FPS_MAX_FRAMES = 768 # maximum sampled frames
72
+
73
+
74
+ # ─────────────────────────────────────────────────────────────────────────────
75
+ # Rounding helpers
76
+ # ─────────────────────────────────────────────────────────────────────────────
77
+
78
+ def round_by_factor(number: float, factor: int) -> int:
79
+ """Closest integer to *number* divisible by *factor*."""
80
+ return round(number / factor) * factor
81
+
82
+
83
+ def ceil_by_factor(number: float, factor: int) -> int:
84
+ """Smallest integer ≥ *number* divisible by *factor*."""
85
+ return math.ceil(number / factor) * factor
86
+
87
+
88
+ def floor_by_factor(number: float, factor: int) -> int:
89
+ """Largest integer ≤ *number* divisible by *factor*."""
90
+ return math.floor(number / factor) * factor
91
+
92
+
93
+ # ─────────────────────────────────────────────────────────────────────────────
94
+ # Image loading
95
+ # ─────────────────────────────────────────────────────────────────────────────
96
+
97
+ def fetch_image(ele: Dict[str, Any]) -> Image.Image:
98
+ """
99
+ Load a single image from various sources.
100
+
101
+ Supported formats for ``ele["image"]``:
102
+ - ``PIL.Image.Image`` instance
103
+ - Local file path (``/path/to/img.jpg``)
104
+ - ``file://`` URI
105
+ - ``http://`` or ``https://`` URL
106
+ - Base64 data URI (``data:image/...;base64,...``)
107
+
108
+ Returns:
109
+ PIL.Image.Image in RGB mode.
110
+ """
111
+ image = ele.get("image") or ele.get("image_url")
112
+ if image is None:
113
+ raise ValueError("Element must contain 'image' or 'image_url' key")
114
+
115
+ image_obj = None
116
+ if isinstance(image, Image.Image):
117
+ image_obj = image
118
+ elif image.startswith("http://") or image.startswith("https://"):
119
+ with requests.get(image, stream=True, timeout=30) as resp:
120
+ resp.raise_for_status()
121
+ image_obj = copy.deepcopy(Image.open(BytesIO(resp.content)))
122
+ elif image.startswith("file://"):
123
+ image_obj = Image.open(image[7:])
124
+ elif image.startswith("data:image"):
125
+ if "base64," in image:
126
+ _, b64 = image.split("base64,", 1)
127
+ image_obj = copy.deepcopy(Image.open(BytesIO(base64.b64decode(b64))))
128
+ else:
129
+ # Treat as local file path
130
+ image_obj = Image.open(image)
131
+
132
+ if image_obj is None:
133
+ raise ValueError(
134
+ f"Unrecognised image input. Supported: local path, file:// URI, "
135
+ f"http(s) URL, base64 data URI, PIL.Image. Got: {image!r:.120}"
136
+ )
137
+
138
+ # Convert to RGB
139
+ if image_obj.mode == "RGBA":
140
+ bg = Image.new("RGB", image_obj.size, (255, 255, 255))
141
+ bg.paste(image_obj, mask=image_obj.split()[3])
142
+ return bg
143
+ return image_obj.convert("RGB")
144
+
145
+
146
+ # ─────────────────────────────────────────────────────────────────────────────
147
+ # Frame sampling
148
+ # ─────────────────────────────────────────────────────────────────────────────
149
+
150
+ def smart_nframes(
151
+ ele: Dict[str, Any],
152
+ total_frames: int,
153
+ video_fps: float,
154
+ ) -> int:
155
+ """
156
+ Compute the number of frames to sample from a video.
157
+
158
+ Follows the same logic as ``qwen_vl_utils.smart_nframes``:
159
+ - If ``ele["nframes"]`` is set, use it directly (rounded to FRAME_FACTOR).
160
+ - Otherwise, sample at ``ele.get("fps", 2.0)`` FPS, clamped to
161
+ ``[min_frames, max_frames]`` and rounded down to FRAME_FACTOR.
162
+
163
+ Args:
164
+ ele: Dict with optional keys ``fps``, ``nframes``, ``min_frames``,
165
+ ``max_frames``.
166
+ total_frames: Total frames in the (possibly trimmed) video.
167
+ video_fps: Original video FPS.
168
+
169
+ Returns:
170
+ Number of frames to sample.
171
+ """
172
+ assert not ("fps" in ele and "nframes" in ele), (
173
+ "Only accept either `fps` or `nframes`, not both"
174
+ )
175
+
176
+ if "nframes" in ele:
177
+ nframes = round_by_factor(ele["nframes"], FRAME_FACTOR)
178
+ else:
179
+ fps = ele.get("fps", FPS)
180
+ min_frames = ceil_by_factor(
181
+ ele.get("min_frames", FPS_MIN_FRAMES), FRAME_FACTOR
182
+ )
183
+ max_frames = floor_by_factor(
184
+ ele.get("max_frames", min(FPS_MAX_FRAMES, total_frames)), FRAME_FACTOR
185
+ )
186
+ nframes = total_frames / video_fps * fps
187
+ if nframes > total_frames:
188
+ logger.warning(
189
+ f"smart_nframes: computed nframes ({nframes:.1f}) > "
190
+ f"total_frames ({total_frames})"
191
+ )
192
+ nframes = min(min(max(nframes, min_frames), max_frames), total_frames)
193
+ nframes = floor_by_factor(nframes, FRAME_FACTOR)
194
+
195
+ if not (FRAME_FACTOR <= nframes <= total_frames):
196
+ raise ValueError(
197
+ f"nframes should be in [{FRAME_FACTOR}, {total_frames}], "
198
+ f"got {nframes}."
199
+ )
200
+ return nframes
201
+
202
+
203
+ def calculate_video_frame_range(
204
+ ele: Dict[str, Any],
205
+ total_frames: int,
206
+ video_fps: float,
207
+ ) -> Tuple[int, int, int]:
208
+ """
209
+ Calculate start/end frame indices from optional ``video_start``/``video_end``
210
+ keys (in seconds).
211
+
212
+ Returns:
213
+ (start_frame, end_frame, frame_count) — end_frame is inclusive.
214
+ """
215
+ if video_fps <= 0:
216
+ raise ValueError("video_fps must be positive")
217
+ if total_frames <= 0:
218
+ raise ValueError("total_frames must be positive")
219
+
220
+ video_start = ele.get("video_start")
221
+ video_end = ele.get("video_end")
222
+
223
+ if video_start is None and video_end is None:
224
+ return 0, total_frames - 1, total_frames
225
+
226
+ max_duration = total_frames / video_fps
227
+
228
+ if video_start is not None:
229
+ start_sec = max(0.0, min(video_start, max_duration))
230
+ start_frame = math.ceil(start_sec * video_fps)
231
+ else:
232
+ start_frame = 0
233
+
234
+ if video_end is not None:
235
+ end_sec = max(0.0, min(video_end, max_duration))
236
+ end_frame = min(math.floor(end_sec * video_fps), total_frames - 1)
237
+ else:
238
+ end_frame = total_frames - 1
239
+
240
+ if start_frame >= end_frame:
241
+ raise ValueError(
242
+ f"Invalid time range: start_frame={start_frame} >= end_frame={end_frame}. "
243
+ f"Video: {max_duration:.2f}s ({total_frames} frames @ {video_fps:.1f}fps)"
244
+ )
245
+
246
+ return start_frame, end_frame, end_frame - start_frame + 1
247
+
248
+
249
+ # ─────────────────────────────────────────────────────────────────────────────
250
+ # Video decoding backends
251
+ # ─────────────────────────────────────────────────────────────────────────────
252
+
253
+ def _read_video_decord(
254
+ ele: Dict[str, Any],
255
+ ) -> Tuple[torch.Tensor, dict, float]:
256
+ """Read video with decord. Returns (video_TCHW, metadata, sample_fps)."""
257
+ import decord
258
+
259
+ video_path = ele["video"]
260
+ if video_path.startswith("file://"):
261
+ video_path = video_path[7:]
262
+
263
+ st = time.time()
264
+ vr = decord.VideoReader(video_path)
265
+ total_frames, video_fps = len(vr), vr.get_avg_fps()
266
+
267
+ start_frame, end_frame, total_frames = calculate_video_frame_range(
268
+ ele, total_frames, video_fps
269
+ )
270
+ nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
271
+ idx = torch.linspace(start_frame, end_frame, nframes).round().long().tolist()
272
+ sample_fps = nframes / max(total_frames, 1e-6) * video_fps
273
+
274
+ video = torch.from_numpy(vr.get_batch(idx).asnumpy()).permute(0, 3, 1, 2) # TCHW
275
+ logger.info(
276
+ f"decord: {video_path}, {total_frames} frames, "
277
+ f"{video_fps:.1f} fps, sampled {nframes}, "
278
+ f"time={time.time() - st:.3f}s"
279
+ )
280
+
281
+ metadata = dict(
282
+ fps=video_fps,
283
+ sample_fps=sample_fps,
284
+ frames_indices=idx,
285
+ total_num_frames=total_frames,
286
+ video_backend="decord",
287
+ )
288
+ return video, metadata, sample_fps
289
+
290
+
291
+ def _read_video_torchvision(
292
+ ele: Dict[str, Any],
293
+ ) -> Tuple[torch.Tensor, dict, float]:
294
+ """Read video with torchvision. Returns (video_TCHW, metadata, sample_fps)."""
295
+ from torchvision import io as tio
296
+
297
+ video_path = ele["video"]
298
+ if video_path.startswith("file://"):
299
+ video_path = video_path[7:]
300
+
301
+ st = time.time()
302
+ video, _audio, info = tio.read_video(
303
+ video_path,
304
+ start_pts=ele.get("video_start", 0.0),
305
+ end_pts=ele.get("video_end"),
306
+ pts_unit="sec",
307
+ output_format="TCHW",
308
+ )
309
+ total_frames, video_fps = video.size(0), info["video_fps"]
310
+
311
+ nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
312
+ idx = torch.linspace(0, total_frames - 1, nframes).round().long()
313
+ sample_fps = nframes / max(total_frames, 1e-6) * video_fps
314
+ video = video[idx]
315
+
316
+ logger.info(
317
+ f"torchvision: {video_path}, {total_frames} frames, "
318
+ f"{video_fps:.1f} fps, sampled {nframes}, "
319
+ f"time={time.time() - st:.3f}s"
320
+ )
321
+
322
+ metadata = dict(
323
+ fps=video_fps,
324
+ sample_fps=sample_fps,
325
+ frames_indices=idx.tolist(),
326
+ total_num_frames=total_frames,
327
+ video_backend="torchvision",
328
+ )
329
+ return video, metadata, sample_fps
330
+
331
+
332
+ def _is_decord_available() -> bool:
333
+ import importlib.util
334
+ return importlib.util.find_spec("decord") is not None
335
+
336
+
337
+ @lru_cache(maxsize=1)
338
+ def _get_video_backend() -> str:
339
+ forced = os.getenv("TAICHU_VIDEO_READER")
340
+ if forced is not None:
341
+ backend = forced
342
+ elif _is_decord_available():
343
+ backend = "decord"
344
+ else:
345
+ backend = "torchvision"
346
+ print(
347
+ f"ZDTaichu-5.0 utilities using {backend} to read video.",
348
+ file=sys.stderr,
349
+ )
350
+ return backend
351
+
352
+
353
+ _VIDEO_BACKENDS = {
354
+ "decord": _read_video_decord,
355
+ "torchvision": _read_video_torchvision,
356
+ }
357
+
358
+
359
+ # ─────────────────────────────────────────────────────────────────────────────
360
+ # fetch_video — main entry point for video loading
361
+ # ─────────────────────────────────────────────────────────────────────────────
362
+
363
+ def fetch_video(
364
+ ele: Dict[str, Any],
365
+ ) -> Tuple[List[Image.Image], float, dict]:
366
+ """
367
+ Load and sample frames from a video.
368
+
369
+ The ``ele["video"]`` value can be:
370
+ - A string path / URI → decoded with decord or torchvision
371
+ - A list of image paths → loaded as pre-extracted frames
372
+
373
+ Returns:
374
+ (frames, sample_fps, metadata)
375
+ - frames: list of PIL.Image.Image in RGB (one per sampled frame)
376
+ - sample_fps: effective sampling rate after frame selection
377
+ - metadata: dict with ``fps``, ``sample_fps``, ``total_num_frames``,
378
+ ``frames_indices``, ``video_backend``
379
+ """
380
+ if isinstance(ele["video"], str):
381
+ # ── Decode from video file ──────────────────────────────────���────
382
+ backend = _get_video_backend()
383
+ try:
384
+ video_tensor, metadata, sample_fps = _VIDEO_BACKENDS[backend](ele)
385
+ except Exception as exc:
386
+ if backend != "torchvision":
387
+ logger.warning(
388
+ f"{backend} failed ({exc}), falling back to torchvision"
389
+ )
390
+ video_tensor, metadata, sample_fps = _read_video_torchvision(ele)
391
+ else:
392
+ raise
393
+
394
+ # Convert TCHW tensor → list of PIL images
395
+ frames = []
396
+ for i in range(video_tensor.size(0)):
397
+ frame_np = video_tensor[i].permute(1, 2, 0).numpy().astype(np.uint8) # HWC
398
+ frames.append(Image.fromarray(frame_np, "RGB"))
399
+
400
+ elif isinstance(ele["video"], (list, tuple)):
401
+ # ── Pre-extracted frames (paths or PIL images) ───────────────────
402
+ frame_elements = ele["video"]
403
+ frames = []
404
+ for item in frame_elements:
405
+ frames.append(fetch_image({"image": item}))
406
+
407
+ # Pad to FRAME_FACTOR multiple
408
+ nframes = ceil_by_factor(len(frames), FRAME_FACTOR)
409
+ while len(frames) < nframes:
410
+ frames.append(frames[-1].copy())
411
+
412
+ sample_fps = ele.get("fps", FPS)
413
+ raw_fps = ele.get("raw_fps", sample_fps)
414
+ metadata = dict(
415
+ fps=raw_fps,
416
+ sample_fps=sample_fps,
417
+ frames_indices=list(range(len(frames))),
418
+ total_num_frames=len(frames),
419
+ video_backend="frames_list",
420
+ )
421
+ else:
422
+ raise TypeError(
423
+ f"ele['video'] must be a string (path) or list (frames), "
424
+ f"got {type(ele['video'])}"
425
+ )
426
+
427
+ return frames, sample_fps, metadata
428
+
429
+
430
+ # ─────────────────────────────────────────────────────────────────────────────
431
+ # Message parsing
432
+ # ─────────────────────────────────────────────────────────────────────────────
433
+
434
+ def extract_vision_info(
435
+ conversations: Union[List[Dict[str, Any]], List[List[Dict[str, Any]]]],
436
+ ) -> List[Dict[str, Any]]:
437
+ """
438
+ Extract all vision elements (image / video dicts) from Qwen-style
439
+ structured messages.
440
+
441
+ Args:
442
+ conversations: Either a single conversation (list of message dicts)
443
+ or a batch of conversations.
444
+
445
+ Returns:
446
+ Flat list of vision element dicts, in order of appearance.
447
+ """
448
+ # Normalise to batch format
449
+ if isinstance(conversations[0], dict):
450
+ conversations = [conversations]
451
+
452
+ vision_infos = []
453
+ for conversation in conversations:
454
+ for message in conversation:
455
+ content = message.get("content")
456
+ if not isinstance(content, list):
457
+ continue
458
+ for ele in content:
459
+ if (
460
+ "image" in ele
461
+ or "image_url" in ele
462
+ or "video" in ele
463
+ or ele.get("type") in ("image", "image_url", "video")
464
+ ):
465
+ vision_infos.append(ele)
466
+ return vision_infos
467
+
468
+
469
+ def process_vision_info(
470
+ conversations: Union[List[Dict[str, Any]], List[List[Dict[str, Any]]]],
471
+ ) -> Tuple[Optional[List[Image.Image]], Optional[List[List[Image.Image]]], Optional[Dict[str, Any]]]:
472
+ """
473
+ Extract and load all images and videos from structured messages.
474
+
475
+ This is the main entry point — equivalent to
476
+ ``qwen_vl_utils.process_vision_info`` — adapted for ZDTaichu-5.0.
477
+
478
+ Args:
479
+ conversations: Qwen-style messages with structured ``content`` lists
480
+ containing ``{"type": "image", "image": ...}`` and/or
481
+ ``{"type": "video", "video": ...}`` elements.
482
+
483
+ Returns:
484
+ (image_inputs, video_inputs, video_kwargs)
485
+ - image_inputs: list of PIL images, or None
486
+ - video_inputs: list of frame-lists (each is ``List[PIL.Image]``),
487
+ or None
488
+ - video_kwargs: dict with ``sample_fps_list`` and ``metadata_list``
489
+
490
+ Example::
491
+
492
+ from vision_utils import process_vision_info
493
+
494
+ messages = [
495
+ {"role": "user", "content": [
496
+ {"type": "video", "video": "clip.mp4", "fps": 2.0},
497
+ {"type": "text", "text": "Describe this video."},
498
+ ]}
499
+ ]
500
+
501
+ images, videos, video_kwargs = process_vision_info(messages)
502
+ # images = None
503
+ # videos = [[PIL.Image, PIL.Image, ...]] (one list of frames per video)
504
+ # video_kwargs = {"sample_fps_list": [2.0], "metadata_list": [...]}
505
+ """
506
+ vision_infos = extract_vision_info(conversations)
507
+
508
+ image_inputs: List[Image.Image] = []
509
+ video_inputs: List[List[Image.Image]] = []
510
+ sample_fps_list: List[float] = []
511
+ metadata_list: List[dict] = []
512
+
513
+ for info in vision_infos:
514
+ if "image" in info or "image_url" in info:
515
+ image_inputs.append(fetch_image(info))
516
+
517
+ elif "video" in info:
518
+ frames, sample_fps, metadata = fetch_video(info)
519
+ video_inputs.append(frames)
520
+ sample_fps_list.append(sample_fps)
521
+ metadata_list.append(metadata)
522
+
523
+ else:
524
+ raise ValueError(
525
+ "Vision element must contain 'image', 'image_url', or 'video' key."
526
+ )
527
+
528
+ video_kwargs = {
529
+ "sample_fps_list": sample_fps_list,
530
+ "metadata_list": metadata_list,
531
+ }
532
+
533
+ return (
534
+ image_inputs if image_inputs else None,
535
+ video_inputs if video_inputs else None,
536
+ video_kwargs,
537
+ )
538
+
539
+
540
+ def build_text_from_messages(
541
+ messages: List[Dict[str, Any]],
542
+ image_token: str = "<|image_pad|>",
543
+ video_token: str = "<|video_pad|>",
544
+ ) -> List[Dict[str, Any]]:
545
+ """
546
+ Convert structured messages (with typed content lists) into plain-text
547
+ messages that ``apply_chat_template`` can handle.
548
+
549
+ Each ``{"type": "image", ...}`` is replaced with ``image_token``.
550
+ Each ``{"type": "video", ...}`` is replaced with ``video_token``.
551
+ Text elements are concatenated.
552
+
553
+ Returns:
554
+ New message list with plain string ``content`` fields.
555
+ """
556
+ output = []
557
+ for msg in messages:
558
+ content = msg.get("content")
559
+ if isinstance(content, str):
560
+ output.append(msg)
561
+ continue
562
+
563
+ parts = []
564
+ for ele in content:
565
+ typ = ele.get("type", "text")
566
+ if typ == "text":
567
+ parts.append(ele.get("text", ""))
568
+ elif typ in ("image", "image_url"):
569
+ parts.append(image_token)
570
+ elif typ == "video":
571
+ parts.append(video_token)
572
+ output.append({**msg, "content": "".join(parts)})
573
+ return output
574
+
575
+
576
+ __all__ = [
577
+ "fetch_image",
578
+ "fetch_video",
579
+ "smart_nframes",
580
+ "extract_vision_info",
581
+ "process_vision_info",
582
+ "build_text_from_messages",
583
+ ]