Upload folder using huggingface_hub
Browse files- .gitattributes +1 -0
- .ms_upload_cache +1 -0
- .vscode/settings.json +15 -0
- README.md +653 -1
- __init__.py +16 -0
- assets/logos/gemini.png +0 -0
- assets/logos/gemma.png +0 -0
- assets/logos/grok.png +0 -0
- assets/logos/qwen.png +0 -0
- assets/logos/stepfun.png +0 -0
- assets/logos/taichu.png +0 -0
- assets/taichu-release-benchmark-comparison.svg +0 -0
- assets/taichu-vs-closed-models.svg +0 -0
- chat_template.jinja +154 -0
- config.json +673 -0
- configuration.py +231 -0
- cradio_config.py +54 -0
- cradio_model.py +699 -0
- generation_config.json +16 -0
- image_processing.py +268 -0
- model.safetensors +3 -0
- modeling.py +1109 -0
- preprocessor_config.json +25 -0
- processing.py +530 -0
- processor_config.json +31 -0
- recipe.yaml +94 -0
- tokenizer.json +3 -0
- tokenizer_config.json +27 -0
- vision_utils.py +583 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
.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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
+

|
| 91 |
+
|
| 92 |
+
**Comparison with closed models**
|
| 93 |
+
|
| 94 |
+

|
| 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 |
+
]
|