Publish model weights, paper figures, and complete model card

#1
by Skywalker0410 - opened
.gitattributes CHANGED
@@ -33,3 +33,13 @@ 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
37
+ assets/decoding.mp4 filter=lfs diff=lfs merge=lfs -text
38
+ assets/demo.mp4 filter=lfs diff=lfs merge=lfs -text
39
+ assets/fig1-teaser.png filter=lfs diff=lfs merge=lfs -text
40
+ assets/fig2-architecture.png filter=lfs diff=lfs merge=lfs -text
41
+ assets/fig4-attention-mask.png filter=lfs diff=lfs merge=lfs -text
42
+ assets/fig6-self-speculative-decoding.png filter=lfs diff=lfs merge=lfs -text
43
+ assets/fig7-grounding-performance.png filter=lfs diff=lfs merge=lfs -text
44
+ assets/logo.png filter=lfs diff=lfs merge=lfs -text
45
+ assets/demo-poster.jpg filter=lfs diff=lfs merge=lfs -text
LICENSE ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Kimi K3 License
2
+
3
+ Copyright (c) 2026 Moonshot AI
4
+
5
+ Permission is hereby granted, free of charge, to any person (the "Licensee")
6
+ obtaining a copy of this software — including the model weights, parameters,
7
+ configuration files, inference and training code, and associated documentation
8
+ (collectively, the "Software") — to deal in the Software without restriction.
9
+ This includes, without limitation, the rights to use, copy, modify, merge,
10
+ publish, distribute, sublicense, and/or sell copies of the Software; to run,
11
+ deploy, fine-tune, or otherwise modify the Software and create derivative works
12
+ from it; and to permit persons to whom the Software is furnished to do so, in
13
+ each case subject to the following conditions:
14
+
15
+ 1. The above copyright notice and this permission notice shall be included in
16
+ all copies or substantial portions of the Software. Licensee's use of the
17
+ Software must comply with applicable laws and regulations.
18
+
19
+ 2. "Model as a Service" means giving a third party access to language model
20
+ inference or fine-tuning (e.g., via API) in a manner that allows such third
21
+ party to exercise meaningful control over the inputs, parameters, or training
22
+ data. This does not include (a) end-user products with model capabilities solely
23
+ embedded within specific features or harnesses, or (b) mere relaying of requests
24
+ to models hosted by others.
25
+
26
+ If the Licensee or any of its affiliates operates a Model as a Service business,
27
+ and the aggregate revenue of the Licensee and its affiliates exceeds 20 million
28
+ US dollars (or the equivalent in other currencies) in total over any consecutive
29
+ 12 months, the Licensee must enter into a separate agreement with Moonshot AI
30
+ before using the Software or its derivative works for any commercial purpose.
31
+
32
+ 3. If the Software (or any derivative works thereof) is used for any of the
33
+ Licensee's commercial products or services that have more than 100 million
34
+ monthly active users, or more than 20 million US dollars (or equivalent in other
35
+ currencies) in monthly revenue, "Kimi K3" must be prominently displayed on the
36
+ user interface of such product or service.
37
+
38
+ 4. The requirements set forth in Sections 2 and 3 do not apply to: (a) internal
39
+ use of the Software, defined as any use that does not make the Software, its
40
+ outputs, or its underlying capabilities available to third parties; or (b) any
41
+ use of the Software accessed through Moonshot AI's official products or
42
+ certified inference partners.
43
+
44
+ 5. THE SOFTWARE AND ANY OUTPUT AND RESULTS THEREFROM ARE PROVIDED ON AN “AS IS”
45
+ BASIS, WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT
46
+ LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE
47
+ AND NONINFRINGEMENT. IN NO EVENT SHALL MOONSHOT AI OR ITS AFFILIATES OR
48
+ COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
49
+ IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
50
+ CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
51
+
52
+ For any questions regarding this license, please contact <license@moonshot.ai>.
README.md CHANGED
@@ -1,3 +1,381 @@
1
  ---
2
- license: apache-2.0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ license: other
3
+ license_name: kimi-k3
4
+ license_link: https://huggingface.co/moonshotai/Kimi-K3/blob/f831ab66814297da540d832a5235f8e904f29d06/LICENSE
5
+ language:
6
+ - en
7
+ - zh
8
+ library_name: transformers
9
+ pipeline_tag: image-text-to-text
10
+ tags:
11
+ - visual-grounding
12
+ - object-detection
13
+ - referring-expression-comprehension
14
+ - pointing
15
+ - ocr
16
+ - document-layout
17
+ - custom-code
18
+ - diffusion-language-model
19
+ inference: false
20
  ---
21
+
22
+ <p align="center"><img src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/logo.png" width="150" alt="GroundAnything logo" /></p>
23
+
24
+ # GroundAnything: Reconciling Parallel Decoding with Precise Visual Grounding at Flash Speed
25
+
26
+ **Model family:** [GroundAnything — DLM / parallel decoding](https://huggingface.co/GroundingPI/GroundAnything) · [GroundAnything-VLM — autoregressive](https://huggingface.co/GroundingPI/GroundAnything-VLM).
27
+
28
+ **This repository contains the GroundAnything DLM checkpoint.** Use entropy-guided decoding for the main benchmark setting, or select optional self-speculative decoding.
29
+
30
+ <p align="center"><img src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/fig1-teaser.png" width="100%" alt="GroundAnything Figure 1: broad visual grounding and parallel visual evidence extraction" /></p>
31
+
32
+ ## 🔗 Quick Links
33
+
34
+ - 🚀 **Online Demo:** Coming soon — XXX.
35
+ - 💻 **GitHub Code:** Coming soon — XXX.
36
+ - 📄 **Paper:** [arXiv:2609.39600](https://arxiv.org/abs/2609.39600).
37
+ - 🧪 **Evaluation data:** Coming soon — XXX.
38
+
39
+ # Model Overview
40
+
41
+ ### Description:
42
+
43
+ Autoregressive (AR) grounding models serialize spatial predictions, introducing sequential latency and imposing a causal order on output tokens. We view grounding as visual evidence extraction: objects, locations, and spatial relations are jointly constrained by the image and query, yet their dependencies do not imply an intrinsic left-to-right generation order. This distinction makes bidirectional diffusion a natural fit, allowing spatial hypotheses to emerge in parallel and be jointly refined through iterative denoising.
44
+
45
+ We introduce **GroundAnything**, a 4B-parameter grounding foundation model that reconciles fast parallel decoding with precise localization through blockwise denoising. Across 30 grounding benchmarks, the autoregressive variant, **GroundAnything-VLM**, establishes a new overall state of the art among similarly sized models at **72.42%**, remaining competitive with GPT-6 Astra (**71.35%**). With **entropy-guided decoding**, GroundAnything also surpasses the prior state of the art at this scale, averaging **61.75%** versus **53.32%** for the fast MTP-based LocateAnything model.
46
+
47
+ An optional **self-speculative mode** achieves a **4.51× speedup** over the AR counterpart with a **0.74 percentage-point drop in COCO F1mIoU** at the paper's reported operating point. Infrastructure experiments show that progressive inference optimizations translate parallel decoding into practical speedups. These support efficient visual grounding in latency-sensitive real-world systems.
48
+
49
+ ### Demo Videos
50
+
51
+ <video controls playsinline preload="none" width="100%" poster="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/demo-poster.jpg" src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/demo.mp4"></video>
52
+
53
+ **Parallel decoding in action**
54
+
55
+ <video controls playsinline preload="none" width="100%" src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/decoding.mp4"></video>
56
+
57
+ ### License/Terms of Use:
58
+
59
+ This checkpoint is released under the **Kimi K3 License**; see the repository's [LICENSE](LICENSE) and the [upstream license](https://huggingface.co/moonshotai/Kimi-K3/blob/f831ab66814297da540d832a5235f8e904f29d06/LICENSE). The model incorporates Kimi K3-derived vision components and implementation code. The license includes additional conditions for certain commercial uses. Third-party components retain their respective licenses and copyright notices.
60
+
61
+ ### Deployment Geography:
62
+
63
+ Global.
64
+
65
+ ### Use Case:
66
+
67
+ - Open-vocabulary object localization and dense-scene grounding.
68
+ - Referring-expression comprehension and visual-prompt-based localization.
69
+ - Point-based localization, spatial reasoning, and GUI element grounding.
70
+ - OCR with text localization and document layout understanding.
71
+ - Perception research for robotics, embodied agents, and autonomous systems.
72
+
73
+ ### Release Date:
74
+
75
+ - **Paper [09/30/2026]:** [GroundAnything](https://arxiv.org/abs/2609.39600).
76
+
77
+ ## References(s):
78
+
79
+ - [GroundAnything paper and supplementary material](https://arxiv.org/abs/2609.39600).
80
+ - [GroundingPI: grounding with visual primitives](https://arxiv.org/abs/2609.39601).
81
+
82
+ <details>
83
+ <summary>Citation</summary>
84
+
85
+ ```bibtex
86
+ @misc{yu2026groundanything,
87
+ title = {GroundAnything: Reconciling Parallel Decoding with Precise Visual Grounding at Flash Speed},
88
+ author = {Qize Yu and Lianrui Fan and Bowen Ping and Xini Ding and Zetian Song and Junbo Niu and Kaixuan Wang and Tianxing Chen and Yue Chen and Minghua He and Yuran Wang and Jie Huang and Haojun Zhang and Min Chen and Hao Li and Wenxuan Song and Ruihai Wu and Xianming Liu and Shilong Liu and Shuchang Zhou and Ping Luo and Shiyu Huang},
89
+ year = {2026},
90
+ eprint = {2609.39600},
91
+ archivePrefix = {arXiv},
92
+ primaryClass = {cs.CV},
93
+ url = {https://arxiv.org/abs/2609.39600},
94
+ }
95
+ ```
96
+
97
+ </details>
98
+
99
+ ## Model Architecture:
100
+
101
+ **Architecture type:** a shared vision-language backbone supporting an autoregressive checkpoint and a blockwise diffusion checkpoint.
102
+
103
+ - **Vision encoder:** MoonViT-V2 / Kimi-K3 vision backbone.
104
+ - **Language backbone:** Qwen3-4B-Instruct-2507.
105
+ - **Multimodal projector:** 2 × 2 spatial aggregation and a two-layer MLP.
106
+ - **Spatial vocabulary:** 1,000 coordinate tokens shared with semantic labels and protocol markers.
107
+ - **DLM conversion:** the shared decoder and vocabulary head support both causal prediction and bidirectional response-block denoising. A mask token is added for diffusion generation.
108
+
109
+ <p align="center"><img src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/fig2-architecture.png" width="100%" alt="GroundAnything Figure 2: vision-language architecture and autoregressive-to-diffusion conversion" /></p>
110
+
111
+ *Figure 2. Shared model architecture and the conversion to parallel grounding.*
112
+
113
+ <p align="center"><img src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/fig4-attention-mask.png" width="640" alt="GroundAnything Figure 4: clean-stream causal and noisy-response block attention" /></p>
114
+
115
+ *Figure 4. The attention mask used during diffusion conversion. The clean stream uses causal attention. A noisy response block attends bidirectionally within its block and reads the clean conditioning context and strictly preceding clean response blocks. Conversion uses B=32; the illustration uses two-token blocks.*
116
+
117
+ ## Input(s):
118
+
119
+ **Input types:** image and text.
120
+
121
+ - **Image:** one RGB image per request. The Python client accepts JPEG, PNG, and WebP files; the supplied processor handles image preparation.
122
+ - **Text:** a natural-language instruction, category list, referring expression, OCR/layout query, or a prompt containing example boxes.
123
+ - **Image encoding for the HTTP API:** an `image_url` content part containing a base64 data URI, alongside a `text` content part.
124
+
125
+ Use the checkpoint's own tokenizer, processor, and chat template. Multiple categories are separated by `</c>`. Reference boxes use the same 0–999 spatial-token vocabulary as outputs.
126
+
127
+ ## Output(s):
128
+
129
+ **Output type:** text containing semantic labels and quantized spatial coordinates.
130
+
131
+ All three released models use the **GAM protocol**: integer coordinate tokens `<0>` through `<999>`, object-reference delimiters, and box delimiters. A bounding box contains `(x1, y1, x2, y2)`; a point contains `(x, y)`. Adjacent coordinate tokens have no intervening spaces. Multiple instances of the same label are comma-separated inside one box wrapper. A missing target is represented by `None`.
132
+
133
+ Illustrative syntax:
134
+
135
+ ```text
136
+ <|object_ref_start|>car<|object_ref_end|><|box_start|><100><200><500><650><|box_end|>
137
+ <|object_ref_start|>car center<|object_ref_end|><|box_start|><300><425><|box_end|>
138
+ <|object_ref_start|>absent object<|object_ref_end|><|box_start|>None<|box_end|>
139
+ ```
140
+
141
+ The client returns parsed predictions together with `raw_output`, `finish_reason`, `usage`, `parse_error`, and `valid`. It maps coordinates back to image pixels for visualization. A truncated or malformed response is marked invalid. For custom API clients, preserve spatial tokens with `skip_special_tokens=false` and avoid inserting spaces between them.
142
+
143
+ ## Software Integration:
144
+
145
+ **Runtime engines:** the repository's custom **SGLang** integration for the DLM and VLM services, plus a native Transformers reference route for the DLM.
146
+
147
+ **Default serving environment:** Linux and Python 3.12 with the source package's serving profile. Use `python3 run.py setup serve` to install the bundled custom engine and its pinned dependencies; installing upstream SGLang alone does not provide the same model and decoding integration.
148
+
149
+ The serving profile pins Torch **2.9.1**, Transformers **5.5.4**, Triton **3.5.1**, sgl-kernel **0.3.20**, and FlashInfer **0.5.3**. Serving and evaluation use separate environments.
150
+
151
+ **Tested hardware:** NVIDIA **B300, B200, H200, H800**, and **PPU**. The provided SGLang serving installer is a CUDA/GPU recipe; PPU requires its matching platform runtime and is not selected by this GPU installer.
152
+
153
+ The default DLM service uses **BF16, Triton attention, eager execution, one GPU, one active request, and two queued requests**. Client concurrency queues requests; it does not imply a multi-request model batch. CUDA Graph and selective FP8 are discussed below as separate infrastructure experiments.
154
+
155
+ ## Model Version(s):
156
+
157
+ | Checkpoint | Generation | Evaluation mode |
158
+ |:---|:---|:---|
159
+ | [GroundAnything](https://huggingface.co/GroundingPI/GroundAnything) | Entropy-guided blockwise diffusion; optional self-speculation | **GAM** |
160
+ | [GroundAnything-VLM](https://huggingface.co/GroundingPI/GroundAnything-VLM) | Autoregressive generation | **GAM** |
161
+
162
+ The main GroundAnything benchmark results use **entropy-guided diffusion without autoregressive verification**. GroundAnything-VLM uses autoregressive decoding. The self-speculative speed–quality operating point is a separate experiment and should not be substituted for the main diffusion benchmark setting.
163
+
164
+ ## Testing and Evaluation Datasets:
165
+
166
+ ### Data Modality:
167
+
168
+ Image and text, with task-specific box, point, text-region, or interaction annotations.
169
+
170
+ ## Evaluation Dataset:
171
+
172
+ **Evaluation data: XXX — download URL coming soon.**
173
+
174
+ The shared evaluation toolkit provides **42 task recipes across eight task families**: Grounding, Referring, Dense, OCR, Layout, GUI, Pointing, and VisualPrompt. These are executable task/split recipes, not a count of distinct datasets. Dataset paths are registered in `configs/datasets.yaml`; task IDs are listed in `configs/eval/tasks.json`.
175
+
176
+ ### Evaluation Modes
177
+
178
+ The evaluator accepts seven mode identifiers. A mode selects the **prompt and output parser** independently of the inference engine and decoding algorithm.
179
+
180
+ | Mode | Supported evaluation interface | Output / coordinates |
181
+ |:---|:---|:---|
182
+ | **`GAM`** | **GroundingPI, GroundAnything, and GroundAnything-VLM** | Native spatial tokens, **0–999** |
183
+ | `DLM` | Legacy compatible diffusion-checkpoint alias | Same GAM spatial-token protocol |
184
+ | `RLV2` | Compatible RL-checkpoint alias | Same GAM spatial-token protocol |
185
+ | `VLM` | Generic VLM baselines | Explicit pixel, 0–1000, or 0–1 coordinate mode |
186
+ | `REXOMNI` | Rex-Omni adapter | Model-specific spatial-token parser |
187
+ | `LOCATEANYTHING` | LocateAnything adapter | Its own output parser and generation-mode setting |
188
+ | `GROUNDINGDINO` | External compatible GroundingDINO bridge | JSON coordinate responses |
189
+
190
+ **Use `mode: GAM` for all three of our checkpoints.** The `-VLM` suffix identifies an autoregressive checkpoint; it does not select the evaluator's generic `VLM` mode. The DLM's entropy-guided or self-speculative algorithm is selected separately through `decoder`.
191
+
192
+ External baseline weights and model services are separate from the evaluation toolkit. Metrics remain task-specific: box localization quality, point accuracy, OCR text-and-region matching, GUI grounding accuracy, and counting error. Full protocols and benchmark tables are provided in the paper and supplementary material.
193
+
194
+ ## Quantitative Evaluation Benchmarks
195
+
196
+ <p align="center"><img src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/fig7-grounding-performance.png" width="100%" alt="GroundAnything Figure 7: GroundAnything and GroundAnything-VLM benchmark overview" /></p>
197
+
198
+ *Figure 7. Paper-reported capability overview. Detailed task-level scores, baseline settings, and evaluation protocols are provided in the paper and supplementary material.*
199
+
200
+ ## Inference:
201
+
202
+ ### Installation
203
+
204
+ Run the commands from the **GroundAnything source repository root** after obtaining the code package, using **Linux x86_64 and Python 3.12**. The code link above is reserved for the public release. Use the serving profile in the supplied source bundle so that the custom model adapter, decoding implementation, and dependencies remain aligned.
205
+
206
+ ```bash
207
+ python3 -m pip install -r requirements.txt huggingface_hub
208
+ python3 run.py setup serve
209
+ ```
210
+
211
+ The installer verifies and extracts its bundled frameworks, creates `.venv-serve`, and records resolved packages. It prepares the custom SGLang implementation and applies the serving profile's dependency order. The profile includes the required cuDNN compatibility selection; its dependency-check report records the known Torch/cuDNN metadata exception.
212
+
213
+ **Choose the checkpoint-specific recipe below.** GroundAnything-VLM users need only the VLM download and launch commands; the DLM commands require the separate GroundAnything weights.
214
+
215
+ ### GroundAnything: Entropy-guided Service
216
+
217
+ Download the **complete DLM model bundle**, including its custom code and tokenizer, then launch:
218
+
219
+ ```bash
220
+ hf download GroundingPI/GroundAnything --local-dir weights/dlm_bundle
221
+ python3 run.py serve --decoder denoise
222
+ ```
223
+
224
+ The published model package is already a DLM bundle, so `prepare-model` is unnecessary for this download. The endpoint is `http://127.0.0.1:8101/v1`, with model ID `groundinganything`.
225
+
226
+ ### GroundAnything: Self-speculative Service
227
+
228
+ Stop the existing DLM service before switching its decoder:
229
+
230
+ ```bash
231
+ python3 run.py serve --decoder speculative
232
+ ```
233
+
234
+ This reuses the same DLM weights and endpoint, with diffusion proposals and greedy causal verification.
235
+
236
+ ### GroundAnything-VLM: Autoregressive Service
237
+
238
+ Use the separate autoregressive checkpoint and its service recipe:
239
+
240
+ ```bash
241
+ hf download GroundingPI/GroundAnything-VLM --local-dir weights/vlm
242
+ python3 run.py serve --config configs/release/vlm_sglang.yaml
243
+ ```
244
+
245
+ This service uses `http://127.0.0.1:8102/v1`, with model ID `groundinganything-vlm`. The two repositories share the spatial interface but have different loading and generation paths.
246
+
247
+ ### Worker (recommended)
248
+
249
+ Start the appropriate service once, then reuse the client:
250
+
251
+ ```python
252
+ from grounding_anything import GroundingAnything, visualize
253
+
254
+ client = GroundingAnything(
255
+ base_url="http://127.0.0.1:8101/v1",
256
+ model="groundinganything",
257
+ )
258
+ result = client.predict("example.jpg", "the red car", task="bbox")
259
+ print(result.to_dict())
260
+ if result.valid:
261
+ visualize("example.jpg", result).save("prediction.png")
262
+
263
+ point_result = client.predict(
264
+ "example.jpg", "the center of the red car", task="point"
265
+ )
266
+ ```
267
+
268
+ The HTTP client retains no model weights. A generic interactive request has no benchmark task identity; it does not automatically receive the benchmark's task-specific sampling and stopping policy. Use the evaluator below to reproduce the reported protocol.
269
+
270
+ ### Supported Tasks & Prompt Templates
271
+
272
+ | Task | Prompt example |
273
+ |:---|:---|
274
+ | Category / dense grounding | `Locate all the instances that match the following categories: car</c>person.` |
275
+ | Referring boxes | `Locate the target referred to by the following description: the red car.` |
276
+ | Category points | `Point to: car</c>person.` |
277
+ | Referring points | `Point to the target referred to by the following description: the red car.` |
278
+ | OCR | `OCR task detect all the text in box format.` |
279
+ | Layout | `Detect all document layout elements that match the following categories: title</c>text.` |
280
+ | GUI | `Point to the UI element to click for the following instruction: open the settings menu.` |
281
+ | Visual prompting | Provide reference boxes in the native spatial-token format, then request similar objects. |
282
+
283
+ ```text
284
+ Given reference boxes <|box_start|><100><200><500><650><|box_end|> indicating one or more objects, find all similar objects in the image and output their bounding boxes.
285
+ ```
286
+
287
+ The convenience client's `predict()` method wraps referring-box and referring-point prompts. Use an OpenAI-compatible request to `/chat/completions` for the other task templates; keep the image and prompt in the same user message.
288
+
289
+ ### Generation Modes
290
+
291
+ #### Entropy-guided decoding
292
+
293
+ Image/query prefill creates a causal prefix cache and a known anchor. The release recipe uses **block size 32, sub-block size 4, and entropy threshold 0.8**. A physical block contains the known anchor and 31 masked positions; sub-blocks are completed from left to right while the forward pass evaluates the physical block.
294
+
295
+ For each still-masked position in the active sub-block, the decoder measures entropy from the unmodified token distribution over the generatable vocabulary, excluding the mask token. Positions at or below the threshold are committed together. If no position qualifies, the lowest-entropy position is committed so decoding makes progress. Committed tokens remain fixed.
296
+
297
+ After a block is complete, a causal forward reconstructs its authoritative KV cache and supplies the next anchor. This cache-building pass **does not verify or reject the generated block**. A block requiring D denoising passes therefore uses D + 1 model forwards, excluding the initial prefill.
298
+
299
+ #### Self-speculative decoding
300
+
301
+ The model uses its own shared weights to draft tokens with bidirectional attention and verify them with causal attention. Verification accepts the **longest consecutive matching prefix**, stops at the first mismatch, applies the causal correction, and discards the rejected suffix cache states. The shipped speculative route uses greedy verification; it is not a general stochastic speculative sampler.
302
+
303
+ <p align="center"><img src="https://huggingface.co/GroundingPI/GroundAnything/resolve/0f8e30894c3ca86378d01ae51ec69c217c78151b/assets/fig6-self-speculative-decoding.png" width="100%" alt="GroundAnything Figure 6: linear and quadratic self-speculative schedules with shared model weights" /></p>
304
+
305
+ *Figure 6. The paper studies two schedules: linear drafting and verification use two model forwards and 2B query tokens per round; quadratic fusion uses one forward after initialization with B(B + 1) query tokens. These counts describe queries and model calls, not total Transformer FLOPs.*
306
+
307
+ The documented `--decoder speculative` service is the linear shared-weight route. Exact greedy verification is relative to the converted model's causal branch; it does not imply identical outputs to the separately trained GroundAnything-VLM checkpoint.
308
+
309
+ ### Benchmark Sampling Policy
310
+
311
+ Entropy-guided evaluation applies the five task profiles from `infer/decode/configs/task_profiles.json`:
312
+
313
+ | Task profile | Recipes | Temperature | Top-p | Output-token budget |
314
+ |:---|---:|---:|---:|---:|
315
+ | Strict single target | 14 | 0 | 1 | 512 |
316
+ | Medium, non-OCR | 12 | 0.1 | 0.95 | 4,096 |
317
+ | Dense, non-OCR | 8 | 0.1 | 0.95 | 8,192 |
318
+ | Medium OCR | 4 | 0.3 | 0.95 | 4,096 |
319
+ | Dense OCR | 4 | 0 | 1 | 4,096 |
320
+
321
+ The evaluator sends outer sampling parameters and the nested `custom_params.gam_dlm_decode` contract, adds the strict-single-target policy when required, and checks the server's actual decoder before sending requests. Unknown task profiles and mismatched decoders fail early. The entropy-guided benchmark policy rejects a global `max_tokens` override because it would replace the task budgets.
322
+
323
+ Self-speculative evaluation uses greedy sampling (`temperature=0`, `top_p=1`, repetition penalty 1). GroundAnything-VLM retains its own autoregressive task recipes. All of these routes use **GAM** prompts and spatial-token parsing.
324
+
325
+ ### Run Evaluation
326
+
327
+ ```bash
328
+ python3 run.py setup eval
329
+
330
+ # Match the command to the service that is already running.
331
+ python3 run.py eval --decoder denoise
332
+ python3 run.py eval --decoder speculative
333
+ python3 run.py eval --config configs/eval/vlm.yaml
334
+ ```
335
+
336
+ Choose one evaluation command for the intended checkpoint/decoder:
337
+
338
+ | Checkpoint / decoder | Recipe | Mode | Endpoint model ID |
339
+ |:---|:---|:---|:---|
340
+ | GroundAnything / entropy-guided | `configs/eval/dlm.yaml` | **GAM** | `groundinganything` |
341
+ | GroundAnything / self-speculative | `configs/eval/dlm_speculative.yaml` | **GAM** | `groundinganything` |
342
+ | GroundAnything-VLM / autoregressive | `configs/eval/vlm.yaml` | **GAM** | `groundinganything-vlm` |
343
+
344
+ Use `service_contract: openai` and the corresponding port (8101 for DLM, 8102 for VLM). For a preflight that checks configured inputs without sending inference requests:
345
+
346
+ ```bash
347
+ .venv-eval/bin/python scripts/evaluate.py configs/eval/dlm.yaml --dry-run
348
+ ```
349
+
350
+ Configure the running service, `data_root`, dataset registry, task list, and a fresh `run_id` before launching. The default `limit: 8` is a smoke test; set **`limit: null`** for full evaluation. Evaluation connects to an existing service and does not start or change the decoder.
351
+
352
+ Each run saves `run.json` (configuration and provenance), `responses.jsonl` (raw responses, finish reasons, and token usage), task logs, and `summary.json` (metrics and completion state). Record the checkpoint revision and task configuration when comparing results.
353
+
354
+ ## Inference Infrastructure
355
+
356
+ ### SGLang Execution
357
+
358
+ The custom SGLang integration coordinates model loading, request scheduling, attention kernels, and KV-cache ownership for blockwise generation. Denoising uses bidirectional attention inside the active block; completed history is retained as causal KV. Self-speculation additionally verifies proposals and removes rejected suffix states. These cache semantics must be preserved when optimizing execution.
359
+
360
+ The supplied recipe selects **BF16 + Triton attention + eager execution**. DLM recipes record the engine source, decoder, effective settings, profile digest, and package versions in `outputs/sglang/<decoder>/engine_runtime.json`; DLM evaluation checks decoder alignment before starting workers. The VLM service records its causal model/runtime configuration separately in `outputs/sglang/vlm/engine_runtime.json`. For higher service concurrency, use independent replicas with distinct devices, ports, and output directories; the default queue is not continuous multi-request batching.
361
+
362
+ ### CUDA Graph and Selective FP8
363
+
364
+ The paper evaluates a progressive infrastructure sequence:
365
+
366
+ | Layer | Purpose | Status in the supplied default recipe |
367
+ |:---|:---|:---|
368
+ | Native PyTorch eager | Reference model execution | Reference implementation |
369
+ | SGLang eager | Integrated scheduling, attention, and cache execution | **Default serving path** |
370
+ | CUDA Graph replay | Replay compatible captured GPU work to reduce repeated launch overhead | Evaluated in the paper; **disabled by the default launcher** |
371
+ | Selective FP8 | Reduce arithmetic cost in eligible language-model linear operations while retaining other components in BF16 | Evaluated in the paper; **default serving remains BF16** |
372
+
373
+ CUDA Graph changes how compatible GPU work is submitted; it does not define a new token-commitment or verification rule. Captured shapes and state updates must remain compatible with the active decoding path. Selective FP8 can change logits and subsequent decoding decisions, so it is a distinct numerical configuration.
374
+
375
+ The optional graph implementation captures fixed-shape block work. Its FlashInfer path uses persistent attention masks and device-resident buffer updates for bidirectional drafting and causal verification. The Triton verifier path separates draft/verification metadata and input buffers while sharing parameters and the real KV pool. Its supported capture case is B32 with one request and no tensor/pipeline/data parallel expansion; prefill and unsupported shapes use eager execution. Optional shadow checks compare cache states, logits, and token choices against eager execution. These implementation paths are not exposed as a public `run.py serve --cuda-graph` switch in this release.
376
+
377
+ The paper's cumulative infrastructure speedups compare execution implementations **within the same decoding mode**. They are separate from the headline comparison of self-speculation against the autoregressive checkpoint. The default release commands above do not enable Graph replay or FP8, and no unsupported activation flags are implied.
378
+
379
+ ## Ethical Considerations:
380
+
381
+ Grounding predictions can miss small or occluded objects, repeat instances, or produce inaccurate text and coordinates. Validate localization quality on the intended task and inspect incomplete responses. GUI points describe image locations; the model does not execute interface actions. Perception outputs require task-specific validation before integration into physical systems.
added_tokens.json ADDED
@@ -0,0 +1,1030 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "</c>": 152669,
3
+ "</think>": 151668,
4
+ "</tool_call>": 151658,
5
+ "</tool_response>": 151666,
6
+ "<0>": 151669,
7
+ "<100>": 151769,
8
+ "<101>": 151770,
9
+ "<102>": 151771,
10
+ "<103>": 151772,
11
+ "<104>": 151773,
12
+ "<105>": 151774,
13
+ "<106>": 151775,
14
+ "<107>": 151776,
15
+ "<108>": 151777,
16
+ "<109>": 151778,
17
+ "<10>": 151679,
18
+ "<110>": 151779,
19
+ "<111>": 151780,
20
+ "<112>": 151781,
21
+ "<113>": 151782,
22
+ "<114>": 151783,
23
+ "<115>": 151784,
24
+ "<116>": 151785,
25
+ "<117>": 151786,
26
+ "<118>": 151787,
27
+ "<119>": 151788,
28
+ "<11>": 151680,
29
+ "<120>": 151789,
30
+ "<121>": 151790,
31
+ "<122>": 151791,
32
+ "<123>": 151792,
33
+ "<124>": 151793,
34
+ "<125>": 151794,
35
+ "<126>": 151795,
36
+ "<127>": 151796,
37
+ "<128>": 151797,
38
+ "<129>": 151798,
39
+ "<12>": 151681,
40
+ "<130>": 151799,
41
+ "<131>": 151800,
42
+ "<132>": 151801,
43
+ "<133>": 151802,
44
+ "<134>": 151803,
45
+ "<135>": 151804,
46
+ "<136>": 151805,
47
+ "<137>": 151806,
48
+ "<138>": 151807,
49
+ "<139>": 151808,
50
+ "<13>": 151682,
51
+ "<140>": 151809,
52
+ "<141>": 151810,
53
+ "<142>": 151811,
54
+ "<143>": 151812,
55
+ "<144>": 151813,
56
+ "<145>": 151814,
57
+ "<146>": 151815,
58
+ "<147>": 151816,
59
+ "<148>": 151817,
60
+ "<149>": 151818,
61
+ "<14>": 151683,
62
+ "<150>": 151819,
63
+ "<151>": 151820,
64
+ "<152>": 151821,
65
+ "<153>": 151822,
66
+ "<154>": 151823,
67
+ "<155>": 151824,
68
+ "<156>": 151825,
69
+ "<157>": 151826,
70
+ "<158>": 151827,
71
+ "<159>": 151828,
72
+ "<15>": 151684,
73
+ "<160>": 151829,
74
+ "<161>": 151830,
75
+ "<162>": 151831,
76
+ "<163>": 151832,
77
+ "<164>": 151833,
78
+ "<165>": 151834,
79
+ "<166>": 151835,
80
+ "<167>": 151836,
81
+ "<168>": 151837,
82
+ "<169>": 151838,
83
+ "<16>": 151685,
84
+ "<170>": 151839,
85
+ "<171>": 151840,
86
+ "<172>": 151841,
87
+ "<173>": 151842,
88
+ "<174>": 151843,
89
+ "<175>": 151844,
90
+ "<176>": 151845,
91
+ "<177>": 151846,
92
+ "<178>": 151847,
93
+ "<179>": 151848,
94
+ "<17>": 151686,
95
+ "<180>": 151849,
96
+ "<181>": 151850,
97
+ "<182>": 151851,
98
+ "<183>": 151852,
99
+ "<184>": 151853,
100
+ "<185>": 151854,
101
+ "<186>": 151855,
102
+ "<187>": 151856,
103
+ "<188>": 151857,
104
+ "<189>": 151858,
105
+ "<18>": 151687,
106
+ "<190>": 151859,
107
+ "<191>": 151860,
108
+ "<192>": 151861,
109
+ "<193>": 151862,
110
+ "<194>": 151863,
111
+ "<195>": 151864,
112
+ "<196>": 151865,
113
+ "<197>": 151866,
114
+ "<198>": 151867,
115
+ "<199>": 151868,
116
+ "<19>": 151688,
117
+ "<1>": 151670,
118
+ "<200>": 151869,
119
+ "<201>": 151870,
120
+ "<202>": 151871,
121
+ "<203>": 151872,
122
+ "<204>": 151873,
123
+ "<205>": 151874,
124
+ "<206>": 151875,
125
+ "<207>": 151876,
126
+ "<208>": 151877,
127
+ "<209>": 151878,
128
+ "<20>": 151689,
129
+ "<210>": 151879,
130
+ "<211>": 151880,
131
+ "<212>": 151881,
132
+ "<213>": 151882,
133
+ "<214>": 151883,
134
+ "<215>": 151884,
135
+ "<216>": 151885,
136
+ "<217>": 151886,
137
+ "<218>": 151887,
138
+ "<219>": 151888,
139
+ "<21>": 151690,
140
+ "<220>": 151889,
141
+ "<221>": 151890,
142
+ "<222>": 151891,
143
+ "<223>": 151892,
144
+ "<224>": 151893,
145
+ "<225>": 151894,
146
+ "<226>": 151895,
147
+ "<227>": 151896,
148
+ "<228>": 151897,
149
+ "<229>": 151898,
150
+ "<22>": 151691,
151
+ "<230>": 151899,
152
+ "<231>": 151900,
153
+ "<232>": 151901,
154
+ "<233>": 151902,
155
+ "<234>": 151903,
156
+ "<235>": 151904,
157
+ "<236>": 151905,
158
+ "<237>": 151906,
159
+ "<238>": 151907,
160
+ "<239>": 151908,
161
+ "<23>": 151692,
162
+ "<240>": 151909,
163
+ "<241>": 151910,
164
+ "<242>": 151911,
165
+ "<243>": 151912,
166
+ "<244>": 151913,
167
+ "<245>": 151914,
168
+ "<246>": 151915,
169
+ "<247>": 151916,
170
+ "<248>": 151917,
171
+ "<249>": 151918,
172
+ "<24>": 151693,
173
+ "<250>": 151919,
174
+ "<251>": 151920,
175
+ "<252>": 151921,
176
+ "<253>": 151922,
177
+ "<254>": 151923,
178
+ "<255>": 151924,
179
+ "<256>": 151925,
180
+ "<257>": 151926,
181
+ "<258>": 151927,
182
+ "<259>": 151928,
183
+ "<25>": 151694,
184
+ "<260>": 151929,
185
+ "<261>": 151930,
186
+ "<262>": 151931,
187
+ "<263>": 151932,
188
+ "<264>": 151933,
189
+ "<265>": 151934,
190
+ "<266>": 151935,
191
+ "<267>": 151936,
192
+ "<268>": 151937,
193
+ "<269>": 151938,
194
+ "<26>": 151695,
195
+ "<270>": 151939,
196
+ "<271>": 151940,
197
+ "<272>": 151941,
198
+ "<273>": 151942,
199
+ "<274>": 151943,
200
+ "<275>": 151944,
201
+ "<276>": 151945,
202
+ "<277>": 151946,
203
+ "<278>": 151947,
204
+ "<279>": 151948,
205
+ "<27>": 151696,
206
+ "<280>": 151949,
207
+ "<281>": 151950,
208
+ "<282>": 151951,
209
+ "<283>": 151952,
210
+ "<284>": 151953,
211
+ "<285>": 151954,
212
+ "<286>": 151955,
213
+ "<287>": 151956,
214
+ "<288>": 151957,
215
+ "<289>": 151958,
216
+ "<28>": 151697,
217
+ "<290>": 151959,
218
+ "<291>": 151960,
219
+ "<292>": 151961,
220
+ "<293>": 151962,
221
+ "<294>": 151963,
222
+ "<295>": 151964,
223
+ "<296>": 151965,
224
+ "<297>": 151966,
225
+ "<298>": 151967,
226
+ "<299>": 151968,
227
+ "<29>": 151698,
228
+ "<2>": 151671,
229
+ "<300>": 151969,
230
+ "<301>": 151970,
231
+ "<302>": 151971,
232
+ "<303>": 151972,
233
+ "<304>": 151973,
234
+ "<305>": 151974,
235
+ "<306>": 151975,
236
+ "<307>": 151976,
237
+ "<308>": 151977,
238
+ "<309>": 151978,
239
+ "<30>": 151699,
240
+ "<310>": 151979,
241
+ "<311>": 151980,
242
+ "<312>": 151981,
243
+ "<313>": 151982,
244
+ "<314>": 151983,
245
+ "<315>": 151984,
246
+ "<316>": 151985,
247
+ "<317>": 151986,
248
+ "<318>": 151987,
249
+ "<319>": 151988,
250
+ "<31>": 151700,
251
+ "<320>": 151989,
252
+ "<321>": 151990,
253
+ "<322>": 151991,
254
+ "<323>": 151992,
255
+ "<324>": 151993,
256
+ "<325>": 151994,
257
+ "<326>": 151995,
258
+ "<327>": 151996,
259
+ "<328>": 151997,
260
+ "<329>": 151998,
261
+ "<32>": 151701,
262
+ "<330>": 151999,
263
+ "<331>": 152000,
264
+ "<332>": 152001,
265
+ "<333>": 152002,
266
+ "<334>": 152003,
267
+ "<335>": 152004,
268
+ "<336>": 152005,
269
+ "<337>": 152006,
270
+ "<338>": 152007,
271
+ "<339>": 152008,
272
+ "<33>": 151702,
273
+ "<340>": 152009,
274
+ "<341>": 152010,
275
+ "<342>": 152011,
276
+ "<343>": 152012,
277
+ "<344>": 152013,
278
+ "<345>": 152014,
279
+ "<346>": 152015,
280
+ "<347>": 152016,
281
+ "<348>": 152017,
282
+ "<349>": 152018,
283
+ "<34>": 151703,
284
+ "<350>": 152019,
285
+ "<351>": 152020,
286
+ "<352>": 152021,
287
+ "<353>": 152022,
288
+ "<354>": 152023,
289
+ "<355>": 152024,
290
+ "<356>": 152025,
291
+ "<357>": 152026,
292
+ "<358>": 152027,
293
+ "<359>": 152028,
294
+ "<35>": 151704,
295
+ "<360>": 152029,
296
+ "<361>": 152030,
297
+ "<362>": 152031,
298
+ "<363>": 152032,
299
+ "<364>": 152033,
300
+ "<365>": 152034,
301
+ "<366>": 152035,
302
+ "<367>": 152036,
303
+ "<368>": 152037,
304
+ "<369>": 152038,
305
+ "<36>": 151705,
306
+ "<370>": 152039,
307
+ "<371>": 152040,
308
+ "<372>": 152041,
309
+ "<373>": 152042,
310
+ "<374>": 152043,
311
+ "<375>": 152044,
312
+ "<376>": 152045,
313
+ "<377>": 152046,
314
+ "<378>": 152047,
315
+ "<379>": 152048,
316
+ "<37>": 151706,
317
+ "<380>": 152049,
318
+ "<381>": 152050,
319
+ "<382>": 152051,
320
+ "<383>": 152052,
321
+ "<384>": 152053,
322
+ "<385>": 152054,
323
+ "<386>": 152055,
324
+ "<387>": 152056,
325
+ "<388>": 152057,
326
+ "<389>": 152058,
327
+ "<38>": 151707,
328
+ "<390>": 152059,
329
+ "<391>": 152060,
330
+ "<392>": 152061,
331
+ "<393>": 152062,
332
+ "<394>": 152063,
333
+ "<395>": 152064,
334
+ "<396>": 152065,
335
+ "<397>": 152066,
336
+ "<398>": 152067,
337
+ "<399>": 152068,
338
+ "<39>": 151708,
339
+ "<3>": 151672,
340
+ "<400>": 152069,
341
+ "<401>": 152070,
342
+ "<402>": 152071,
343
+ "<403>": 152072,
344
+ "<404>": 152073,
345
+ "<405>": 152074,
346
+ "<406>": 152075,
347
+ "<407>": 152076,
348
+ "<408>": 152077,
349
+ "<409>": 152078,
350
+ "<40>": 151709,
351
+ "<410>": 152079,
352
+ "<411>": 152080,
353
+ "<412>": 152081,
354
+ "<413>": 152082,
355
+ "<414>": 152083,
356
+ "<415>": 152084,
357
+ "<416>": 152085,
358
+ "<417>": 152086,
359
+ "<418>": 152087,
360
+ "<419>": 152088,
361
+ "<41>": 151710,
362
+ "<420>": 152089,
363
+ "<421>": 152090,
364
+ "<422>": 152091,
365
+ "<423>": 152092,
366
+ "<424>": 152093,
367
+ "<425>": 152094,
368
+ "<426>": 152095,
369
+ "<427>": 152096,
370
+ "<428>": 152097,
371
+ "<429>": 152098,
372
+ "<42>": 151711,
373
+ "<430>": 152099,
374
+ "<431>": 152100,
375
+ "<432>": 152101,
376
+ "<433>": 152102,
377
+ "<434>": 152103,
378
+ "<435>": 152104,
379
+ "<436>": 152105,
380
+ "<437>": 152106,
381
+ "<438>": 152107,
382
+ "<439>": 152108,
383
+ "<43>": 151712,
384
+ "<440>": 152109,
385
+ "<441>": 152110,
386
+ "<442>": 152111,
387
+ "<443>": 152112,
388
+ "<444>": 152113,
389
+ "<445>": 152114,
390
+ "<446>": 152115,
391
+ "<447>": 152116,
392
+ "<448>": 152117,
393
+ "<449>": 152118,
394
+ "<44>": 151713,
395
+ "<450>": 152119,
396
+ "<451>": 152120,
397
+ "<452>": 152121,
398
+ "<453>": 152122,
399
+ "<454>": 152123,
400
+ "<455>": 152124,
401
+ "<456>": 152125,
402
+ "<457>": 152126,
403
+ "<458>": 152127,
404
+ "<459>": 152128,
405
+ "<45>": 151714,
406
+ "<460>": 152129,
407
+ "<461>": 152130,
408
+ "<462>": 152131,
409
+ "<463>": 152132,
410
+ "<464>": 152133,
411
+ "<465>": 152134,
412
+ "<466>": 152135,
413
+ "<467>": 152136,
414
+ "<468>": 152137,
415
+ "<469>": 152138,
416
+ "<46>": 151715,
417
+ "<470>": 152139,
418
+ "<471>": 152140,
419
+ "<472>": 152141,
420
+ "<473>": 152142,
421
+ "<474>": 152143,
422
+ "<475>": 152144,
423
+ "<476>": 152145,
424
+ "<477>": 152146,
425
+ "<478>": 152147,
426
+ "<479>": 152148,
427
+ "<47>": 151716,
428
+ "<480>": 152149,
429
+ "<481>": 152150,
430
+ "<482>": 152151,
431
+ "<483>": 152152,
432
+ "<484>": 152153,
433
+ "<485>": 152154,
434
+ "<486>": 152155,
435
+ "<487>": 152156,
436
+ "<488>": 152157,
437
+ "<489>": 152158,
438
+ "<48>": 151717,
439
+ "<490>": 152159,
440
+ "<491>": 152160,
441
+ "<492>": 152161,
442
+ "<493>": 152162,
443
+ "<494>": 152163,
444
+ "<495>": 152164,
445
+ "<496>": 152165,
446
+ "<497>": 152166,
447
+ "<498>": 152167,
448
+ "<499>": 152168,
449
+ "<49>": 151718,
450
+ "<4>": 151673,
451
+ "<500>": 152169,
452
+ "<501>": 152170,
453
+ "<502>": 152171,
454
+ "<503>": 152172,
455
+ "<504>": 152173,
456
+ "<505>": 152174,
457
+ "<506>": 152175,
458
+ "<507>": 152176,
459
+ "<508>": 152177,
460
+ "<509>": 152178,
461
+ "<50>": 151719,
462
+ "<510>": 152179,
463
+ "<511>": 152180,
464
+ "<512>": 152181,
465
+ "<513>": 152182,
466
+ "<514>": 152183,
467
+ "<515>": 152184,
468
+ "<516>": 152185,
469
+ "<517>": 152186,
470
+ "<518>": 152187,
471
+ "<519>": 152188,
472
+ "<51>": 151720,
473
+ "<520>": 152189,
474
+ "<521>": 152190,
475
+ "<522>": 152191,
476
+ "<523>": 152192,
477
+ "<524>": 152193,
478
+ "<525>": 152194,
479
+ "<526>": 152195,
480
+ "<527>": 152196,
481
+ "<528>": 152197,
482
+ "<529>": 152198,
483
+ "<52>": 151721,
484
+ "<530>": 152199,
485
+ "<531>": 152200,
486
+ "<532>": 152201,
487
+ "<533>": 152202,
488
+ "<534>": 152203,
489
+ "<535>": 152204,
490
+ "<536>": 152205,
491
+ "<537>": 152206,
492
+ "<538>": 152207,
493
+ "<539>": 152208,
494
+ "<53>": 151722,
495
+ "<540>": 152209,
496
+ "<541>": 152210,
497
+ "<542>": 152211,
498
+ "<543>": 152212,
499
+ "<544>": 152213,
500
+ "<545>": 152214,
501
+ "<546>": 152215,
502
+ "<547>": 152216,
503
+ "<548>": 152217,
504
+ "<549>": 152218,
505
+ "<54>": 151723,
506
+ "<550>": 152219,
507
+ "<551>": 152220,
508
+ "<552>": 152221,
509
+ "<553>": 152222,
510
+ "<554>": 152223,
511
+ "<555>": 152224,
512
+ "<556>": 152225,
513
+ "<557>": 152226,
514
+ "<558>": 152227,
515
+ "<559>": 152228,
516
+ "<55>": 151724,
517
+ "<560>": 152229,
518
+ "<561>": 152230,
519
+ "<562>": 152231,
520
+ "<563>": 152232,
521
+ "<564>": 152233,
522
+ "<565>": 152234,
523
+ "<566>": 152235,
524
+ "<567>": 152236,
525
+ "<568>": 152237,
526
+ "<569>": 152238,
527
+ "<56>": 151725,
528
+ "<570>": 152239,
529
+ "<571>": 152240,
530
+ "<572>": 152241,
531
+ "<573>": 152242,
532
+ "<574>": 152243,
533
+ "<575>": 152244,
534
+ "<576>": 152245,
535
+ "<577>": 152246,
536
+ "<578>": 152247,
537
+ "<579>": 152248,
538
+ "<57>": 151726,
539
+ "<580>": 152249,
540
+ "<581>": 152250,
541
+ "<582>": 152251,
542
+ "<583>": 152252,
543
+ "<584>": 152253,
544
+ "<585>": 152254,
545
+ "<586>": 152255,
546
+ "<587>": 152256,
547
+ "<588>": 152257,
548
+ "<589>": 152258,
549
+ "<58>": 151727,
550
+ "<590>": 152259,
551
+ "<591>": 152260,
552
+ "<592>": 152261,
553
+ "<593>": 152262,
554
+ "<594>": 152263,
555
+ "<595>": 152264,
556
+ "<596>": 152265,
557
+ "<597>": 152266,
558
+ "<598>": 152267,
559
+ "<599>": 152268,
560
+ "<59>": 151728,
561
+ "<5>": 151674,
562
+ "<600>": 152269,
563
+ "<601>": 152270,
564
+ "<602>": 152271,
565
+ "<603>": 152272,
566
+ "<604>": 152273,
567
+ "<605>": 152274,
568
+ "<606>": 152275,
569
+ "<607>": 152276,
570
+ "<608>": 152277,
571
+ "<609>": 152278,
572
+ "<60>": 151729,
573
+ "<610>": 152279,
574
+ "<611>": 152280,
575
+ "<612>": 152281,
576
+ "<613>": 152282,
577
+ "<614>": 152283,
578
+ "<615>": 152284,
579
+ "<616>": 152285,
580
+ "<617>": 152286,
581
+ "<618>": 152287,
582
+ "<619>": 152288,
583
+ "<61>": 151730,
584
+ "<620>": 152289,
585
+ "<621>": 152290,
586
+ "<622>": 152291,
587
+ "<623>": 152292,
588
+ "<624>": 152293,
589
+ "<625>": 152294,
590
+ "<626>": 152295,
591
+ "<627>": 152296,
592
+ "<628>": 152297,
593
+ "<629>": 152298,
594
+ "<62>": 151731,
595
+ "<630>": 152299,
596
+ "<631>": 152300,
597
+ "<632>": 152301,
598
+ "<633>": 152302,
599
+ "<634>": 152303,
600
+ "<635>": 152304,
601
+ "<636>": 152305,
602
+ "<637>": 152306,
603
+ "<638>": 152307,
604
+ "<639>": 152308,
605
+ "<63>": 151732,
606
+ "<640>": 152309,
607
+ "<641>": 152310,
608
+ "<642>": 152311,
609
+ "<643>": 152312,
610
+ "<644>": 152313,
611
+ "<645>": 152314,
612
+ "<646>": 152315,
613
+ "<647>": 152316,
614
+ "<648>": 152317,
615
+ "<649>": 152318,
616
+ "<64>": 151733,
617
+ "<650>": 152319,
618
+ "<651>": 152320,
619
+ "<652>": 152321,
620
+ "<653>": 152322,
621
+ "<654>": 152323,
622
+ "<655>": 152324,
623
+ "<656>": 152325,
624
+ "<657>": 152326,
625
+ "<658>": 152327,
626
+ "<659>": 152328,
627
+ "<65>": 151734,
628
+ "<660>": 152329,
629
+ "<661>": 152330,
630
+ "<662>": 152331,
631
+ "<663>": 152332,
632
+ "<664>": 152333,
633
+ "<665>": 152334,
634
+ "<666>": 152335,
635
+ "<667>": 152336,
636
+ "<668>": 152337,
637
+ "<669>": 152338,
638
+ "<66>": 151735,
639
+ "<670>": 152339,
640
+ "<671>": 152340,
641
+ "<672>": 152341,
642
+ "<673>": 152342,
643
+ "<674>": 152343,
644
+ "<675>": 152344,
645
+ "<676>": 152345,
646
+ "<677>": 152346,
647
+ "<678>": 152347,
648
+ "<679>": 152348,
649
+ "<67>": 151736,
650
+ "<680>": 152349,
651
+ "<681>": 152350,
652
+ "<682>": 152351,
653
+ "<683>": 152352,
654
+ "<684>": 152353,
655
+ "<685>": 152354,
656
+ "<686>": 152355,
657
+ "<687>": 152356,
658
+ "<688>": 152357,
659
+ "<689>": 152358,
660
+ "<68>": 151737,
661
+ "<690>": 152359,
662
+ "<691>": 152360,
663
+ "<692>": 152361,
664
+ "<693>": 152362,
665
+ "<694>": 152363,
666
+ "<695>": 152364,
667
+ "<696>": 152365,
668
+ "<697>": 152366,
669
+ "<698>": 152367,
670
+ "<699>": 152368,
671
+ "<69>": 151738,
672
+ "<6>": 151675,
673
+ "<700>": 152369,
674
+ "<701>": 152370,
675
+ "<702>": 152371,
676
+ "<703>": 152372,
677
+ "<704>": 152373,
678
+ "<705>": 152374,
679
+ "<706>": 152375,
680
+ "<707>": 152376,
681
+ "<708>": 152377,
682
+ "<709>": 152378,
683
+ "<70>": 151739,
684
+ "<710>": 152379,
685
+ "<711>": 152380,
686
+ "<712>": 152381,
687
+ "<713>": 152382,
688
+ "<714>": 152383,
689
+ "<715>": 152384,
690
+ "<716>": 152385,
691
+ "<717>": 152386,
692
+ "<718>": 152387,
693
+ "<719>": 152388,
694
+ "<71>": 151740,
695
+ "<720>": 152389,
696
+ "<721>": 152390,
697
+ "<722>": 152391,
698
+ "<723>": 152392,
699
+ "<724>": 152393,
700
+ "<725>": 152394,
701
+ "<726>": 152395,
702
+ "<727>": 152396,
703
+ "<728>": 152397,
704
+ "<729>": 152398,
705
+ "<72>": 151741,
706
+ "<730>": 152399,
707
+ "<731>": 152400,
708
+ "<732>": 152401,
709
+ "<733>": 152402,
710
+ "<734>": 152403,
711
+ "<735>": 152404,
712
+ "<736>": 152405,
713
+ "<737>": 152406,
714
+ "<738>": 152407,
715
+ "<739>": 152408,
716
+ "<73>": 151742,
717
+ "<740>": 152409,
718
+ "<741>": 152410,
719
+ "<742>": 152411,
720
+ "<743>": 152412,
721
+ "<744>": 152413,
722
+ "<745>": 152414,
723
+ "<746>": 152415,
724
+ "<747>": 152416,
725
+ "<748>": 152417,
726
+ "<749>": 152418,
727
+ "<74>": 151743,
728
+ "<750>": 152419,
729
+ "<751>": 152420,
730
+ "<752>": 152421,
731
+ "<753>": 152422,
732
+ "<754>": 152423,
733
+ "<755>": 152424,
734
+ "<756>": 152425,
735
+ "<757>": 152426,
736
+ "<758>": 152427,
737
+ "<759>": 152428,
738
+ "<75>": 151744,
739
+ "<760>": 152429,
740
+ "<761>": 152430,
741
+ "<762>": 152431,
742
+ "<763>": 152432,
743
+ "<764>": 152433,
744
+ "<765>": 152434,
745
+ "<766>": 152435,
746
+ "<767>": 152436,
747
+ "<768>": 152437,
748
+ "<769>": 152438,
749
+ "<76>": 151745,
750
+ "<770>": 152439,
751
+ "<771>": 152440,
752
+ "<772>": 152441,
753
+ "<773>": 152442,
754
+ "<774>": 152443,
755
+ "<775>": 152444,
756
+ "<776>": 152445,
757
+ "<777>": 152446,
758
+ "<778>": 152447,
759
+ "<779>": 152448,
760
+ "<77>": 151746,
761
+ "<780>": 152449,
762
+ "<781>": 152450,
763
+ "<782>": 152451,
764
+ "<783>": 152452,
765
+ "<784>": 152453,
766
+ "<785>": 152454,
767
+ "<786>": 152455,
768
+ "<787>": 152456,
769
+ "<788>": 152457,
770
+ "<789>": 152458,
771
+ "<78>": 151747,
772
+ "<790>": 152459,
773
+ "<791>": 152460,
774
+ "<792>": 152461,
775
+ "<793>": 152462,
776
+ "<794>": 152463,
777
+ "<795>": 152464,
778
+ "<796>": 152465,
779
+ "<797>": 152466,
780
+ "<798>": 152467,
781
+ "<799>": 152468,
782
+ "<79>": 151748,
783
+ "<7>": 151676,
784
+ "<800>": 152469,
785
+ "<801>": 152470,
786
+ "<802>": 152471,
787
+ "<803>": 152472,
788
+ "<804>": 152473,
789
+ "<805>": 152474,
790
+ "<806>": 152475,
791
+ "<807>": 152476,
792
+ "<808>": 152477,
793
+ "<809>": 152478,
794
+ "<80>": 151749,
795
+ "<810>": 152479,
796
+ "<811>": 152480,
797
+ "<812>": 152481,
798
+ "<813>": 152482,
799
+ "<814>": 152483,
800
+ "<815>": 152484,
801
+ "<816>": 152485,
802
+ "<817>": 152486,
803
+ "<818>": 152487,
804
+ "<819>": 152488,
805
+ "<81>": 151750,
806
+ "<820>": 152489,
807
+ "<821>": 152490,
808
+ "<822>": 152491,
809
+ "<823>": 152492,
810
+ "<824>": 152493,
811
+ "<825>": 152494,
812
+ "<826>": 152495,
813
+ "<827>": 152496,
814
+ "<828>": 152497,
815
+ "<829>": 152498,
816
+ "<82>": 151751,
817
+ "<830>": 152499,
818
+ "<831>": 152500,
819
+ "<832>": 152501,
820
+ "<833>": 152502,
821
+ "<834>": 152503,
822
+ "<835>": 152504,
823
+ "<836>": 152505,
824
+ "<837>": 152506,
825
+ "<838>": 152507,
826
+ "<839>": 152508,
827
+ "<83>": 151752,
828
+ "<840>": 152509,
829
+ "<841>": 152510,
830
+ "<842>": 152511,
831
+ "<843>": 152512,
832
+ "<844>": 152513,
833
+ "<845>": 152514,
834
+ "<846>": 152515,
835
+ "<847>": 152516,
836
+ "<848>": 152517,
837
+ "<849>": 152518,
838
+ "<84>": 151753,
839
+ "<850>": 152519,
840
+ "<851>": 152520,
841
+ "<852>": 152521,
842
+ "<853>": 152522,
843
+ "<854>": 152523,
844
+ "<855>": 152524,
845
+ "<856>": 152525,
846
+ "<857>": 152526,
847
+ "<858>": 152527,
848
+ "<859>": 152528,
849
+ "<85>": 151754,
850
+ "<860>": 152529,
851
+ "<861>": 152530,
852
+ "<862>": 152531,
853
+ "<863>": 152532,
854
+ "<864>": 152533,
855
+ "<865>": 152534,
856
+ "<866>": 152535,
857
+ "<867>": 152536,
858
+ "<868>": 152537,
859
+ "<869>": 152538,
860
+ "<86>": 151755,
861
+ "<870>": 152539,
862
+ "<871>": 152540,
863
+ "<872>": 152541,
864
+ "<873>": 152542,
865
+ "<874>": 152543,
866
+ "<875>": 152544,
867
+ "<876>": 152545,
868
+ "<877>": 152546,
869
+ "<878>": 152547,
870
+ "<879>": 152548,
871
+ "<87>": 151756,
872
+ "<880>": 152549,
873
+ "<881>": 152550,
874
+ "<882>": 152551,
875
+ "<883>": 152552,
876
+ "<884>": 152553,
877
+ "<885>": 152554,
878
+ "<886>": 152555,
879
+ "<887>": 152556,
880
+ "<888>": 152557,
881
+ "<889>": 152558,
882
+ "<88>": 151757,
883
+ "<890>": 152559,
884
+ "<891>": 152560,
885
+ "<892>": 152561,
886
+ "<893>": 152562,
887
+ "<894>": 152563,
888
+ "<895>": 152564,
889
+ "<896>": 152565,
890
+ "<897>": 152566,
891
+ "<898>": 152567,
892
+ "<899>": 152568,
893
+ "<89>": 151758,
894
+ "<8>": 151677,
895
+ "<900>": 152569,
896
+ "<901>": 152570,
897
+ "<902>": 152571,
898
+ "<903>": 152572,
899
+ "<904>": 152573,
900
+ "<905>": 152574,
901
+ "<906>": 152575,
902
+ "<907>": 152576,
903
+ "<908>": 152577,
904
+ "<909>": 152578,
905
+ "<90>": 151759,
906
+ "<910>": 152579,
907
+ "<911>": 152580,
908
+ "<912>": 152581,
909
+ "<913>": 152582,
910
+ "<914>": 152583,
911
+ "<915>": 152584,
912
+ "<916>": 152585,
913
+ "<917>": 152586,
914
+ "<918>": 152587,
915
+ "<919>": 152588,
916
+ "<91>": 151760,
917
+ "<920>": 152589,
918
+ "<921>": 152590,
919
+ "<922>": 152591,
920
+ "<923>": 152592,
921
+ "<924>": 152593,
922
+ "<925>": 152594,
923
+ "<926>": 152595,
924
+ "<927>": 152596,
925
+ "<928>": 152597,
926
+ "<929>": 152598,
927
+ "<92>": 151761,
928
+ "<930>": 152599,
929
+ "<931>": 152600,
930
+ "<932>": 152601,
931
+ "<933>": 152602,
932
+ "<934>": 152603,
933
+ "<935>": 152604,
934
+ "<936>": 152605,
935
+ "<937>": 152606,
936
+ "<938>": 152607,
937
+ "<939>": 152608,
938
+ "<93>": 151762,
939
+ "<940>": 152609,
940
+ "<941>": 152610,
941
+ "<942>": 152611,
942
+ "<943>": 152612,
943
+ "<944>": 152613,
944
+ "<945>": 152614,
945
+ "<946>": 152615,
946
+ "<947>": 152616,
947
+ "<948>": 152617,
948
+ "<949>": 152618,
949
+ "<94>": 151763,
950
+ "<950>": 152619,
951
+ "<951>": 152620,
952
+ "<952>": 152621,
953
+ "<953>": 152622,
954
+ "<954>": 152623,
955
+ "<955>": 152624,
956
+ "<956>": 152625,
957
+ "<957>": 152626,
958
+ "<958>": 152627,
959
+ "<959>": 152628,
960
+ "<95>": 151764,
961
+ "<960>": 152629,
962
+ "<961>": 152630,
963
+ "<962>": 152631,
964
+ "<963>": 152632,
965
+ "<964>": 152633,
966
+ "<965>": 152634,
967
+ "<966>": 152635,
968
+ "<967>": 152636,
969
+ "<968>": 152637,
970
+ "<969>": 152638,
971
+ "<96>": 151765,
972
+ "<970>": 152639,
973
+ "<971>": 152640,
974
+ "<972>": 152641,
975
+ "<973>": 152642,
976
+ "<974>": 152643,
977
+ "<975>": 152644,
978
+ "<976>": 152645,
979
+ "<977>": 152646,
980
+ "<978>": 152647,
981
+ "<979>": 152648,
982
+ "<97>": 151766,
983
+ "<980>": 152649,
984
+ "<981>": 152650,
985
+ "<982>": 152651,
986
+ "<983>": 152652,
987
+ "<984>": 152653,
988
+ "<985>": 152654,
989
+ "<986>": 152655,
990
+ "<987>": 152656,
991
+ "<988>": 152657,
992
+ "<989>": 152658,
993
+ "<98>": 151767,
994
+ "<990>": 152659,
995
+ "<991>": 152660,
996
+ "<992>": 152661,
997
+ "<993>": 152662,
998
+ "<994>": 152663,
999
+ "<995>": 152664,
1000
+ "<996>": 152665,
1001
+ "<997>": 152666,
1002
+ "<998>": 152667,
1003
+ "<999>": 152668,
1004
+ "<99>": 151768,
1005
+ "<9>": 151678,
1006
+ "<think>": 151667,
1007
+ "<tool_call>": 151657,
1008
+ "<tool_response>": 151665,
1009
+ "<|box_end|>": 151649,
1010
+ "<|box_start|>": 151648,
1011
+ "<|endoftext|>": 151643,
1012
+ "<|file_sep|>": 151664,
1013
+ "<|fim_middle|>": 151660,
1014
+ "<|fim_pad|>": 151662,
1015
+ "<|fim_prefix|>": 151659,
1016
+ "<|fim_suffix|>": 151661,
1017
+ "<|im_end|>": 151645,
1018
+ "<|im_start|>": 151644,
1019
+ "<|image_pad|>": 151655,
1020
+ "<|object_ref_end|>": 151647,
1021
+ "<|object_ref_start|>": 151646,
1022
+ "<|quad_end|>": 151651,
1023
+ "<|quad_start|>": 151650,
1024
+ "<|repo_name|>": 151663,
1025
+ "<|video_pad|>": 151656,
1026
+ "<|vision_end|>": 151653,
1027
+ "<|vision_pad|>": 151654,
1028
+ "<|vision_start|>": 151652,
1029
+ "|<MASK>|": 152670
1030
+ }
assets/decoding.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5f6fb9d57b88af79cbe78ee9f66cbba22519732daa2068484fc1e449d8fb138a
3
+ size 8080597
assets/demo-poster.jpg ADDED

Git LFS Details

  • SHA256: 288c6121ec41e3114fe98221cf795381b1aace7c99660da945164be88af331f0
  • Pointer size: 131 Bytes
  • Size of remote file: 217 kB
assets/demo.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1d987aac49b390c0411faecc33892873fb0b8b811ff1f961d8a87cb39c406285
3
+ size 22955151
assets/fig1-teaser.png ADDED

Git LFS Details

  • SHA256: 5e426112a17e140f12a9b8a489ffc2a11476db26f389b4b814bf9f9a787867c2
  • Pointer size: 132 Bytes
  • Size of remote file: 6.24 MB
assets/fig2-architecture.png ADDED

Git LFS Details

  • SHA256: 65128fe5b23dbd8521f61b0afc93faefa135a9da8413c33a5e73d654814f6187
  • Pointer size: 131 Bytes
  • Size of remote file: 413 kB
assets/fig4-attention-mask.png ADDED

Git LFS Details

  • SHA256: 71bd095e5e01deae2177d698f277ecc978d0760df078fcaccd4c7713b4d1b3dd
  • Pointer size: 131 Bytes
  • Size of remote file: 372 kB
assets/fig6-self-speculative-decoding.png ADDED

Git LFS Details

  • SHA256: 7b0af08a3c9d62bb996856846b8110dc8ff3bc9b0daa5425b792f975080d866f
  • Pointer size: 131 Bytes
  • Size of remote file: 281 kB
assets/fig7-grounding-performance.png ADDED

Git LFS Details

  • SHA256: f42986c53a440d04156ab1f36088253b2f29b00e899f38bfc046775a8d76e380
  • Pointer size: 131 Bytes
  • Size of remote file: 819 kB
assets/logo.png ADDED

Git LFS Details

  • SHA256: 67e70060c339a3a4444f2081efb9c41b0addf42fa0a2643654dc5dba78c43ae0
  • Pointer size: 131 Bytes
  • Size of remote file: 416 kB
chat_template.jinja ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {% set image_count = namespace(value=0) %}{% set video_count = namespace(value=0) %}{% for message in messages %}{% if loop.first and message['role'] != 'system' %}<|im_start|>system
2
+ You are a helpful assistant.<|im_end|>
3
+ {% endif %}<|im_start|>{{ message['role'] }}
4
+ {% if message['content'] is string %}{{ message['content'] }}<|im_end|>
5
+ {% else %}{% for content in message['content'] %}{% if content['type'] == 'image' or 'image' in content or 'image_url' in content %}{% set image_count.value = image_count.value + 1 %}{% if add_vision_id %}Picture {{ image_count.value }}: {% endif %}<|vision_start|><|image_pad|><|vision_end|>{% elif content['type'] == 'video' or 'video' in content %}{% set video_count.value = video_count.value + 1 %}{% if add_vision_id %}Video {{ video_count.value }}: {% endif %}<|vision_start|><|video_pad|><|vision_end|>{% elif 'text' in content %}{{ content['text'] }}{% endif %}{% endfor %}<|im_end|>
6
+ {% endif %}{% endfor %}{% if add_generation_prompt %}<|im_start|>assistant
7
+ {% endif %}
checksums.sha256 ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ 20c797ce19af0c17de52c6afb144644768a591c521655f5ebf5712c9850f2887 LICENSE
2
+ 235136e6f31b16cd2012fa1279cf97141b5eef6916c775ba075cdc7f42c043ba README.md
3
+ 669bb095cca2e86ddc821926f4c3b42dd7389f2ad3ec49c8edca8fe6294aae7f added_tokens.json
4
+ a0bc6f6fc7a29a80017a433e8f03a1cc1236e838a944a2d034295a60c4f2fddb chat_template.jinja
5
+ 22c369e2bdac19723fe7db7e4c64b224462744b1f1aae7bd5daf05307db3b9fd config.json
6
+ 7ad33532c846fdd81dd41a6ba54db54e3c47cada50472d732ba91eaede725323 configuration_groundinganything.py
7
+ 6b255ff08426effe581833b3809bf2232cac9eb93d1e5cf83fff76ce62791e38 configuration_groundinganything_vision.py
8
+ 42e0cb3d5a6a9e5c93c20555739f1038753ba1adb139c5c82a3a9f3b39ed5eab generation_config.json
9
+ 33cb222671104ec4f0d3d234a3db7561b94d11543089d5f250c14a2f1b8651c3 image_processing_groundinganything.py
10
+ 78403540328f9847d6b7ebc5c44eb2e6a752863de0afb7d0710728bb161dc60d media_utils.py
11
+ 8831e4f1a044471340f7c0a83d7bd71306a5b867e95fd870f74d0c5308a904d5 merges.txt
12
+ 3cf09355cde8cf4877cdc76e0a72301d572c1d6263b2831f8a03368285bb65e2 model.safetensors
13
+ 86ba0522e32bcf5d512f8e7dcc529abd22a39f10d55fa8327efcb65cbd758a67 modeling_groundinganything.py
14
+ a839e10631620ae5fe70ec3aa9d009af521f7ea03faa66b32adce287105003be modeling_groundinganything_vision.py
15
+ 11ddb518de0bcafe3f46f49310887bf7424b4f533d300c5f4749052c9d6465d1 preprocessor_config.json
16
+ 90e533272ff79d068ce922e9156e5d83f3cf2cc9e565af37b135c09e9fe57b81 processing_groundinganything.py
17
+ 534931ef99997dec0735da46a570a45696cfee6dace291cc9fa9a2da24420685 requirements.txt
18
+ d948513463b339ae85d6be6c99ecee787cc9c7a96be09180aa521a32a216eff4 special_tokens_map.json
19
+ 5d9a9d7525aeecc0360ffd43f891ce9d2220297df5092a293a98e2fd31e9d99d streammind_gate.py
20
+ 7e0398aca93659140fd6345deb335db2330dee89a2aad1228d8604c1952f6dc8 tokenizer.json
21
+ 2aa39a017b159c2ac86edb3baf6a78a77549fc7b26ea3cee9c37128b31fec589 tokenizer_config.json
22
+ ca10d7e9fb3ed18575dd1e277a2579c16d108e32f27439684afa0e10b1440910 vocab.json
config.json ADDED
@@ -0,0 +1,141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "GroundAnythingForConditionalGeneration"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_groundinganything.GroundAnythingConfig",
7
+ "AutoModel": "modeling_groundinganything.GroundAnythingModel",
8
+ "AutoModelForCausalLM": "modeling_groundinganything.GroundAnythingForConditionalGeneration",
9
+ "AutoModelForImageTextToText": "modeling_groundinganything.GroundAnythingForConditionalGeneration",
10
+ "AutoProcessor": "processing_groundinganything.GroundAnythingProcessor"
11
+ },
12
+ "bos_token_id": null,
13
+ "dtype": "bfloat16",
14
+ "eos_token_id": 151645,
15
+ "hidden_size": 2560,
16
+ "image_token_id": 151655,
17
+ "model_type": "groundinganything",
18
+ "pad_token_id": 151643,
19
+ "text_config": {
20
+ "_name_or_path": "Qwen3-4B-Instruct-2507",
21
+ "architectures": [
22
+ "Qwen3ForCausalLM"
23
+ ],
24
+ "attention_bias": false,
25
+ "attention_dropout": 0.0,
26
+ "bos_token_id": 151643,
27
+ "dtype": "bfloat16",
28
+ "eos_token_id": 151645,
29
+ "head_dim": 128,
30
+ "hidden_act": "silu",
31
+ "hidden_size": 2560,
32
+ "initializer_range": 0.02,
33
+ "intermediate_size": 9728,
34
+ "layer_types": [
35
+ "full_attention",
36
+ "full_attention",
37
+ "full_attention",
38
+ "full_attention",
39
+ "full_attention",
40
+ "full_attention",
41
+ "full_attention",
42
+ "full_attention",
43
+ "full_attention",
44
+ "full_attention",
45
+ "full_attention",
46
+ "full_attention",
47
+ "full_attention",
48
+ "full_attention",
49
+ "full_attention",
50
+ "full_attention",
51
+ "full_attention",
52
+ "full_attention",
53
+ "full_attention",
54
+ "full_attention",
55
+ "full_attention",
56
+ "full_attention",
57
+ "full_attention",
58
+ "full_attention",
59
+ "full_attention",
60
+ "full_attention",
61
+ "full_attention",
62
+ "full_attention",
63
+ "full_attention",
64
+ "full_attention",
65
+ "full_attention",
66
+ "full_attention",
67
+ "full_attention",
68
+ "full_attention",
69
+ "full_attention",
70
+ "full_attention"
71
+ ],
72
+ "max_position_embeddings": 262144,
73
+ "max_window_layers": 36,
74
+ "model_type": "qwen3",
75
+ "num_attention_heads": 32,
76
+ "num_hidden_layers": 36,
77
+ "num_key_value_heads": 8,
78
+ "pad_token_id": 151643,
79
+ "rms_norm_eps": 1e-06,
80
+ "rope_parameters": {
81
+ "rope_theta": 5000000,
82
+ "rope_type": "default"
83
+ },
84
+ "sliding_window": null,
85
+ "tie_word_embeddings": false,
86
+ "use_cache": false,
87
+ "use_sliding_window": false,
88
+ "vocab_size": 152670
89
+ },
90
+ "tie_word_embeddings": false,
91
+ "transformers_version": "5.7.0",
92
+ "use_cache": false,
93
+ "video_token_id": 151656,
94
+ "vision_config": {
95
+ "activation_func": "gelu_pytorch_tanh",
96
+ "attention_dropout": 0.0,
97
+ "attn_bias": false,
98
+ "dtype": "bfloat16",
99
+ "frame_windows_size": 4,
100
+ "hidden_act": "gelu",
101
+ "hidden_size": 1024,
102
+ "image_size": 448,
103
+ "init_pos_emb_height": 64,
104
+ "init_pos_emb_time": 4,
105
+ "init_pos_emb_width": 64,
106
+ "initializer_range": 0.02,
107
+ "intermediate_size": 4096,
108
+ "layer_norm_eps": 1e-06,
109
+ "layer_norm_type": "layer_norm",
110
+ "linear_bias": false,
111
+ "max_position_embeddings": 8192,
112
+ "merge_kernel_size": [
113
+ 2,
114
+ 2
115
+ ],
116
+ "merge_type": "sd2_tpool",
117
+ "mlp_type": "mlp2",
118
+ "model_type": "groundinganything_vision",
119
+ "norm_type": "rmsnorm",
120
+ "num_attention_heads": 12,
121
+ "num_channels": 3,
122
+ "num_hidden_layers": 27,
123
+ "out_hidden_size": 2560,
124
+ "patch_embed_proj_bias": false,
125
+ "patch_position_encoding_type": "absolute",
126
+ "patch_size": 14,
127
+ "pos_emb_interpolation_mode": "bilinear",
128
+ "pos_emb_type": "divided_fixed",
129
+ "projector_hidden_act": "gelu",
130
+ "projector_hidden_size": 4096,
131
+ "projector_ln_eps": 1e-05,
132
+ "qkv_hidden_size": 1536,
133
+ "rope_theta": 10000.0,
134
+ "spatial_merge_size": 2,
135
+ "tokens_per_second": 1,
136
+ "use_head": false,
137
+ "use_patch_position_encoding": false
138
+ },
139
+ "vision_end_token_id": 151653,
140
+ "vision_start_token_id": 151652
141
+ }
configuration_groundinganything.py ADDED
@@ -0,0 +1,171 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from dataclasses import dataclass
2
+
3
+ from transformers.models.qwen3.configuration_qwen3 import Qwen3Config
4
+ try:
5
+ from transformers.configuration_utils import PreTrainedConfig
6
+ except ImportError:
7
+ from transformers.configuration_utils import PretrainedConfig as PreTrainedConfig
8
+
9
+
10
+ @dataclass(init=False)
11
+ class GroundAnythingVLMVisionConfig(PreTrainedConfig):
12
+ model_type = "groundinganything_vision"
13
+ base_config_key = "vision_config"
14
+
15
+ def __init__(self, **kwargs):
16
+ super().__init__(**kwargs)
17
+
18
+ hidden_size: int = 1024
19
+ intermediate_size: int = 4096
20
+ num_hidden_layers: int = 24
21
+ num_attention_heads: int = 16
22
+ num_channels: int = 3
23
+ image_size: int = 448
24
+ patch_size: int = 14
25
+ hidden_act: str = "gelu"
26
+ layer_norm_eps: float = 1e-6
27
+ layer_norm_type: str = "layer_norm"
28
+ attention_dropout: float = 0.0
29
+ initializer_range: float = 0.02
30
+ rope_theta: float = 10000.0
31
+ use_head: bool = False
32
+ out_hidden_size: int = 1024
33
+ spatial_merge_size: int = 2
34
+ tokens_per_second: int = 1
35
+ frame_windows_size: int = 4
36
+ use_patch_position_encoding: bool = False
37
+ patch_position_encoding_type: str = "absolute"
38
+ max_position_embeddings: int = 8192
39
+ init_pos_emb_height: int = 64
40
+ init_pos_emb_width: int = 64
41
+ init_pos_emb_time: int = 4
42
+ pos_emb_type: str = "divided_fixed"
43
+ merge_kernel_size: list | None = None
44
+ merge_type: str = "sd2_tpool"
45
+ qkv_hidden_size: int = 1536
46
+ norm_type: str = "rmsnorm"
47
+ attn_bias: bool = False
48
+ patch_embed_proj_bias: bool = False
49
+ mlp_type: str = "mlp2"
50
+ linear_bias: bool = False
51
+ activation_func: str = "gelu_pytorch_tanh"
52
+ pos_emb_interpolation_mode: str = "bilinear"
53
+ projector_hidden_size: int = 4096
54
+ projector_hidden_act: str = "gelu"
55
+ projector_ln_eps: float = 1e-5
56
+
57
+
58
+ @dataclass(init=False)
59
+ class GroundAnythingVLMConfig(PreTrainedConfig):
60
+ r"""
61
+ This is the configuration class to store the configuration of a [`GroundAnythingVLMBaseModel`]. It is used to instantiate a
62
+ GroundAnythingVLMBaseModel model according to the specified arguments, defining the model architecture. Instantiating a configuration
63
+ with the defaults will yield a GroundAnything-VLM configuration.
64
+
65
+ Configuration objects inherit from [`PreTrainedConfig`] and can be used to control the model outputs. Read the
66
+ documentation from [`PreTrainedConfig`] for more information.
67
+
68
+ Args:
69
+ text_config (`Union[PreTrainedConfig, dict]`, *optional*, defaults to `Qwen3Config`):
70
+ The config object or dictionary of the text backbone.
71
+ vision_config (`Union[PreTrainedConfig, dict]`, *optional*, defaults to `GroundAnythingVLMVisionConfig`):
72
+ The config object or dictionary of the vision backbone.
73
+ image_token_id (`int`, *optional*, defaults to 151655):
74
+ The image token index to encode the image prompt.
75
+ video_token_id (`int`, *optional*, defaults to 151656):
76
+ The video token index to encode the image prompt.
77
+ vision_start_token_id (`int`, *optional*, defaults to 151652):
78
+ The token index to denote start of vision input.
79
+ vision_end_token_id (`int`, *optional*, defaults to 151653):
80
+ The token index to denote end of vision input.
81
+ """
82
+
83
+ model_type = "groundinganything_vlm"
84
+ # `text_config` is resolved dynamically based on its `model_type` (defaults to `qwen3`),
85
+ # so we use `AutoConfig` here as a placeholder; `__post_init__` swaps it for the
86
+ # concrete config class via `CONFIG_MAPPING`.
87
+ sub_configs = {"vision_config": GroundAnythingVLMVisionConfig, "text_config": Qwen3Config}
88
+ keys_to_ignore_at_inference = ["past_key_values"]
89
+
90
+
91
+ def __init__(self, **kwargs):
92
+ self.text_config = kwargs.pop("text_config", None)
93
+ self.vision_config = kwargs.pop("vision_config", None)
94
+ self.image_token_id = kwargs.pop("image_token_id", 151655)
95
+ self.video_token_id = kwargs.pop("video_token_id", 151656)
96
+ self.vision_start_token_id = kwargs.pop("vision_start_token_id", 151652)
97
+ self.vision_end_token_id = kwargs.pop("vision_end_token_id", 151653)
98
+ self.tie_word_embeddings = kwargs.pop("tie_word_embeddings", False)
99
+ self.bos_token_id = kwargs.pop("bos_token_id", None)
100
+ self.eos_token_id = kwargs.pop("eos_token_id", None)
101
+ self.pad_token_id = kwargs.pop("pad_token_id", None)
102
+ if isinstance(self.vision_config, dict):
103
+ self.vision_config = GroundAnythingVLMVisionConfig(**self.vision_config)
104
+ if isinstance(self.text_config, dict):
105
+ text_model_type = self.text_config.get("model_type", "qwen3")
106
+ if text_model_type != "qwen3":
107
+ raise ValueError(f"unsupported text model type: {text_model_type}")
108
+ text_config_cls = Qwen3Config
109
+ self.sub_configs["text_config"] = text_config_cls
110
+ self.text_config = text_config_cls(**self.text_config)
111
+ # Transformers 5.x uses a generated dataclass __init__ that dispatches
112
+ # to this subclass' __post_init__. Call the base hook explicitly so its
113
+ # private attention/config state is initialized. Keep 4.x compatible.
114
+ if hasattr(PreTrainedConfig, "__post_init__"):
115
+ PreTrainedConfig.__post_init__(self, **kwargs)
116
+ else:
117
+ super().__init__(**kwargs)
118
+ self.__post_init__()
119
+
120
+ text_config: dict | PreTrainedConfig | None = None
121
+ vision_config: dict | PreTrainedConfig | None = None
122
+ image_token_id: int = 151655
123
+ video_token_id: int = 151656
124
+ vision_start_token_id: int = 151652
125
+ vision_end_token_id: int = 151653
126
+ tie_word_embeddings: bool = False
127
+ # Generation-related token ids are mirrored from `text_config` in `__post_init__`
128
+ # so downstream tools (e.g. `generate`, vLLM) that read them at the top level keep working.
129
+ bos_token_id: int | None = None
130
+ eos_token_id: int | list[int] | None = None
131
+ pad_token_id: int | None = None
132
+
133
+ def __post_init__(self, **kwargs):
134
+ # Resolve vision_config
135
+ if isinstance(self.vision_config, dict):
136
+ self.vision_config = self.sub_configs["vision_config"](**self.vision_config)
137
+ elif self.vision_config is None:
138
+ self.vision_config = self.sub_configs["vision_config"]()
139
+
140
+ # Resolve text_config dynamically via CONFIG_MAPPING (defaults to qwen3)
141
+ if isinstance(self.text_config, dict):
142
+ text_model_type = self.text_config.get("model_type", "qwen3")
143
+ self.text_config["model_type"] = text_model_type
144
+ if text_model_type != "qwen3":
145
+ raise ValueError(f"unsupported text model type: {text_model_type}")
146
+ text_config_cls = Qwen3Config
147
+ self.sub_configs["text_config"] = text_config_cls
148
+ self.text_config = text_config_cls(**self.text_config)
149
+ elif self.text_config is None:
150
+ text_config_cls = Qwen3Config
151
+ self.sub_configs["text_config"] = text_config_cls
152
+ self.text_config = text_config_cls()
153
+
154
+ # Mirror generation-related token ids from text_config to the top level so
155
+ # downstream tools (e.g. `generate`, chat templates, vLLM) that read them
156
+ # from the top-level config keep working.
157
+ for tok_key in ("bos_token_id", "eos_token_id", "pad_token_id"):
158
+ text_val = getattr(self.text_config, tok_key, None)
159
+ if text_val is not None and getattr(self, tok_key, None) is None:
160
+ setattr(self, tok_key, text_val)
161
+
162
+ del kwargs
163
+
164
+
165
+ __all__ = ["GroundAnythingVLMConfig", "GroundAnythingVLMVisionConfig", "GroundAnythingConfig"]
166
+
167
+
168
+ class GroundAnythingConfig(GroundAnythingVLMConfig):
169
+ """DLM release identity for the shared Qwen3 VLM configuration."""
170
+
171
+ model_type = "groundinganything"
configuration_groundinganything_vision.py ADDED
@@ -0,0 +1,285 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Optional
2
+
3
+ from transformers.configuration_utils import PretrainedConfig
4
+
5
+
6
+ class GroundAnythingBackboneLinearConfig(PretrainedConfig):
7
+ model_type = "kimi_linear"
8
+ keys_to_ignore_at_inference = ["past_key_values"]
9
+
10
+ def __init__(
11
+ self,
12
+ model_type="kimi_linear",
13
+ vocab_size=163840,
14
+ hidden_size=4096,
15
+ head_dim=None,
16
+ intermediate_size=11008,
17
+ num_hidden_layers=32,
18
+ num_attention_heads=32,
19
+ num_key_value_heads=None,
20
+ hidden_act="silu",
21
+ initializer_range=0.02,
22
+ rms_norm_eps=1e-6,
23
+ use_cache=True,
24
+ pad_token_id=0,
25
+ bos_token_id=1,
26
+ eos_token_id=2,
27
+ rope_theta=10000.0,
28
+ rope_scaling=None,
29
+ tie_word_embeddings=False,
30
+ moe_intermediate_size: Optional[int] = None,
31
+ moe_renormalize: bool = True,
32
+ moe_router_activation_func: str = "sigmoid",
33
+ num_experts: Optional[int] = None,
34
+ num_experts_per_token: Optional[int] = None,
35
+ num_shared_experts: int = 0,
36
+ routed_scaling_factor: float = 1.0,
37
+ first_k_dense_replace: int = 0,
38
+ moe_layer_freq: int = 1,
39
+ use_grouped_topk: bool = True,
40
+ num_expert_group: int = 1,
41
+ topk_group: int = 1,
42
+ q_lora_rank: Optional[int] = None,
43
+ kv_lora_rank: Optional[int] = None,
44
+ qk_nope_head_dim: Optional[int] = None,
45
+ qk_rope_head_dim: Optional[int] = None,
46
+ v_head_dim: Optional[int] = None,
47
+ mla_use_nope: Optional[bool] = False,
48
+ mla_use_output_gate: Optional[bool] = False,
49
+ num_nextn_predict_layers: int = 0,
50
+ linear_attn_config: Optional[dict] = None,
51
+ attn_res_block_size: Optional[int] = None,
52
+ latent_moe_use_norm: bool = False,
53
+ activation_situ_beta: Optional[float] = None,
54
+ activation_situ_linear_beta: Optional[float] = None,
55
+ max_position_embeddings: int = 4096,
56
+ routed_expert_hidden_size: Optional[int] = None,
57
+ topk_method: str = "noaux_tc",
58
+ **kwargs,
59
+ ):
60
+ self.model_type = model_type
61
+ self.vocab_size = vocab_size
62
+ self.hidden_size = hidden_size
63
+ self.head_dim = (
64
+ head_dim if head_dim is not None else hidden_size // num_attention_heads
65
+ )
66
+ self.intermediate_size = intermediate_size
67
+ self.num_hidden_layers = num_hidden_layers
68
+ self.num_attention_heads = num_attention_heads
69
+
70
+ # for backward compatibility
71
+ if num_key_value_heads is None:
72
+ num_key_value_heads = num_attention_heads
73
+
74
+ self.num_key_value_heads = num_key_value_heads
75
+ self.hidden_act = hidden_act
76
+ self.initializer_range = initializer_range
77
+ self.rms_norm_eps = rms_norm_eps
78
+ self.use_cache = use_cache
79
+ self.rope_theta = rope_theta
80
+ self.rope_scaling = rope_scaling
81
+
82
+ self.q_lora_rank = q_lora_rank
83
+ self.kv_lora_rank = kv_lora_rank
84
+ self.qk_nope_head_dim = qk_nope_head_dim
85
+ self.qk_rope_head_dim = qk_rope_head_dim
86
+ self.v_head_dim = v_head_dim
87
+ self.mla_use_nope = mla_use_nope
88
+ self.mla_use_output_gate = mla_use_output_gate
89
+ # moe config
90
+ self.num_experts = num_experts
91
+ self.num_experts_per_token = num_experts_per_token
92
+ self.moe_renormalize = moe_renormalize
93
+ self.num_shared_experts = num_shared_experts
94
+ self.routed_scaling_factor = routed_scaling_factor
95
+ self.moe_router_activation_func = moe_router_activation_func
96
+ assert self.moe_router_activation_func in ("softmax", "sigmoid")
97
+ self.moe_intermediate_size = moe_intermediate_size
98
+ self.first_k_dense_replace = first_k_dense_replace
99
+ self.moe_layer_freq = moe_layer_freq
100
+ self.use_grouped_topk = use_grouped_topk
101
+ self.num_expert_group = num_expert_group
102
+ self.topk_group = topk_group
103
+ self.num_nextn_predict_layers = num_nextn_predict_layers
104
+
105
+ self.attn_res_block_size = attn_res_block_size
106
+ self.latent_moe_use_norm = latent_moe_use_norm
107
+ self.activation_situ_beta = activation_situ_beta
108
+ self.activation_situ_linear_beta = activation_situ_linear_beta
109
+ self.max_position_embeddings = max_position_embeddings
110
+ self.routed_expert_hidden_size = routed_expert_hidden_size
111
+ self.topk_method = topk_method
112
+
113
+ if linear_attn_config is not None:
114
+ assert linear_attn_config["kda_layers"] is not None
115
+ assert linear_attn_config["full_attn_layers"] is not None
116
+ self.linear_attn_config = linear_attn_config
117
+
118
+ super().__init__(
119
+ pad_token_id=pad_token_id,
120
+ bos_token_id=bos_token_id,
121
+ eos_token_id=eos_token_id,
122
+ tie_word_embeddings=tie_word_embeddings,
123
+ **kwargs,
124
+ )
125
+
126
+ @property
127
+ def is_mla(self):
128
+ return (
129
+ self.q_lora_rank is not None
130
+ or self.kv_lora_rank is not None
131
+ or self.qk_nope_head_dim is not None
132
+ or self.qk_rope_head_dim is not None
133
+ or self.v_head_dim is not None
134
+ or self.mla_use_nope is True
135
+ )
136
+
137
+ @property
138
+ def is_moe(self):
139
+ return self.num_experts is not None
140
+
141
+ @property
142
+ def is_linear_attn(self) -> bool:
143
+ return not (
144
+ self.linear_attn_config is None
145
+ or (
146
+ isinstance(self.linear_attn_config, dict)
147
+ and self.linear_attn_config["kda_layers"] is not None
148
+ and len(self.linear_attn_config["kda_layers"]) == 0
149
+ )
150
+ )
151
+
152
+ def is_kda_layer(self, layer_idx: int):
153
+ return (
154
+ self.linear_attn_config is not None
155
+ and (layer_idx + 1) in self.linear_attn_config["kda_layers"]
156
+ )
157
+
158
+
159
+ class GroundAnythingBackboneVisionConfig(PretrainedConfig):
160
+
161
+ def __init__(
162
+ self,
163
+ patch_size: int = 14,
164
+ init_pos_emb_height: int = 64,
165
+ init_pos_emb_width: int = 64,
166
+ init_pos_emb_time: int = 4,
167
+ pos_emb_type: str = 'divided_fixed',
168
+ vt_num_attention_heads: int = 12,
169
+ vt_num_hidden_layers: int = 27,
170
+ vt_hidden_size: int = 1024,
171
+ vt_intermediate_size: int = 4096,
172
+ merge_kernel_size: tuple = (2, 2),
173
+ merge_type: str = 'sd2_tpool',
174
+ _attn_implementation: str = 'flash_attention_2',
175
+ # MM Projector parameters
176
+ mm_projector_type: str = 'patchmergerv2',
177
+ mm_hidden_size: int | None = None,
178
+ projector_hidden_act: str = "gelu",
179
+ projector_ln_eps: float = 1e-5,
180
+ # vision tower parameters
181
+ qkv_hidden_size: int = 1536,
182
+ norm_type: str = 'rmsnorm',
183
+ attn_bias: bool = False,
184
+ patch_embed_proj_bias: bool = False,
185
+ mlp_type: str = 'mlp2',
186
+ linear_bias: bool = False,
187
+ activation_func: str = 'gelu_pytorch_tanh',
188
+ pos_emb_interpolation_mode: str = 'bilinear',
189
+ # Other parameters
190
+ ignore_index: int = -100,
191
+ media_placeholder_token_id: int = 163605,
192
+ pad_token_id: int = 0,
193
+ text_hidden_size=7168,
194
+ **kwargs):
195
+
196
+ self.patch_size = patch_size
197
+ self.init_pos_emb_height = init_pos_emb_height
198
+ self.init_pos_emb_width = init_pos_emb_width
199
+ self.init_pos_emb_time = init_pos_emb_time
200
+ self.pos_emb_type = pos_emb_type
201
+ self.vt_num_attention_heads = vt_num_attention_heads
202
+ self.vt_num_hidden_layers = vt_num_hidden_layers
203
+ self.vt_hidden_size = vt_hidden_size
204
+ self.vt_intermediate_size = vt_intermediate_size
205
+ self.merge_kernel_size = merge_kernel_size
206
+ self.merge_type = merge_type
207
+ self._attn_implementation = _attn_implementation
208
+
209
+ # MM Projector config
210
+ self.mm_projector_type = mm_projector_type
211
+ self.mm_hidden_size = mm_hidden_size if mm_hidden_size is not None else vt_hidden_size
212
+ self.projector_hidden_act = projector_hidden_act
213
+ self.projector_ln_eps = projector_ln_eps
214
+ self.text_hidden_size = text_hidden_size
215
+
216
+ # vision tower parameters
217
+ self.qkv_hidden_size = qkv_hidden_size
218
+ self.norm_type = norm_type
219
+ self.attn_bias = attn_bias
220
+ self.patch_embed_proj_bias = patch_embed_proj_bias
221
+ self.mlp_type = mlp_type
222
+ self.linear_bias = linear_bias
223
+ self.activation_func = activation_func
224
+ self.pos_emb_interpolation_mode = pos_emb_interpolation_mode
225
+
226
+ super().__init__(**kwargs)
227
+
228
+
229
+ class GroundAnythingBackboneConfig(PretrainedConfig):
230
+ """Kimi-K3 model configuration.
231
+
232
+ Args:
233
+ text_config (dict | GroundAnythingBackboneLinearConfig): Configuration for the text model.
234
+
235
+ Vision Tower Parameters (from MoonViT3dConfig):
236
+ patch_size (int): Patch size for vision tower.
237
+ init_pos_emb_height (int): Initial position embedding height.
238
+ init_pos_emb_width (int): Initial position embedding width.
239
+ init_pos_emb_time (int): Initial position embedding time dimension.
240
+ pos_emb_type (str): Type of position embedding.
241
+ vt_num_attention_heads (int): Number of attention heads in vision tower.
242
+ vt_num_hidden_layers (int): Number of hidden layers in vision tower.
243
+ vt_hidden_size (int): Hidden size of vision tower.
244
+ vt_intermediate_size (int): Intermediate size in vision tower FFN.
245
+ merge_kernel_size (tuple): Kernel size for patch merging.
246
+ merge_type (str): Type of merge operation.
247
+ _attn_implementation (str): Attention implementation type.
248
+
249
+ MM Projector Parameters (from MultiModalProjectorConfig):
250
+ mm_projector_type (str): Type of multimodal projector.
251
+ mm_hidden_size (int): Hidden size from vision tower (should match vt_hidden_size).
252
+ projector_hidden_act (str): Activation function for projector.
253
+ projector_ln_eps (float): Layer norm epsilon for projector.
254
+
255
+ Other Parameters:
256
+ ignore_index (int): The ignore index for the loss function.
257
+ media_placeholder_token_id (int): The token ID to use for media placeholders.
258
+ pad_token_id (int): The token ID to use for padding.
259
+ """
260
+
261
+ model_type = "groundinganything_backbone"
262
+
263
+ def __init__(
264
+ self,
265
+ text_config: dict | GroundAnythingBackboneLinearConfig = None,
266
+ vision_config: dict | GroundAnythingBackboneVisionConfig = None,
267
+ # Other parameters
268
+ ignore_index: int = -100,
269
+ media_placeholder_token_id: int = 163605,
270
+ pad_token_id: int = 0,
271
+ **kwargs,
272
+ ):
273
+ if isinstance(text_config, dict):
274
+ text_config = GroundAnythingBackboneLinearConfig(**text_config)
275
+ if isinstance(vision_config, dict):
276
+ vision_config = GroundAnythingBackboneVisionConfig(**vision_config)
277
+ self.text_config = text_config
278
+ self.vision_config = vision_config
279
+ # Other config
280
+ self.ignore_index = ignore_index
281
+ self.media_placeholder_token_id = media_placeholder_token_id
282
+ if getattr(self.text_config, "quantization_config", None) is not None:
283
+ self.quantization_config = self.text_config.quantization_config
284
+
285
+ super().__init__(pad_token_id=pad_token_id, **kwargs)
generation_config.json ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 151643,
4
+ "eos_token_id": [
5
+ 151645,
6
+ 151643
7
+ ],
8
+ "output_attentions": false,
9
+ "output_hidden_states": false,
10
+ "transformers_version": "5.7.0",
11
+ "use_cache": true
12
+ }
image_processing_groundinganything.py ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Image processor class for Kimi-K3.
2
+ """
3
+
4
+ import json
5
+ import os
6
+ from typing import Any, Dict, Optional, Union
7
+
8
+ import numpy as np
9
+ import torch
10
+ from PIL import Image
11
+ from transformers.image_processing_utils import (BaseImageProcessor,
12
+ BatchFeature)
13
+ from transformers.utils import TensorType
14
+
15
+ from .media_utils import (MediaInput, TransparentBgConfig, _to_tensor,
16
+ ensure_media_type, image_to_np, navit_patchify,
17
+ navit_resize_image, normalize)
18
+
19
+
20
+ class GroundAnythingVLMImageProcessor(BaseImageProcessor):
21
+ model_type = "groundinganything_backbone"
22
+
23
+ def __init__(
24
+ self,
25
+ media_proc_cfg: dict,
26
+ **kwargs,
27
+ ):
28
+ super().__init__(**kwargs)
29
+ self.media_proc_cfg = dict(media_proc_cfg)
30
+ self.patch_size = int(self.media_proc_cfg.get("patch_size", 14))
31
+ self.merge_size = int(self.media_proc_cfg.get("merge_kernel_size", 2))
32
+ max_output_tokens = os.getenv("IMAGE_MAX_TOKEN_NUM")
33
+ if max_output_tokens:
34
+ max_input_patches = int(max_output_tokens) * self.merge_size**2
35
+ self.media_proc_cfg["in_patch_limit"] = min(
36
+ int(self.media_proc_cfg["in_patch_limit"]), max_input_patches
37
+ )
38
+ self.max_output_tokens = (
39
+ int(self.media_proc_cfg["in_patch_limit"]) // self.merge_size**2
40
+ )
41
+
42
+ @property
43
+ def _transparent_bg_config(self) -> Optional[TransparentBgConfig]:
44
+ cfg = self.media_proc_cfg.get("transparent_bg_config")
45
+ if cfg is None:
46
+ return None
47
+ if isinstance(cfg, TransparentBgConfig):
48
+ return cfg
49
+ return TransparentBgConfig(**cfg)
50
+
51
+ @property
52
+ def _transparent_bg_fill_stage(self) -> str:
53
+ return self.media_proc_cfg.get("transparent_bg_fill_stage",
54
+ "before_resize")
55
+
56
+ def media_tokens_calculator(self, media: MediaInput):
57
+ media = ensure_media_type(
58
+ media,
59
+ transparent_bg_config=self._transparent_bg_config,
60
+ transparent_bg_fill_stage=self._transparent_bg_fill_stage,
61
+ )
62
+ ret = self.get_resize_config(media)
63
+ return ret['num_tokens']
64
+
65
+ @classmethod
66
+ def make_image_prompt(cls, width: int, height: int) -> str:
67
+ """Build the K3 image placeholder with resolution info."""
68
+ return (f"<|media_begin|>image {width}x{height}"
69
+ f"<|media_content|><|media_pad|><|media_end|>")
70
+
71
+ def get_resize_config(self, media_input: MediaInput) -> dict:
72
+ if media_input['type'] == 'image':
73
+ w, h = media_input['image'].size
74
+ input_patch_limit = int(self.media_proc_cfg['in_patch_limit'])
75
+ while True:
76
+ ret = navit_resize_image(
77
+ w, h, self.media_proc_cfg['patch_size'],
78
+ self.media_proc_cfg['merge_kernel_size'],
79
+ input_patch_limit,
80
+ self.media_proc_cfg['patch_limit_on_one_side'],
81
+ self.media_proc_cfg['fixed_output_tokens'])
82
+ if ret['num_tokens'] <= self.max_output_tokens:
83
+ return ret
84
+ reduced_limit = int(
85
+ input_patch_limit * self.max_output_tokens /
86
+ ret['num_tokens'])
87
+ input_patch_limit = min(input_patch_limit - 1,
88
+ reduced_limit)
89
+ else:
90
+ raise ValueError("Unsupported type: {}".format(
91
+ media_input['type']))
92
+
93
+ def resize_image(self, image: Image.Image, new_width: int, new_height: int,
94
+ pad_width: int, pad_height: int) -> np.ndarray:
95
+ image_np = image_to_np(
96
+ image,
97
+ (new_width, new_height),
98
+ "resize",
99
+ transparent_bg_config=self._transparent_bg_config,
100
+ transparent_bg_fill_stage=self._transparent_bg_fill_stage,
101
+ )
102
+ image_np = np.pad(
103
+ image_np,
104
+ ((0, pad_height), (0, pad_width), (0, 0)),
105
+ mode="constant",
106
+ constant_values=0,
107
+ )
108
+ return image_np
109
+
110
+ def preprocess(
111
+ self,
112
+ medias: Optional[list[MediaInput]] = None,
113
+ return_tensors: Optional[Union[str, TensorType]] = None,
114
+ images=None,
115
+ do_resize: Optional[bool] = None,
116
+ **kwargs,
117
+ ) -> BatchFeature:
118
+ """
119
+ Preprocess a atom vision input (images) into model-ready tensors.
120
+
121
+ Args:
122
+ medias: List of MediaInput.
123
+ return_tensors: Desired output format ('pt', 'np', 'tf', or None).
124
+
125
+ Returns:
126
+ BatchFeature containing 'pixel_values' and 'grid_thws' tensors.
127
+ """
128
+ del do_resize, kwargs
129
+ if medias is None:
130
+ medias = images
131
+ if medias is None:
132
+ medias = []
133
+ if not isinstance(medias, list):
134
+ medias = [medias]
135
+ medias = [item if isinstance(item, dict) else {"type": "image", "image": item} for item in medias]
136
+ if medias:
137
+ pixel_values = []
138
+ for item in medias:
139
+ item = ensure_media_type(
140
+ item,
141
+ transparent_bg_config=self._transparent_bg_config,
142
+ transparent_bg_fill_stage=self._transparent_bg_fill_stage,
143
+ )
144
+ resize_config = self.get_resize_config(item)
145
+ new_width, new_height, pad_width, pad_height = resize_config[
146
+ 'new_width'], resize_config['new_height'], resize_config[
147
+ 'pad_width'], resize_config['pad_height']
148
+ if item['type'] == 'image':
149
+ image = item['image']
150
+ image_np = self.resize_image(image, new_width, new_height,
151
+ pad_width, pad_height)
152
+ pixel_values.append(np.expand_dims(image_np, axis=0))
153
+ else:
154
+ raise ValueError("Unsupported type: {}".format(
155
+ item['type']))
156
+ normalized_pixel_values = []
157
+ image_std_inv = 1.0 / np.array(self.media_proc_cfg['image_std'])
158
+ image_mean = np.array(self.media_proc_cfg['image_mean'])
159
+ for pixels in pixel_values:
160
+ pixels = normalize(pixels, image_mean, image_std_inv)
161
+ pixels_and_thw = navit_patchify(
162
+ pixels,
163
+ self.media_proc_cfg['patch_size'],
164
+ )
165
+ normalized_pixel_values.append(pixels_and_thw)
166
+
167
+ pixel_values = torch.cat([
168
+ _to_tensor(pixel_value['pixel_values'])
169
+ for pixel_value in normalized_pixel_values
170
+ ])
171
+ grid_thws = torch.cat([
172
+ _to_tensor(pixel_value['grid_thw'],
173
+ dtype=torch.int64).unsqueeze(0)
174
+ for pixel_value in normalized_pixel_values
175
+ ])
176
+
177
+ data = {
178
+ 'pixel_values': pixel_values,
179
+ 'image_grid_thw': grid_thws,
180
+ }
181
+
182
+ else:
183
+ data = {}
184
+
185
+ return BatchFeature(data=data, tensor_type=return_tensors)
186
+
187
+ def __repr__(self):
188
+ return f"GroundAnythingVLMImageProcessor(media_proc_cfg={self.media_proc_cfg})"
189
+
190
+ def to_dict(self) -> Dict[str, Any]:
191
+ output = super().to_dict()
192
+ output["media_proc_cfg"] = self.media_proc_cfg
193
+ if "media_processor" in output:
194
+ del output["media_processor"]
195
+ return output
196
+
197
+ @classmethod
198
+ def from_dict(cls, config_dict: Dict[str, Any], **kwargs):
199
+ config = config_dict.copy()
200
+ media_proc_cfg = config.pop("media_proc_cfg", {})
201
+ return cls(media_proc_cfg=media_proc_cfg, **config, **kwargs)
202
+
203
+ def to_json_string(self):
204
+ dictionary = self.to_dict()
205
+ for key, value in dictionary.items():
206
+ if hasattr(value, 'tolist'):
207
+ dictionary[key] = value.tolist()
208
+ return json.dumps(dictionary, indent=2, sort_keys=True) + "\n"
media_utils.py ADDED
@@ -0,0 +1,376 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import functools
3
+ import io
4
+ import math
5
+ from dataclasses import dataclass
6
+ from typing import Literal, TypedDict
7
+
8
+ import numpy as np
9
+ from PIL import Image
10
+
11
+
12
+ class ImageInput(TypedDict):
13
+ type: Literal['image']
14
+ image: Image.Image
15
+
16
+
17
+ MediaInput = ImageInput
18
+
19
+
20
+ @dataclass
21
+ class TransparentBgConfig:
22
+ """The config of the transparent background."""
23
+
24
+ pattern: Literal["white", "black", "gray", "chessboard"] = "black"
25
+ """The pattern of the transparent background."""
26
+
27
+ chessboard_square_size: int = 16
28
+ """The size of the squares in the chessboard background."""
29
+
30
+ chessboard_square_on_top_left: bool = True
31
+ """Whether to start the chessboard with a white square on the top left."""
32
+
33
+ chessboard_white_value: int = 255
34
+ """The value of the white pixels in the background."""
35
+
36
+ chessboard_gray_value: int = 200
37
+ """The value of the gray pixels in the background."""
38
+
39
+
40
+ @functools.lru_cache(maxsize=256)
41
+ def _create_chessboard_background(
42
+ height: int,
43
+ width: int,
44
+ square_size: int,
45
+ square_on_top_left: bool,
46
+ white_value: int,
47
+ gray_value: int,
48
+ ) -> np.ndarray:
49
+ """Create a chessboard background."""
50
+ bg = np.ones((height, width, 3), dtype=np.uint8) * white_value
51
+ for y in range(0, height, square_size):
52
+ for x in range(0, width, square_size):
53
+ if (y // square_size + x // square_size) % 2 == (
54
+ 1 if square_on_top_left else 0):
55
+ bg[y:y + square_size, x:x + square_size] = gray_value
56
+ return bg
57
+
58
+
59
+ def fill_transparent_bg_with(
60
+ image: Image.Image,
61
+ transparent_bg_config: TransparentBgConfig | None = None,
62
+ ) -> Image.Image:
63
+ """Composite a (possibly) transparent image onto a configured background.
64
+
65
+ When ``transparent_bg_config`` is ``None``, the image is simply converted
66
+ to RGB (preserving the historical behavior). Otherwise the alpha channel
67
+ is alpha-composited over a background generated according to the config.
68
+ """
69
+ if transparent_bg_config is None:
70
+ return image.convert("RGB")
71
+
72
+ if image.mode == "RGB":
73
+ return image
74
+
75
+ has_alpha = "A" in image.getbands() or "transparency" in image.info
76
+ if not has_alpha:
77
+ return image.convert("RGB")
78
+
79
+ img = np.array(image.convert("RGBA"))
80
+ height, width = img.shape[:2]
81
+ bg_pattern = transparent_bg_config.pattern
82
+ if bg_pattern == "white":
83
+ bg = np.full((height, width, 3), 255, dtype=np.uint8)
84
+ elif bg_pattern == "black":
85
+ bg = np.zeros((height, width, 3), dtype=np.uint8)
86
+ elif bg_pattern == "gray":
87
+ bg = np.full((height, width, 3), 128, dtype=np.uint8)
88
+ elif bg_pattern == "chessboard":
89
+ bg = _create_chessboard_background(
90
+ height,
91
+ width,
92
+ transparent_bg_config.chessboard_square_size,
93
+ transparent_bg_config.chessboard_square_on_top_left,
94
+ transparent_bg_config.chessboard_white_value,
95
+ transparent_bg_config.chessboard_gray_value,
96
+ )
97
+ else:
98
+ raise ValueError(f"Invalid background pattern: {bg_pattern}")
99
+
100
+ alpha = img[:, :, 3]
101
+ img_rgb = img[:, :, :3]
102
+ alpha_normalized = alpha.astype(np.float32) / 255.0
103
+ alpha_3d = np.stack([alpha_normalized] * 3, axis=2)
104
+ result = alpha_3d * img_rgb + (1 - alpha_3d) * bg
105
+ result = result.astype(np.uint8)
106
+ return Image.fromarray(result)
107
+
108
+
109
+ def navit_resize_image(
110
+ width: int,
111
+ height: int,
112
+ patch_size: int,
113
+ merge_kernel_size: int,
114
+ in_patch_limit: int,
115
+ patch_limit_on_one_side: int,
116
+ fixed_output_tokens: int | None,
117
+ ):
118
+ # Apply the patch limits.
119
+ s1 = math.sqrt(
120
+ in_patch_limit /
121
+ (max(1.0, width // patch_size) * max(1.0, height // patch_size)))
122
+ s2 = patch_limit_on_one_side * patch_size / width
123
+ s3 = patch_limit_on_one_side * patch_size / height
124
+ scale = min(1.0, s1, s2, s3)
125
+ new_w, new_h = max(1, int(width * scale)), max(1, int(height * scale))
126
+ new_w = min(new_w, patch_limit_on_one_side * patch_size)
127
+ new_h = min(new_h, patch_limit_on_one_side * patch_size)
128
+
129
+ # Calculate the padding to make the height and width divisible by the merge kernel size and patch size.
130
+ factor = merge_kernel_size * patch_size
131
+
132
+ pad_height = (factor - new_h % factor) % factor
133
+ pad_width = (factor - new_w % factor) % factor
134
+
135
+ if fixed_output_tokens is not None:
136
+ num_tokens = fixed_output_tokens
137
+ else:
138
+ # Calculate new dimensions after padding and patching
139
+ token_height = (new_h + pad_height) // factor
140
+ token_width = (new_w + pad_width) // factor
141
+
142
+ assert token_height * merge_kernel_size <= patch_limit_on_one_side, (
143
+ f"token_height {token_height} * merge_kernel_size {merge_kernel_size} > patch_limit_on_one_side {patch_limit_on_one_side}"
144
+ )
145
+ assert token_width * merge_kernel_size <= patch_limit_on_one_side, (
146
+ f"token_width {token_width} * merge_kernel_size {merge_kernel_size} > patch_limit_on_one_side {patch_limit_on_one_side}"
147
+ )
148
+
149
+ num_tokens = token_height * token_width
150
+ return {
151
+ "num_tokens": num_tokens,
152
+ "new_width": new_w,
153
+ "new_height": new_h,
154
+ "pad_width": pad_width,
155
+ "pad_height": pad_height,
156
+ "sampled_nframes": 1,
157
+ }
158
+
159
+
160
+ def _to_pil(
161
+ data: str | bytes | Image.Image,
162
+ transparent_bg_config: TransparentBgConfig | None = None,
163
+ to_rgb: bool = True,
164
+ ) -> Image.Image:
165
+ """Load an image and (optionally) composite its transparent background.
166
+
167
+ Args:
168
+ data: A PIL Image, a base64 ``data:`` URL, a file path, or raw bytes.
169
+ transparent_bg_config: The config used to fill the transparent
170
+ background. ``None`` keeps the historical behavior of converting
171
+ to RGB without compositing.
172
+ to_rgb: If ``False`` the image is returned as-is (the
173
+ ``transparent_bg_config`` is ignored). The caller is then
174
+ expected to call :func:`fill_transparent_bg_with` later — e.g.
175
+ after a resize.
176
+ """
177
+ if isinstance(data, Image.Image):
178
+ image = data
179
+ elif isinstance(data, str):
180
+ if data.startswith("data:"):
181
+ raw_base64 = data.split(",")[1]
182
+ image = Image.open(io.BytesIO(base64.b64decode(raw_base64)))
183
+ else:
184
+ image = Image.open(data)
185
+ elif isinstance(data, bytes):
186
+ image = Image.open(io.BytesIO(data))
187
+ else:
188
+ raise ValueError(f"Unsupported data type: {type(data)}")
189
+
190
+ if not to_rgb:
191
+ return image
192
+
193
+ return fill_transparent_bg_with(image, transparent_bg_config)
194
+
195
+
196
+ def ensure_media_type(
197
+ media: MediaInput,
198
+ transparent_bg_config: TransparentBgConfig | None = None,
199
+ transparent_bg_fill_stage: Literal["before_resize",
200
+ "after_resize"] = "before_resize",
201
+ ) -> MediaInput:
202
+ if media['type'] == 'image':
203
+ media['image'] = _to_pil(
204
+ media['image'],
205
+ transparent_bg_config=transparent_bg_config,
206
+ to_rgb=transparent_bg_fill_stage == "before_resize",
207
+ )
208
+ return media
209
+ else:
210
+ raise ValueError(f"Unsupported media type: {media['type']}")
211
+
212
+
213
+ def image_to_np(
214
+ image: Image.Image,
215
+ resize_to: tuple[int, int] | None = None,
216
+ mode: str = "resize",
217
+ raise_error_for_ill_resize: bool = True,
218
+ transparent_bg_config: TransparentBgConfig | None = None,
219
+ transparent_bg_fill_stage: Literal["before_resize",
220
+ "after_resize"] = "before_resize",
221
+ ) -> np.ndarray:
222
+ """Convert an image to a numpy array.
223
+
224
+ Args:
225
+ content: The image to convert.
226
+ resize_to: The size to resize the image to.
227
+ mode: The mode to resize the image to.
228
+ raise_error_for_ill_resize: Whether to raise an error for ill-sized resize.
229
+ transparent_bg_config: The config of the transparent background. Only
230
+ used when ``transparent_bg_fill_stage == "after_resize"`` (the
231
+ caller is responsible for filling before resize otherwise).
232
+ transparent_bg_fill_stage: When to composite the transparent
233
+ background — before or after the resize step.
234
+
235
+ Returns:
236
+ A numpy array.
237
+ """
238
+ assert isinstance(image, Image.Image), "image must be a PIL Image"
239
+ if resize_to is not None:
240
+ if mode == "resize":
241
+ image = image.resize(resize_to, resample=Image.Resampling.BICUBIC)
242
+ if transparent_bg_fill_stage == "after_resize":
243
+ image = fill_transparent_bg_with(image, transparent_bg_config)
244
+
245
+ elif mode == "rescale_and_pad_to_center":
246
+ scale = min(resize_to[0] / image.width,
247
+ resize_to[1] / image.height, 1.0)
248
+ new_width = round(image.width * scale)
249
+ new_height = round(image.height * scale)
250
+ if new_width == 0 or new_height == 0:
251
+ if raise_error_for_ill_resize:
252
+ raise ValueError(
253
+ f"Invalid resize to: {resize_to}, from image size: {image.size}"
254
+ )
255
+ else:
256
+ return np.zeros((resize_to[1], resize_to[0], 3),
257
+ dtype=np.uint8)
258
+
259
+ image = image.resize((new_width, new_height),
260
+ resample=Image.Resampling.BICUBIC)
261
+ if transparent_bg_fill_stage == "after_resize":
262
+ image = fill_transparent_bg_with(image, transparent_bg_config)
263
+ padding_left = (resize_to[0] - new_width) // 2
264
+ padding_right = resize_to[0] - new_width - padding_left
265
+ padding_top = (resize_to[1] - new_height) // 2
266
+ padding_bottom = resize_to[1] - new_height - padding_top
267
+ image = np.asarray(image)
268
+ image = np.pad(
269
+ image,
270
+ ((padding_top, padding_bottom), (padding_left, padding_right),
271
+ (0, 0)),
272
+ mode="constant",
273
+ constant_values=0,
274
+ )
275
+ assert image.shape == (resize_to[1], resize_to[0], 3)
276
+
277
+ elif mode == "rescale_and_pad_to_rightbottom":
278
+ scale = min(resize_to[0] / image.width,
279
+ resize_to[1] / image.height, 1.0)
280
+ new_width = round(image.width * scale)
281
+ new_height = round(image.height * scale)
282
+ if new_width == 0 or new_height == 0:
283
+ if raise_error_for_ill_resize:
284
+ raise ValueError(
285
+ f"Invalid resize to: {resize_to}, from image size: {image.size}"
286
+ )
287
+ else:
288
+ return np.zeros((resize_to[1], resize_to[0], 3),
289
+ dtype=np.uint8)
290
+
291
+ image = image.resize((new_width, new_height),
292
+ resample=Image.Resampling.BICUBIC)
293
+ if transparent_bg_fill_stage == "after_resize":
294
+ image = fill_transparent_bg_with(image, transparent_bg_config)
295
+ padding_right = resize_to[0] - new_width
296
+ padding_bottom = resize_to[1] - new_height
297
+ image = np.asarray(image)
298
+ image = np.pad(
299
+ image,
300
+ ((0, padding_bottom), (0, padding_right), (0, 0)),
301
+ mode="constant",
302
+ constant_values=0,
303
+ )
304
+ assert image.shape == (resize_to[1], resize_to[0], 3)
305
+
306
+ else:
307
+ raise ValueError(f"Invalid mode: {mode}")
308
+
309
+ if isinstance(image, Image.Image):
310
+ return np.asarray(image)
311
+ else:
312
+ return image
313
+
314
+
315
+ def navit_patchify(pixel_values: np.ndarray,
316
+ patch_size: int) -> dict[str, np.ndarray]:
317
+ """Reshape the pixel values to a navit shape.
318
+
319
+ Args:
320
+ pixel_values: np.ndarray, shape (t, h, w, c)
321
+ patch_size: int
322
+
323
+ Returns:
324
+ dict[str, np.ndarray]
325
+ - patches: np.ndarray, shape (t * h//patch_size * w//patch_size, c, patch_size, patch_size)
326
+ - grid_thw: np.ndarray, (t, h//patch_size, w//patch_size)
327
+ """
328
+ T, H, W, C = pixel_values.shape
329
+ assert C == 3, "pixel_values must have 3 channels"
330
+
331
+ patches = pixel_values.reshape(T, H // patch_size, patch_size,
332
+ W // patch_size, patch_size, C)
333
+ # (T, H//patch_size, W//patch_size, C, patch_size, patch_size)
334
+ patches = patches.transpose(0, 1, 3, 5, 2, 4)
335
+ patches = patches.reshape(-1, C, patch_size, patch_size)
336
+ grid_thw = np.array([T, H // patch_size, W // patch_size])
337
+ return {"pixel_values": patches, "grid_thw": grid_thw}
338
+
339
+
340
+ def normalize(x: np.ndarray,
341
+ mean,
342
+ std_inv,
343
+ pixels_dtype: np.dtype = np.float32) -> np.ndarray:
344
+ """Normalize the image.
345
+
346
+ Args:
347
+ x: The image to normalize. The shape is (..., 3). The dtype is uint8. The range is [0, 255].
348
+ mean: The mean of the image.
349
+ std_inv: The inverse of the std of the image.
350
+ pixels_dtype: The dtype of the image.
351
+ Returns:
352
+ The normalized image. The shape is (..., 3). The dtype is determined by the pixels_dtype.
353
+ """
354
+ x = (x / 255.0).astype(pixels_dtype)
355
+ x -= mean
356
+ x *= std_inv
357
+ return x
358
+
359
+
360
+ def _to_tensor(data, **kwargs):
361
+ import torch
362
+
363
+ if isinstance(data, np.ndarray):
364
+ return torch.from_numpy(data).to(**kwargs)
365
+ elif isinstance(data, torch.Tensor):
366
+ return data.to(**kwargs)
367
+ elif isinstance(data, list):
368
+ return [_to_tensor(item, **kwargs) for item in data]
369
+ elif isinstance(data, tuple):
370
+ return tuple(_to_tensor(item, **kwargs) for item in data)
371
+ elif isinstance(data, dict):
372
+ return {k: _to_tensor(v, **kwargs) for k, v in data.items()}
373
+ elif data is None:
374
+ return None
375
+ else:
376
+ raise ValueError(f"Unsupported data type: {type(data)}")
merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3cf09355cde8cf4877cdc76e0a72301d572c1d6263b2831f8a03368285bb65e2
3
+ size 9687419736
modeling_groundinganything.py ADDED
@@ -0,0 +1,1799 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from collections.abc import Callable
3
+ from dataclasses import dataclass
4
+ from typing import Any, Optional, Union
5
+ from pathlib import Path
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ from torch.nn import LayerNorm
10
+
11
+ from transformers import AutoModel
12
+ from transformers.cache_utils import Cache
13
+ from transformers.generation import GenerationMixin
14
+ from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling, ModelOutput
15
+ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
16
+ from transformers.models.siglip.modeling_siglip import SiglipMLP
17
+ from transformers.processing_utils import Unpack
18
+ from transformers.utils import (
19
+ TransformersKwargs,
20
+ auto_docstring,
21
+ can_return_tuple,
22
+ replace_return_docstrings,
23
+ )
24
+ try:
25
+ from transformers.utils.generic import is_flash_attention_requested
26
+ except ImportError:
27
+ def is_flash_attention_requested(config):
28
+ return getattr(config, "_attn_implementation", None) == "flash_attention_2"
29
+
30
+ from .configuration_groundinganything import GroundAnythingVLMConfig, GroundAnythingVLMVisionConfig, GroundAnythingConfig
31
+ from .modeling_groundinganything_vision import MoonViT3dPretrainedModel
32
+
33
+
34
+ @dataclass
35
+ @auto_docstring(
36
+ custom_intro="""
37
+ Base class for GroundAnything-VLM outputs, with hidden states and attentions.
38
+ """
39
+ )
40
+ class GroundAnythingVLMBaseModelOutputWithPast(ModelOutput):
41
+ r"""
42
+ past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
43
+ It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).
44
+
45
+ Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
46
+ `past_key_values` input) to speed up sequential decoding.
47
+ """
48
+
49
+ last_hidden_state: Optional[torch.FloatTensor] = None
50
+ past_key_values: Optional[Cache] = None
51
+ hidden_states: Optional[tuple[torch.FloatTensor]] = None
52
+ attentions: Optional[tuple[torch.FloatTensor]] = None
53
+
54
+
55
+ @dataclass
56
+ @auto_docstring(
57
+ custom_intro="""
58
+ Base class for GroundAnything-VLM causal language model (or autoregressive) outputs.
59
+ """
60
+ )
61
+ class GroundAnythingVLMCausalLMOutputWithPast(ModelOutput):
62
+ r"""
63
+ loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
64
+ Language modeling loss (for next-token prediction).
65
+ logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
66
+ Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
67
+ past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
68
+ It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).
69
+
70
+ Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
71
+ `past_key_values` input) to speed up sequential decoding.
72
+ """
73
+
74
+ loss: Optional[torch.FloatTensor] = None
75
+ logits: Optional[torch.FloatTensor] = None
76
+ past_key_values: Optional[Cache] = None
77
+ hidden_states: Optional[tuple[torch.FloatTensor]] = None
78
+ attentions: Optional[tuple[torch.FloatTensor]] = None
79
+
80
+
81
+ # ---------------------------------------------------------------------------
82
+ # Vision Rotary Embedding
83
+ # ---------------------------------------------------------------------------
84
+
85
+
86
+ class VisionRotaryEmbedding(nn.Module):
87
+ """
88
+ 3D (T,H,W) Rotary frequency constructor with 4:6:6 split.
89
+ Supports both grid_thw-based and explicit position-based RoPE computation.
90
+ """
91
+
92
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
93
+ super().__init__()
94
+ head_dim = config.hidden_size // config.num_attention_heads
95
+ base = config.rope_theta
96
+
97
+ assert head_dim % 2 == 0, "head_dim must be even for rotary."
98
+ assert head_dim % 16 == 0, "head_dim must be divisible by 16."
99
+ half = head_dim // 2
100
+ assert half % 16 == 0, "head_dim//2 must also be divisible by 16 to split into 4:6:6."
101
+
102
+ self.head_dim = head_dim
103
+ self.half = half
104
+ self.base = base
105
+
106
+ # 4:6:6 split for T:H:W
107
+ unit = half // 16
108
+ self.t_size = 4 * unit
109
+ self.h_size = 6 * unit
110
+ self.w_size = 6 * unit
111
+
112
+ self.register_buffer(
113
+ "inv_freq_t",
114
+ 1.0 / (base ** (torch.arange(self.t_size, dtype=torch.float32) / self.t_size)),
115
+ persistent=False,
116
+ )
117
+ self.register_buffer(
118
+ "inv_freq_h",
119
+ 1.0 / (base ** (torch.arange(self.h_size, dtype=torch.float32) / self.h_size)),
120
+ persistent=False,
121
+ )
122
+ self.register_buffer(
123
+ "inv_freq_w",
124
+ 1.0 / (base ** (torch.arange(self.w_size, dtype=torch.float32) / self.w_size)),
125
+ persistent=False,
126
+ )
127
+
128
+ def forward(self, grid_thw: torch.Tensor) -> torch.Tensor:
129
+ """
130
+ Compute rotary position embeddings from grid_thw (Qwen2VL style).
131
+
132
+ Args:
133
+ grid_thw: [num_samples, 3] tensor with [t, h, w] for each sample
134
+
135
+ Returns:
136
+ freqs: [total_seq_len, half] tensor of position frequencies
137
+ """
138
+ device = grid_thw.device
139
+ inv_t = self.inv_freq_t.to(device=device)
140
+ inv_h = self.inv_freq_h.to(device=device)
141
+ inv_w = self.inv_freq_w.to(device=device)
142
+
143
+ all_freqs = []
144
+ for sample_thw in grid_thw:
145
+ t, h, w = sample_thw[0].item(), sample_thw[1].item(), sample_thw[2].item()
146
+
147
+ # Compute frequency tables
148
+ ft = torch.outer(torch.arange(t, device=device, dtype=torch.float32), inv_t)
149
+ fh = torch.outer(torch.arange(h, device=device, dtype=torch.float32), inv_h)
150
+ fw = torch.outer(torch.arange(w, device=device, dtype=torch.float32), inv_w)
151
+
152
+ # Build position indices for this sample
153
+ t_ids = torch.arange(t, device=device).repeat_interleave(h * w)
154
+ h_ids = torch.arange(h, device=device).repeat_interleave(w).repeat(t)
155
+ w_ids = torch.arange(w, device=device).repeat(h).repeat(t)
156
+
157
+ # Concatenate frequencies: [seq_len, half]
158
+ sample_freqs = torch.cat([ft[t_ids], fh[h_ids], fw[w_ids]], dim=-1)
159
+ all_freqs.append(sample_freqs)
160
+
161
+ return torch.cat(all_freqs, dim=0)
162
+
163
+ def forward_from_positions(self, patch_positions: torch.Tensor) -> torch.Tensor:
164
+ """
165
+ Compute rotary position embeddings from explicit patch positions.
166
+
167
+ Args:
168
+ patch_positions: [seq_len, 3] tensor with [t, h, w] positions for each patch
169
+
170
+ Returns:
171
+ freqs: [seq_len, half] tensor of position frequencies
172
+ """
173
+ device = patch_positions.device
174
+ inv_t = self.inv_freq_t.to(device=device)
175
+ inv_h = self.inv_freq_h.to(device=device)
176
+ inv_w = self.inv_freq_w.to(device=device)
177
+
178
+ t_pos = patch_positions[:, 0].float()
179
+ h_pos = patch_positions[:, 1].float()
180
+ w_pos = patch_positions[:, 2].float()
181
+
182
+ ft = torch.outer(t_pos, inv_t)
183
+ fh = torch.outer(h_pos, inv_h)
184
+ fw = torch.outer(w_pos, inv_w)
185
+
186
+ return torch.cat([ft, fh, fw], dim=-1)
187
+
188
+ def forward_with_thw(self, t: int, h: int, w: int, device=None) -> torch.Tensor:
189
+ """
190
+ Compute rotary position embeddings from explicit t, h, w dimensions.
191
+
192
+ Args:
193
+ t: Number of temporal frames
194
+ h: Number of height patches
195
+ w: Number of width patches
196
+ device: Target device
197
+
198
+ Returns:
199
+ freqs: [t*h*w, half] tensor of position frequencies
200
+ """
201
+ if device is None:
202
+ device = self.inv_freq_t.device
203
+
204
+ inv_t = self.inv_freq_t.to(device=device)
205
+ inv_h = self.inv_freq_h.to(device=device)
206
+ inv_w = self.inv_freq_w.to(device=device)
207
+
208
+ ft = torch.outer(torch.arange(t, device=device, dtype=torch.float32), inv_t)
209
+ fh = torch.outer(torch.arange(h, device=device, dtype=torch.float32), inv_h)
210
+ fw = torch.outer(torch.arange(w, device=device, dtype=torch.float32), inv_w)
211
+
212
+ t_ids = torch.arange(t, device=device).repeat_interleave(h * w)
213
+ h_ids = torch.arange(h, device=device).repeat_interleave(w).repeat(t)
214
+ w_ids = torch.arange(w, device=device).repeat(h).repeat(t)
215
+
216
+ freqs = torch.cat([ft[t_ids], fh[h_ids], fw[w_ids]], dim=-1)
217
+ return freqs
218
+
219
+
220
+ # ---------------------------------------------------------------------------
221
+ # Patch Embedding
222
+ # ---------------------------------------------------------------------------
223
+
224
+
225
+ class GroundAnythingVLMVisionEmbeddings(nn.Module):
226
+ """
227
+ Patch embedding layer that converts pre-processed patches to embeddings.
228
+
229
+ This module is designed to receive patches that have already been extracted
230
+ and arranged by the Qwen2VL image processor in 2x2 block spatial order.
231
+
232
+ Input format: [total_patches, num_channels, patch_size, patch_size]
233
+ Output format: [total_patches, embed_dim]
234
+ """
235
+
236
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
237
+ super().__init__()
238
+ self.config = config
239
+ self.embed_dim = config.hidden_size
240
+ self.image_size = config.image_size
241
+ self.patch_size = config.patch_size
242
+ self.in_channels = config.num_channels
243
+
244
+ self.patch_embedding = nn.Conv2d(
245
+ in_channels=config.num_channels,
246
+ out_channels=self.embed_dim,
247
+ kernel_size=self.patch_size,
248
+ stride=self.patch_size,
249
+ bias=False,
250
+ )
251
+
252
+ def forward(self, hidden_states: torch.FloatTensor) -> torch.Tensor:
253
+ target_dtype = self.patch_embedding.weight.dtype
254
+ hidden_states = hidden_states.view(-1, self.in_channels, self.patch_size, self.patch_size)
255
+ hidden_states = self.patch_embedding(hidden_states.to(dtype=target_dtype)).view(-1, self.embed_dim)
256
+
257
+ return hidden_states
258
+
259
+
260
+ # ---------------------------------------------------------------------------
261
+ # Patch Merger
262
+ # ---------------------------------------------------------------------------
263
+
264
+
265
+ class GroundAnythingVLMVisionPatchMerger(nn.Module):
266
+ """
267
+ Patch merger that merges spatial_merge_size x spatial_merge_size patches into one.
268
+
269
+ This module is designed to work with Qwen2VL-style patch processing where patches
270
+ are already arranged in 2x2 block order by the image processor.
271
+ """
272
+
273
+ def __init__(
274
+ self,
275
+ dim: int,
276
+ context_dim: int,
277
+ spatial_merge_size: int = 2,
278
+ layer_norm_eps: float = 1e-05,
279
+ use_patch_position_encoding: bool = False,
280
+ patch_position_encoding_type: str = "absolute",
281
+ max_position_embeddings: int = 8192,
282
+ ) -> None:
283
+ super().__init__()
284
+ self.hidden_size = context_dim * (spatial_merge_size**2)
285
+ self.ln_q = LayerNorm(context_dim, eps=layer_norm_eps)
286
+ self.mlp = nn.Sequential(
287
+ nn.Linear(self.hidden_size, self.hidden_size),
288
+ nn.GELU(),
289
+ nn.Linear(self.hidden_size, dim),
290
+ )
291
+ self.spatial_merge_size = spatial_merge_size
292
+ self.use_patch_position_encoding = use_patch_position_encoding
293
+ self.patch_position_encoding_type = patch_position_encoding_type
294
+
295
+ if self.use_patch_position_encoding:
296
+ if self.patch_position_encoding_type != "absolute":
297
+ raise ValueError(
298
+ f"Unknown patch_position_encoding_type: {self.patch_position_encoding_type}. "
299
+ "Only 'absolute' is supported."
300
+ )
301
+ self.pos_emb_h = nn.Embedding(max_position_embeddings, dim)
302
+ self.pos_emb_w = nn.Embedding(max_position_embeddings, dim)
303
+
304
+ def forward(self, x: torch.Tensor, patch_positions: Optional[torch.Tensor] = None) -> torch.Tensor:
305
+ """
306
+ Merge patches from Qwen2VL-style input.
307
+
308
+ The input patches are already arranged in 2x2 block order by the image processor,
309
+ so we simply need to apply LayerNorm, reshape, and project through MLP.
310
+
311
+ Args:
312
+ x: Input tensor of shape [batch_size, seq_len, hidden_size] or [seq_len, hidden_size]
313
+ where seq_len = t * h * w (patches in 2x2 block order)
314
+
315
+ Returns:
316
+ Merged tensor of shape [batch_size, seq_len // spatial_merge_size^2, dim]
317
+ or [seq_len // spatial_merge_size^2, dim]
318
+ """
319
+ if patch_positions is not None and patch_positions.dim() == 3:
320
+ patch_positions = patch_positions.squeeze(0)
321
+
322
+ x = self.ln_q(x).view(-1, self.hidden_size)
323
+ x = self.mlp(x)
324
+
325
+ if self.use_patch_position_encoding and patch_positions is not None:
326
+ pp = patch_positions.view(-1, self.spatial_merge_size**2, 3)
327
+ pp = pp[:, 0, :]
328
+ pp = (pp // self.spatial_merge_size).long()
329
+
330
+ x = x + self.pos_emb_h(pp[:, 1]) + self.pos_emb_w(pp[:, 2])
331
+
332
+ return x
333
+
334
+
335
+ def rotate_half(x):
336
+ """
337
+ Interleaved rotation to match Source model's implementation.
338
+ (x1, x2, x3, x4) -> (-x2, x1, -x4, x3)
339
+ """
340
+ x_even = x[..., ::2]
341
+ x_odd = x[..., 1::2]
342
+ return torch.stack((-x_odd, x_even), dim=-1).flatten(-2)
343
+
344
+
345
+ def get_norm_layer(config):
346
+ if config.layer_norm_type == "rms_norm":
347
+ return nn.RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
348
+ else:
349
+ return nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
350
+
351
+
352
+ def apply_rotary_pos_emb(q, k, freqs):
353
+ # q, k: (B, H, L, D)
354
+ # freqs: (B, L, D)
355
+ orig_q_dtype = q.dtype
356
+ orig_k_dtype = k.dtype
357
+ q, k = q.float(), k.float()
358
+ # We need to broadcast freqs to match heads
359
+ # (B, L, D) -> (B, 1, L, D)
360
+ # Keep the same dtype as q, k to avoid memory doubling from float32 promotion
361
+ cos = freqs.cos().unsqueeze(1).float()
362
+ sin = freqs.sin().unsqueeze(1).float()
363
+
364
+ q_embed = (q * cos) + (rotate_half(q) * sin)
365
+ k_embed = (k * cos) + (rotate_half(k) * sin)
366
+ q_embed = q_embed.to(orig_q_dtype)
367
+ k_embed = k_embed.to(orig_k_dtype)
368
+ return q_embed, k_embed
369
+
370
+
371
+ def eager_attention_forward(
372
+ module: nn.Module,
373
+ query: torch.Tensor,
374
+ key: torch.Tensor,
375
+ value: torch.Tensor,
376
+ attention_mask: Optional[torch.Tensor],
377
+ scaling: float,
378
+ dropout: float = 0.0,
379
+ **kwargs,
380
+ ):
381
+ """Eager attention; query/key/value are expected as ``(B, H, L, D)``."""
382
+ attn_weights = torch.matmul(query, key.transpose(2, 3)) * scaling
383
+ if attention_mask is not None:
384
+ attn_weights = attn_weights + attention_mask
385
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
386
+ attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)
387
+ attn_output = torch.matmul(attn_weights, value)
388
+ attn_output = attn_output.transpose(1, 2).contiguous() # (B, L, H, D)
389
+ return attn_output, attn_weights
390
+
391
+
392
+ class GroundAnythingVLMVisionAttention(nn.Module):
393
+ """
394
+ Multi-headed attention with RoPE support, dispatched through
395
+ :data:`ALL_ATTENTION_FUNCTIONS` (``eager`` / ``sdpa`` / ``flash_attention_2``)
396
+ based on ``config._attn_implementation``.
397
+ """
398
+
399
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
400
+ super().__init__()
401
+ self.config = config
402
+ self.embed_dim = config.hidden_size
403
+ self.num_heads = config.num_attention_heads
404
+ self.head_dim = self.embed_dim // self.num_heads
405
+ if self.head_dim * self.num_heads != self.embed_dim:
406
+ raise ValueError(
407
+ f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`: {self.num_heads})."
408
+ )
409
+
410
+ self.num_key_value_groups = 1 # required by repeat_kv-aware eager paths
411
+ self.scale = self.head_dim**-0.5
412
+ self.scaling = self.scale # alias expected by some attention interfaces
413
+ self.attention_dropout = config.attention_dropout
414
+ self.is_causal = False
415
+ self.qkv = nn.Linear(self.embed_dim, self.embed_dim * 3)
416
+ self.proj = nn.Linear(self.embed_dim, self.embed_dim)
417
+
418
+ def forward(
419
+ self,
420
+ hidden_states: torch.Tensor,
421
+ attention_mask: Optional[torch.Tensor] = None,
422
+ rotary_pos_emb: Optional[torch.Tensor] = None,
423
+ output_attentions: bool = False,
424
+ cu_seqlens: Optional[torch.Tensor] = None,
425
+ max_seqlen: Optional[int] = None,
426
+ **kwargs,
427
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
428
+ batch_size, q_len, _ = hidden_states.size()
429
+ # (B, L, 3*H*D) -> (B, L, 3, H, D) -> 3 x (B, L, H, D) -> 3 x (B, H, L, D)
430
+ q, k, v = (
431
+ self.qkv(hidden_states)
432
+ .reshape(batch_size, q_len, 3, self.num_heads, self.head_dim)
433
+ .permute(2, 0, 1, 3, 4)
434
+ .unbind(0)
435
+ )
436
+ query_states = q.transpose(1, 2)
437
+ key_states = k.transpose(1, 2)
438
+ value_states = v.transpose(1, 2)
439
+
440
+ if rotary_pos_emb is not None:
441
+ query_states, key_states = apply_rotary_pos_emb(query_states, key_states, rotary_pos_emb)
442
+
443
+ attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
444
+ self.config._attn_implementation, eager_attention_forward
445
+ )
446
+ dropout = 0.0 if not self.training else self.attention_dropout
447
+
448
+ if cu_seqlens is not None and is_flash_attention_requested(self.config):
449
+ # Flash Attention varlen path: pass cu_seq_lens / max_length kwargs.
450
+ if max_seqlen is None:
451
+ max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max()
452
+ attn_output, _ = attention_interface(
453
+ self,
454
+ query_states,
455
+ key_states,
456
+ value_states,
457
+ attention_mask=None,
458
+ scaling=self.scale,
459
+ dropout=dropout,
460
+ cu_seq_lens_q=cu_seqlens,
461
+ cu_seq_lens_k=cu_seqlens,
462
+ max_length_q=max_seqlen,
463
+ max_length_k=max_seqlen,
464
+ is_causal=False,
465
+ **kwargs,
466
+ )
467
+ elif cu_seqlens is not None:
468
+ # Non-FA implementations do not understand cu_seqlens directly; mirror
469
+ # Qwen3-VL by splitting the packed sequence into per-sample chunks
470
+ # along the L dim of (B, H, L, D) and running attention per chunk.
471
+ lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).tolist()
472
+ splits = [torch.split(t, lengths, dim=2) for t in (query_states, key_states, value_states)]
473
+ attn_outputs = [
474
+ attention_interface(
475
+ self,
476
+ q_chunk,
477
+ k_chunk,
478
+ v_chunk,
479
+ attention_mask=None,
480
+ scaling=self.scale,
481
+ dropout=dropout,
482
+ is_causal=False,
483
+ **kwargs,
484
+ )[0]
485
+ for q_chunk, k_chunk, v_chunk in zip(*splits)
486
+ ]
487
+ # interface output is (B, l_i, H, D); concat along the L axis
488
+ attn_output = torch.cat(attn_outputs, dim=1)
489
+ else:
490
+ attn_mask = None
491
+ if attention_mask is not None:
492
+ attn_mask = attention_mask
493
+ if attn_mask.dim() == 2:
494
+ attn_mask = attn_mask.unsqueeze(0)
495
+ if attn_mask.shape[0] == 1 and batch_size > 1:
496
+ attn_mask = attn_mask.expand(batch_size, -1, -1)
497
+ attn_mask = attn_mask.unsqueeze(1) # (B, 1, L, L)
498
+ attn_output, _ = attention_interface(
499
+ self,
500
+ query_states,
501
+ key_states,
502
+ value_states,
503
+ attention_mask=attn_mask,
504
+ scaling=self.scale,
505
+ dropout=dropout,
506
+ is_causal=False,
507
+ **kwargs,
508
+ )
509
+
510
+ attn_output = attn_output.reshape(batch_size, q_len, self.embed_dim)
511
+ attn_output = self.proj(attn_output)
512
+
513
+ return attn_output, None
514
+
515
+
516
+ class GroundAnythingVLMVisionEncoderLayer(nn.Module):
517
+ """Vision encoder layer with pre-norm and Flash Attention 2."""
518
+
519
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
520
+ super().__init__()
521
+ self.embed_dim = config.hidden_size
522
+ self.self_attn = GroundAnythingVLMVisionAttention(config)
523
+ self.layer_norm1 = get_norm_layer(config)
524
+ self.mlp = SiglipMLP(config)
525
+ self.layer_norm2 = get_norm_layer(config)
526
+
527
+ def forward(
528
+ self,
529
+ hidden_states: torch.Tensor,
530
+ attention_mask: Optional[torch.Tensor] = None,
531
+ rotary_pos_emb: Optional[torch.Tensor] = None,
532
+ output_attentions: bool = False,
533
+ cu_seqlens: Optional[torch.Tensor] = None,
534
+ max_seqlen: Optional[int] = None,
535
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
536
+ residual = hidden_states
537
+ hidden_states = self.layer_norm1(hidden_states)
538
+
539
+ hidden_states, attn_weights = self.self_attn(
540
+ hidden_states=hidden_states,
541
+ attention_mask=attention_mask,
542
+ rotary_pos_emb=rotary_pos_emb,
543
+ output_attentions=output_attentions,
544
+ cu_seqlens=cu_seqlens,
545
+ max_seqlen=max_seqlen,
546
+ )
547
+ hidden_states = residual + hidden_states
548
+
549
+ residual = hidden_states
550
+ hidden_states = self.layer_norm2(hidden_states)
551
+ hidden_states = self.mlp(hidden_states)
552
+ hidden_states = residual + hidden_states
553
+
554
+ outputs = (hidden_states, attn_weights) if output_attentions else (hidden_states,)
555
+ return outputs
556
+
557
+
558
+ class GroundAnythingVLMVisionEncoder(nn.Module):
559
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
560
+ super().__init__()
561
+ self.config = config
562
+ self.layers = nn.ModuleList([GroundAnythingVLMVisionEncoderLayer(config) for _ in range(config.num_hidden_layers)])
563
+ # Gradient checkpointing support
564
+ self.gradient_checkpointing = False
565
+
566
+ def forward(
567
+ self,
568
+ hidden_states: torch.Tensor,
569
+ attention_mask: Optional[torch.Tensor] = None,
570
+ rotary_pos_emb: Optional[torch.Tensor] = None,
571
+ output_attentions: bool = False,
572
+ output_hidden_states: bool = False,
573
+ return_dict: bool = True,
574
+ cu_seqlens: Optional[torch.Tensor] = None,
575
+ max_seqlen: Optional[int] = None,
576
+ ) -> Union[tuple, BaseModelOutput]:
577
+ all_hidden_states = () if output_hidden_states else None
578
+ all_self_attentions = () if output_attentions else None
579
+
580
+ for layer in self.layers:
581
+ if output_hidden_states:
582
+ all_hidden_states = all_hidden_states + (hidden_states,)
583
+
584
+ if self.gradient_checkpointing and self.training:
585
+ layer_outputs = self._gradient_checkpointing_func(
586
+ layer.__call__,
587
+ hidden_states,
588
+ attention_mask,
589
+ rotary_pos_emb,
590
+ output_attentions,
591
+ cu_seqlens,
592
+ max_seqlen,
593
+ )
594
+ else:
595
+ layer_outputs = layer(
596
+ hidden_states,
597
+ attention_mask=attention_mask,
598
+ rotary_pos_emb=rotary_pos_emb,
599
+ output_attentions=output_attentions,
600
+ cu_seqlens=cu_seqlens,
601
+ max_seqlen=max_seqlen,
602
+ )
603
+
604
+ hidden_states = layer_outputs[0]
605
+
606
+ if output_attentions:
607
+ all_self_attentions = all_self_attentions + (layer_outputs[1],)
608
+
609
+ if output_hidden_states:
610
+ all_hidden_states = all_hidden_states + (hidden_states,)
611
+
612
+ if not return_dict:
613
+ return tuple(v for v in [hidden_states, all_hidden_states, all_self_attentions] if v is not None)
614
+
615
+ return BaseModelOutput(
616
+ last_hidden_state=hidden_states,
617
+ hidden_states=all_hidden_states,
618
+ attentions=all_self_attentions,
619
+ )
620
+
621
+
622
+ class GroundAnythingVLMPreTrainedModel(PreTrainedModel):
623
+ _supports_attention_backend = True
624
+ config_class = GroundAnythingVLMConfig
625
+ base_model_prefix = "model"
626
+ input_modalities = ("image", "video", "text")
627
+ supports_gradient_checkpointing = True
628
+ _no_split_modules = ["MoonViTEncoderLayer", "Qwen3DecoderLayer"]
629
+ _skip_keys_device_placement = "past_key_values"
630
+ _supports_flash_attn = True
631
+ _supports_sdpa = True
632
+
633
+ def _init_weights(self, module):
634
+ super()._init_weights(module)
635
+ # Re-initialize VisionRotaryEmbedding inv_freq buffers.
636
+ # These are registered with persistent=False, so they are not in the checkpoint
637
+ # state_dict. When ``from_pretrained`` materializes the model from meta tensors,
638
+ # the values in these buffers end up uninitialized. Mirror Qwen3-VL by explicitly
639
+ # filling them here so RoPE produces the correct frequencies post-load.
640
+ if isinstance(module, VisionRotaryEmbedding):
641
+ base = module.base
642
+ with torch.no_grad():
643
+ inv_t = 1.0 / (base ** (torch.arange(module.t_size, dtype=torch.float32) / module.t_size))
644
+ inv_h = 1.0 / (base ** (torch.arange(module.h_size, dtype=torch.float32) / module.h_size))
645
+ inv_w = 1.0 / (base ** (torch.arange(module.w_size, dtype=torch.float32) / module.w_size))
646
+ module.inv_freq_t.copy_(inv_t.to(module.inv_freq_t.device))
647
+ module.inv_freq_h.copy_(inv_h.to(module.inv_freq_h.device))
648
+ module.inv_freq_w.copy_(inv_w.to(module.inv_freq_w.device))
649
+
650
+
651
+ class Siglip2MultiheadAttentionPoolingHead(nn.Module):
652
+ """
653
+ Multi-Head Attention Pooling with a learned probe (PMA-style).
654
+ """
655
+
656
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
657
+ super().__init__()
658
+ self.embed_dim = config.hidden_size
659
+ self.probe = nn.Parameter(torch.randn(1, 1, config.hidden_size))
660
+ self.attention = nn.MultiheadAttention(config.hidden_size, config.num_attention_heads, batch_first=True)
661
+ self.norm = nn.RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
662
+ self.mlp = SiglipMLP(config)
663
+
664
+ def forward(self, hidden_states):
665
+ batch_size = hidden_states.shape[0]
666
+ probe = self.probe.repeat(batch_size, 1, 1)
667
+
668
+ attn_output, _ = self.attention(probe, hidden_states, hidden_states)
669
+
670
+ residual = attn_output
671
+ attn_output = self.norm(attn_output)
672
+ attn_output = residual + self.mlp(attn_output)
673
+
674
+ return attn_output[:, 0]
675
+
676
+
677
+ # ---------------------------------------------------------------------------
678
+ # Vision Model
679
+ # ---------------------------------------------------------------------------
680
+
681
+
682
+ class GroundAnythingVLMVisionPretrainedModel(GroundAnythingVLMPreTrainedModel):
683
+ """
684
+ GroundAnything-VLM Vision Model.
685
+
686
+ This vision model is designed to work with Qwen2VL-style image processing:
687
+ - Receives pre-processed patches in 2x2 block spatial order
688
+ - Applies RoPE with matching 2x2 block layout conversion
689
+ - Accepts explicit patch_positions for RoPE computation
690
+
691
+ Input format:
692
+ hidden_state: [total_patches, num_channels, patch_size, patch_size]
693
+ grid_thw: [num_samples, 3] with [t, h, w] for each sample
694
+ """
695
+
696
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
697
+ super().__init__(config)
698
+ self.config = config
699
+ self.spatial_merge_size = config.spatial_merge_size
700
+
701
+ # Vision components
702
+ self.embeddings = GroundAnythingVLMVisionEmbeddings(config)
703
+ self.layernorm_pre = get_norm_layer(config)
704
+ self.encoder = GroundAnythingVLMVisionEncoder(config)
705
+ self.video_rope = VisionRotaryEmbedding(config)
706
+
707
+ if config.use_head:
708
+ self.layernorm_post = get_norm_layer(config)
709
+ self.head = Siglip2MultiheadAttentionPoolingHead(config)
710
+ else:
711
+ self.layernorm_post = None
712
+ self.head = None
713
+
714
+ self.merger = GroundAnythingVLMVisionPatchMerger(
715
+ dim=config.out_hidden_size,
716
+ context_dim=config.hidden_size,
717
+ spatial_merge_size=config.spatial_merge_size,
718
+ layer_norm_eps=config.layer_norm_eps,
719
+ use_patch_position_encoding=getattr(config, "use_patch_position_encoding", False),
720
+ patch_position_encoding_type=getattr(config, "patch_position_encoding_type", "absolute"),
721
+ max_position_embeddings=getattr(config, "max_position_embeddings", 8192),
722
+ )
723
+
724
+ self.post_init()
725
+
726
+ def _build_cu_seqlens(
727
+ self,
728
+ grid_thw: torch.Tensor,
729
+ total_patches: int,
730
+ fixed_t: Optional[int] = 4,
731
+ device: Optional[torch.device] = None,
732
+ ) -> tuple[torch.Tensor, int]:
733
+ if grid_thw is None or grid_thw.numel() == 0:
734
+ # Fallback for no grid_thw: treat as single sequence
735
+ return torch.tensor([0, total_patches], dtype=torch.int32, device=device), total_patches
736
+
737
+ if device is None:
738
+ device = grid_thw.device
739
+
740
+ cu_seqlens = [0]
741
+ max_seqlen = 0
742
+ total_entries = grid_thw.shape[0]
743
+ current_len = 0
744
+
745
+ # Calculate cumulative lengths: split sequences based on fixed_t if provided
746
+ for idx in range(total_entries):
747
+ t_val = grid_thw[idx, 0].item()
748
+ h_val = grid_thw[idx, 1].item()
749
+ w_val = grid_thw[idx, 2].item()
750
+
751
+ if fixed_t is not None and fixed_t > 0 and t_val > fixed_t:
752
+ # Split large t into chunks of fixed_t
753
+ num_full_windows = t_val // fixed_t
754
+ remainder = t_val % fixed_t
755
+
756
+ # Add full windows
757
+ for _ in range(num_full_windows):
758
+ chunk_patches = fixed_t * int(h_val) * int(w_val)
759
+ current_len += chunk_patches
760
+ max_seqlen = max(max_seqlen, chunk_patches)
761
+ cu_seqlens.append(current_len)
762
+
763
+ # Add remainder if any
764
+ if remainder > 0:
765
+ chunk_patches = remainder * int(h_val) * int(w_val)
766
+ current_len += chunk_patches
767
+ max_seqlen = max(max_seqlen, chunk_patches)
768
+ cu_seqlens.append(current_len)
769
+ else:
770
+ # Standard case: add as one chunk
771
+ chunk_patches = t_val * int(h_val) * int(w_val)
772
+ current_len += chunk_patches
773
+ max_seqlen = max(max_seqlen, chunk_patches)
774
+ cu_seqlens.append(current_len)
775
+
776
+ last_len = cu_seqlens[-1]
777
+ if last_len != total_patches:
778
+ raise ValueError(
779
+ "cu_seqlens calculation mismatch:\n"
780
+ f"- total_patches: {total_patches}\n"
781
+ f"- calculated total: {last_len}\n"
782
+ f"- grid_thw: {grid_thw}"
783
+ )
784
+
785
+ return torch.tensor(cu_seqlens, dtype=torch.int32, device=device), max_seqlen
786
+
787
+ def _build_block_attention_mask(
788
+ self,
789
+ grid_thw: torch.Tensor,
790
+ total_patches: int,
791
+ fixed_t: Optional[int] = 4,
792
+ device: Optional[torch.device] = None,
793
+ ) -> Optional[torch.Tensor]:
794
+ if grid_thw is None or grid_thw.numel() == 0:
795
+ return None
796
+
797
+ if device is None:
798
+ device = grid_thw.device
799
+
800
+ lengths = []
801
+ total_entries = grid_thw.shape[0]
802
+
803
+ for idx in range(total_entries):
804
+ t_val = grid_thw[idx, 0].item()
805
+ h_val = grid_thw[idx, 1].item()
806
+ w_val = grid_thw[idx, 2].item()
807
+
808
+ if fixed_t is not None and fixed_t > 0 and t_val > fixed_t:
809
+ # Split large t into chunks of fixed_t
810
+ num_full_windows = t_val // fixed_t
811
+ remainder = t_val % fixed_t
812
+
813
+ # Add full windows
814
+ for _ in range(num_full_windows):
815
+ lengths.append(fixed_t * int(h_val) * int(w_val))
816
+
817
+ # Add remainder if any
818
+ if remainder > 0:
819
+ lengths.append(remainder * int(h_val) * int(w_val))
820
+ else:
821
+ lengths.append(t_val * int(h_val) * int(w_val))
822
+
823
+ total_len = sum(lengths)
824
+ if total_len != total_patches:
825
+ raise ValueError(
826
+ "Block attention mask length mismatch:\n"
827
+ f"- total_patches: {total_patches}\n"
828
+ f"- total_len: {total_len}\n"
829
+ f"- grid_thw: {grid_thw}"
830
+ )
831
+
832
+ attn_mask = torch.ones((total_len, total_len), dtype=torch.bool, device=device)
833
+ start = 0
834
+ for size in lengths:
835
+ end = start + size
836
+ attn_mask[start:end, start:end] = False
837
+ start = end
838
+
839
+ return attn_mask
840
+
841
+ @replace_return_docstrings(output_type=BaseModelOutputWithPooling, config_class=GroundAnythingVLMVisionConfig)
842
+ def forward(
843
+ self,
844
+ hidden_state: torch.Tensor,
845
+ grid_thw: Optional[torch.Tensor] = None,
846
+ patch_positions: Optional[torch.Tensor] = None,
847
+ output_attentions: Optional[bool] = None,
848
+ output_hidden_states: Optional[bool] = None,
849
+ return_dict: Optional[bool] = None,
850
+ skip_merger: Optional[bool] = False,
851
+ ) -> Union[tuple, BaseModelOutputWithPooling]:
852
+ r"""
853
+ Forward pass for vision model.
854
+
855
+ This method accepts pre-processed patches from Qwen2VL image processor and applies
856
+ RoPE (Rotary Position Embedding) in 2x2 block layout to match the spatial arrangement
857
+ of patches.
858
+
859
+ Args:
860
+ hidden_state: Pre-processed patches from Qwen2VL processor.
861
+ Shape: [total_patches, num_channels, patch_size, patch_size]
862
+ grid_thw: Grid sizes tensor of shape [num_samples, 3] with [t, h, w] for each sample.
863
+ Required for computing RoPE and handling visible indices.
864
+ patch_positions: Optional explicit patch positions for RoPE computation.
865
+ output_attentions: Whether to return attention weights.
866
+ output_hidden_states: Whether to return all hidden states.
867
+ return_dict: Whether to return a ModelOutput instead of tuple.
868
+ skip_merger: If True, skip patch merger (useful for consistency checking).
869
+
870
+ Returns:
871
+ BaseModelOutputWithPooling with last_hidden_state containing merged features.
872
+ """
873
+ output_attentions = (
874
+ output_attentions if output_attentions is not None else getattr(self.config, "output_attentions", False)
875
+ )
876
+ output_hidden_states = (
877
+ output_hidden_states
878
+ if output_hidden_states is not None
879
+ else getattr(self.config, "output_hidden_states", False)
880
+ )
881
+ return_dict = True if return_dict is None else return_dict
882
+
883
+ # 1. Embeddings
884
+ # Note: embeddings returns [total_patches, embed_dim], we need to add batch dimension
885
+ hidden_states = self.embeddings(hidden_state)
886
+ if hidden_states.dim() == 2:
887
+ hidden_states = hidden_states.unsqueeze(0) # [1, total_patches, embed_dim]
888
+ batch_size, total_patches, _ = hidden_states.shape
889
+
890
+ # 2. RoPE Construction
891
+ if patch_positions is not None and patch_positions.dim() == 3:
892
+ patch_positions = patch_positions.squeeze(0)
893
+ freqs_visible = self.video_rope.forward_from_positions(patch_positions)
894
+
895
+ # Concatenate D/2 + D/2 -> D for applying rope
896
+ freqs_visible = torch.cat([freqs_visible, freqs_visible], dim=-1)
897
+ if freqs_visible.dim() == 2:
898
+ freqs_visible = freqs_visible.unsqueeze(0)
899
+
900
+ # 3. Pre-Norm & Encoder
901
+ hidden_states = self.layernorm_pre(hidden_states)
902
+
903
+ cu_seqlens, max_seqlen = self._build_cu_seqlens(
904
+ grid_thw=grid_thw,
905
+ total_patches=total_patches,
906
+ fixed_t=getattr(self.config, "frame_windows_size", 4),
907
+ device=hidden_states.device,
908
+ )
909
+
910
+ encoder_outputs = self.encoder(
911
+ hidden_states,
912
+ attention_mask=None,
913
+ rotary_pos_emb=freqs_visible,
914
+ output_attentions=output_attentions,
915
+ output_hidden_states=True, # Always get hidden states to use -2 layer
916
+ return_dict=True,
917
+ cu_seqlens=cu_seqlens,
918
+ max_seqlen=max_seqlen,
919
+ )
920
+
921
+ # Use second-to-last layer output for better feature representation
922
+ if encoder_outputs.hidden_states is not None and len(encoder_outputs.hidden_states) >= 2 and not skip_merger:
923
+ sequence_output = encoder_outputs.hidden_states[-1]
924
+ else:
925
+ sequence_output = encoder_outputs[0]
926
+
927
+ # Post-Norm
928
+ if self.layernorm_post is not None:
929
+ sequence_output = self.layernorm_post(sequence_output)
930
+
931
+ # Skip merger for consistency check with original ViT
932
+ if skip_merger:
933
+ pooled_output = None
934
+ if self.head is not None:
935
+ pooled_output = self.head(sequence_output)
936
+
937
+ if not return_dict:
938
+ return (sequence_output, pooled_output) + (
939
+ encoder_outputs.hidden_states if output_hidden_states else None,
940
+ )
941
+ return BaseModelOutputWithPooling(
942
+ last_hidden_state=sequence_output,
943
+ pooler_output=pooled_output,
944
+ hidden_states=encoder_outputs.hidden_states if output_hidden_states else None,
945
+ attentions=encoder_outputs.attentions if output_attentions else None,
946
+ )
947
+
948
+ # Patch merger: input patches are already in 2x2 block order from Qwen2VL processor
949
+ merged_output = self.merger(sequence_output, patch_positions=patch_positions)
950
+
951
+ if not return_dict:
952
+ return (merged_output,) + (encoder_outputs.hidden_states if output_hidden_states else None,)
953
+
954
+ return BaseModelOutputWithPooling(
955
+ last_hidden_state=merged_output,
956
+ pooler_output=None,
957
+ hidden_states=encoder_outputs.hidden_states if output_hidden_states else None,
958
+ attentions=encoder_outputs.attentions if output_attentions else None,
959
+ )
960
+
961
+
962
+ class GroundAnythingVLMTwoLayerProjector(nn.Module):
963
+ """K3-native 2x2 patch merger followed by a two-layer MLP."""
964
+
965
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
966
+ super().__init__()
967
+ merge_area = int(config.spatial_merge_size) ** 2
968
+ input_size = int(config.hidden_size) * merge_area
969
+ hidden_size = int(config.projector_hidden_size)
970
+ output_size = int(config.out_hidden_size)
971
+ self.input_size = input_size
972
+ self.pre_norm = nn.LayerNorm(config.hidden_size, eps=config.projector_ln_eps)
973
+ self.fc1 = nn.Linear(input_size, hidden_size, bias=False)
974
+ self.act = nn.GELU()
975
+ self.fc2 = nn.Linear(hidden_size, output_size, bias=False)
976
+ self.post_norm = nn.RMSNorm(output_size, eps=config.projector_ln_eps)
977
+ for layer in (self.fc1, self.fc2):
978
+ nn.init.trunc_normal_(layer.weight, std=math.sqrt(2.0 / layer.in_features))
979
+
980
+ def forward(self, features):
981
+ outputs = []
982
+ for item in features:
983
+ if item.ndim != 3 or item.shape[1] * item.shape[2] != self.input_size:
984
+ raise ValueError(
985
+ f"Expected K3 merged features [tokens, 4, 1024], got {tuple(item.shape)}"
986
+ )
987
+ item = self.pre_norm(item).reshape(item.shape[0], self.input_size)
988
+ outputs.append(self.post_norm(self.fc2(self.act(self.fc1(item)))))
989
+ return outputs
990
+
991
+
992
+ class GroundAnythingVLMVisionModel(nn.Module):
993
+ """MoonViT3D plus the freshly initialized GroundAnything projection connector."""
994
+
995
+ def __init__(self, config: GroundAnythingVLMVisionConfig):
996
+ super().__init__()
997
+ self.config = config
998
+ self.spatial_merge_size = int(config.spatial_merge_size)
999
+ self.vision_tower = MoonViT3dPretrainedModel(config)
1000
+ self.projector = GroundAnythingVLMTwoLayerProjector(config)
1001
+
1002
+ def forward(self, pixel_values, grid_thw=None, patch_positions=None, **kwargs):
1003
+ del patch_positions, kwargs
1004
+ if grid_thw is None:
1005
+ raise ValueError("image_grid_thw is required for K3 MoonViT3D")
1006
+ target_dtype = self.vision_tower.patch_embed.proj.weight.dtype
1007
+ features = self.vision_tower(pixel_values.to(dtype=target_dtype), grid_thw)
1008
+ projected = self.projector(features)
1009
+ merged = torch.cat(projected, dim=0)
1010
+ return BaseModelOutputWithPooling(last_hidden_state=merged)
1011
+
1012
+
1013
+ @auto_docstring
1014
+ class GroundAnythingVLMBaseModel(GroundAnythingVLMPreTrainedModel):
1015
+ base_model_prefix = ""
1016
+ # Reference: fix gemma3 grad acc #37208
1017
+ accepts_loss_kwargs = False
1018
+ config: GroundAnythingVLMConfig
1019
+ _no_split_modules = ["MoonViTEncoderLayer", "Qwen3DecoderLayer"]
1020
+
1021
+ def __init__(self, config: GroundAnythingVLMConfig):
1022
+ super().__init__(config)
1023
+ self.visual = GroundAnythingVLMVisionModel(config.vision_config)
1024
+ self.language_model = AutoModel.from_config(config.text_config)
1025
+ self.streammind_gate = None
1026
+ self._streammind_model_path = None
1027
+
1028
+ # Initialize weights and apply final processing
1029
+ self.post_init()
1030
+
1031
+ def _load_streammind_gate(self):
1032
+ if self.streammind_gate is not None:
1033
+ return self.streammind_gate
1034
+ from safetensors.torch import load_file
1035
+ from .streammind_gate import StreamMindGate
1036
+
1037
+ if not self._streammind_model_path:
1038
+ raise RuntimeError(
1039
+ "StreamMind gate path is unavailable. Load the model with "
1040
+ "GroundAnythingVLMStreamMindForConditionalGeneration.from_pretrained()."
1041
+ )
1042
+ gate = StreamMindGate(self.config.text_config.hidden_size)
1043
+ gate_path = Path(self._streammind_model_path) / "streammind_gate.safetensors"
1044
+ if not gate_path.is_file():
1045
+ from huggingface_hub import hf_hub_download
1046
+
1047
+ gate_path = Path(
1048
+ hf_hub_download(
1049
+ repo_id=self._streammind_model_path,
1050
+ filename="streammind_gate.safetensors",
1051
+ )
1052
+ )
1053
+ state = load_file(str(gate_path))
1054
+ gate.load_state_dict(state, strict=True)
1055
+ gate.to(
1056
+ device=next(self.visual.parameters()).device,
1057
+ dtype=next(self.visual.parameters()).dtype,
1058
+ ).eval()
1059
+ self.streammind_gate = gate
1060
+ return gate
1061
+
1062
+ def _streammind_vision_tokens(self, pixel_values, image_grid_thw, patch_positions=None):
1063
+ pixel_values = pixel_values.type(self.visual.vision_tower.patch_embed.proj.weight.dtype)
1064
+ rope = self.visual.video_rope
1065
+ saved_rope = (rope.inv_freq_t, rope.inv_freq_h, rope.inv_freq_w)
1066
+ try:
1067
+ # The StreamMind checkpoint was trained after the whole OV model,
1068
+ # including non-persistent RoPE buffers, was cast to BF16.
1069
+ rope.inv_freq_t = rope.inv_freq_t.to(pixel_values.dtype)
1070
+ rope.inv_freq_h = rope.inv_freq_h.to(pixel_values.dtype)
1071
+ rope.inv_freq_w = rope.inv_freq_w.to(pixel_values.dtype)
1072
+ vision_output = self.visual(
1073
+ pixel_values,
1074
+ grid_thw=image_grid_thw,
1075
+ patch_positions=patch_positions,
1076
+ )
1077
+ finally:
1078
+ rope.inv_freq_t, rope.inv_freq_h, rope.inv_freq_w = saved_rope
1079
+ merged = vision_output.last_hidden_state.reshape(-1, vision_output.last_hidden_state.shape[-1])
1080
+ merge = self.visual.spatial_merge_size
1081
+ time = int(image_grid_thw[:, 0].sum().item())
1082
+ patches_per_time = int(
1083
+ (image_grid_thw[0, 1].item() // merge)
1084
+ * (image_grid_thw[0, 2].item() // merge)
1085
+ )
1086
+ vision_tokens = merged.reshape(1, time, patches_per_time, -1)
1087
+ return vision_tokens
1088
+
1089
+ def streammind_gate_forward(self, pixel_values, image_grid_thw, patch_positions=None):
1090
+ """Run the gate on one segment without changing the base LLM visual path."""
1091
+ vision_tokens = self._streammind_vision_tokens(
1092
+ pixel_values, image_grid_thw, patch_positions=patch_positions
1093
+ )
1094
+ return self._load_streammind_gate()(vision_tokens)
1095
+
1096
+ def streammind_gate_forward_segments(self, segments):
1097
+ """Run EPFE continuously over a list of codec segments from one stream."""
1098
+ tokens = [
1099
+ self._streammind_vision_tokens(
1100
+ segment["pixel_values"],
1101
+ segment["image_grid_thw"],
1102
+ patch_positions=segment.get("patch_positions"),
1103
+ )
1104
+ for segment in segments
1105
+ ]
1106
+ lengths = [token.shape[1] for token in tokens]
1107
+ boundaries = torch.tensor(lengths).cumsum(0).tolist()
1108
+ return self._load_streammind_gate()(
1109
+ torch.cat(tokens, dim=1), response_positions=boundaries
1110
+ )
1111
+
1112
+ def get_input_embeddings(self):
1113
+ return self.language_model.get_input_embeddings()
1114
+
1115
+ def set_input_embeddings(self, value):
1116
+ self.language_model.set_input_embeddings(value)
1117
+
1118
+ def set_decoder(self, decoder):
1119
+ self.language_model = decoder
1120
+
1121
+ def get_decoder(self):
1122
+ return self.language_model
1123
+
1124
+ def get_video_features(
1125
+ self,
1126
+ pixel_values_videos: torch.FloatTensor,
1127
+ video_grid_thw: Optional[torch.LongTensor] = None,
1128
+ patch_positions=None,
1129
+ ):
1130
+ """
1131
+ Encodes videos into continuous embeddings that can be forwarded to the language model.
1132
+
1133
+ Args:
1134
+ pixel_values_videos: Pre-processed patches from Qwen2VL processor.
1135
+ `torch.FloatTensor` of shape `(total_patches, num_channels, patch_size, patch_size)`
1136
+ video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
1137
+ The temporal, height and width of feature shape of each video in LLM.
1138
+ """
1139
+ # Convert to correct dtype
1140
+ pixel_values_videos = pixel_values_videos.type(self.visual.vision_tower.patch_embed.proj.weight.dtype)
1141
+
1142
+ # Forward through vision model with grid_thw
1143
+ vision_output = self.visual(pixel_values_videos, grid_thw=video_grid_thw, patch_positions=patch_positions)
1144
+
1145
+ # Extract the actual tensor from BaseModelOutputWithPooling
1146
+ if hasattr(vision_output, "last_hidden_state"):
1147
+ video_embeds = vision_output.last_hidden_state
1148
+ else:
1149
+ video_embeds = vision_output[0] # Fallback for tuple output
1150
+
1151
+ # Compute split sizes from video_grid_thw or from input shape
1152
+ if video_grid_thw is not None:
1153
+ split_sizes = (video_grid_thw.prod(-1) // self.visual.spatial_merge_size**2).tolist()
1154
+ else:
1155
+ # Compute from input shape
1156
+ batch_size = pixel_values_videos.shape[0]
1157
+ split_sizes = [video_embeds.shape[1]] * batch_size
1158
+
1159
+ # Split embeddings per video
1160
+ if len(split_sizes) > 1:
1161
+ video_embeds = torch.split(video_embeds.view(-1, video_embeds.shape[-1]), split_sizes)
1162
+ else:
1163
+ video_embeds = [video_embeds.view(-1, video_embeds.shape[-1])]
1164
+
1165
+ return video_embeds
1166
+
1167
+ def get_image_features(
1168
+ self, pixel_values, image_grid_thw: Optional[torch.LongTensor] = None, patch_positions=None
1169
+ ):
1170
+ """
1171
+ Encodes images into continuous embeddings that can be forwarded to the language model.
1172
+
1173
+ Args:
1174
+ pixel_values: Pre-processed patches from Qwen2VL processor.
1175
+ - `torch.FloatTensor` of shape `(total_patches, num_channels, patch_size, patch_size)`
1176
+ image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
1177
+ The temporal, height and width of feature shape of each image in LLM.
1178
+ """
1179
+ # Kimi-K3 processor emits already-unfolded image patches as
1180
+ # [total_patches, channels, patch_height, patch_width].
1181
+ if pixel_values.dim() == 4:
1182
+ # Convert to correct dtype
1183
+ pixel_values = pixel_values.type(self.visual.vision_tower.patch_embed.proj.weight.dtype)
1184
+
1185
+ # Forward through vision model with grid_thw
1186
+ vision_output = self.visual(pixel_values, grid_thw=image_grid_thw, patch_positions=patch_positions)
1187
+
1188
+ # Extract the actual tensor from BaseModelOutputWithPooling
1189
+ if hasattr(vision_output, "last_hidden_state"):
1190
+ image_embeds = vision_output.last_hidden_state
1191
+ else:
1192
+ image_embeds = vision_output[0]
1193
+
1194
+ # Compute split sizes from grid_thw
1195
+ if image_grid_thw is not None:
1196
+ split_sizes = (image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2).tolist()
1197
+ else:
1198
+ # Fallback: assume single image
1199
+ split_sizes = [image_embeds.shape[0] if image_embeds.dim() == 2 else image_embeds.shape[1]]
1200
+
1201
+ # Split embeddings per image
1202
+ image_embeds_flat = image_embeds.view(-1, image_embeds.shape[-1])
1203
+ if len(split_sizes) > 1:
1204
+ image_embeds = list(torch.split(image_embeds_flat, split_sizes))
1205
+ else:
1206
+ image_embeds = [image_embeds_flat]
1207
+
1208
+ return image_embeds
1209
+ else:
1210
+ raise ValueError(
1211
+ f"Unsupported pixel_values shape: expected 4D tensor [total_patches, C, H, W], "
1212
+ f"got {pixel_values.shape if hasattr(pixel_values, 'shape') else type(pixel_values)}"
1213
+ )
1214
+
1215
+ def get_placeholder_mask(
1216
+ self,
1217
+ input_ids: torch.LongTensor,
1218
+ inputs_embeds: torch.FloatTensor,
1219
+ image_features: Optional[torch.FloatTensor] = None,
1220
+ video_features: Optional[torch.FloatTensor] = None,
1221
+ ):
1222
+ """
1223
+ Obtains multimodal placeholder mask from `input_ids` or `inputs_embeds`, and checks that the placeholder token count is
1224
+ equal to the length of multimodal features. If the lengths are different, an error is raised.
1225
+ """
1226
+ if input_ids is None:
1227
+ special_image_mask = inputs_embeds == self.get_input_embeddings()(
1228
+ torch.tensor(self.config.image_token_id, dtype=torch.long, device=inputs_embeds.device)
1229
+ )
1230
+ special_image_mask = special_image_mask.all(-1)
1231
+ special_video_mask = inputs_embeds == self.get_input_embeddings()(
1232
+ torch.tensor(self.config.video_token_id, dtype=torch.long, device=inputs_embeds.device)
1233
+ )
1234
+ special_video_mask = special_video_mask.all(-1)
1235
+ else:
1236
+ special_image_mask = input_ids == self.config.image_token_id
1237
+ special_video_mask = input_ids == self.config.video_token_id
1238
+
1239
+ n_image_tokens = special_image_mask.sum()
1240
+ special_image_mask = special_image_mask.unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device)
1241
+ if image_features is not None and inputs_embeds[special_image_mask].numel() != image_features.numel():
1242
+ raise ValueError(
1243
+ f"Image features and image tokens do not match: tokens: {n_image_tokens}, features {image_features.shape[0]}"
1244
+ )
1245
+
1246
+ n_video_tokens = special_video_mask.sum()
1247
+ special_video_mask = special_video_mask.unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device)
1248
+ if video_features is not None and inputs_embeds[special_video_mask].numel() != video_features.numel():
1249
+ raise ValueError(
1250
+ f"Videos features and video tokens do not match: tokens: {n_video_tokens}, features {video_features.shape[0]}"
1251
+ )
1252
+
1253
+ return special_image_mask, special_video_mask
1254
+
1255
+ @auto_docstring
1256
+ def forward(
1257
+ self,
1258
+ input_ids: Optional[torch.LongTensor] = None,
1259
+ attention_mask: Optional[torch.Tensor] = None,
1260
+ position_ids: Optional[torch.LongTensor] = None,
1261
+ past_key_values: Optional[Cache] = None,
1262
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1263
+ use_cache: Optional[bool] = None,
1264
+ output_attentions: Optional[bool] = None,
1265
+ output_hidden_states: Optional[bool] = None,
1266
+ return_dict: Optional[bool] = None,
1267
+ pixel_values: Optional[torch.Tensor] = None,
1268
+ pixel_values_videos: Optional[torch.FloatTensor] = None,
1269
+ image_grid_thw: Optional[torch.LongTensor] = None,
1270
+ patch_positions: Optional[torch.LongTensor] = None,
1271
+ video_grid_thw: Optional[torch.LongTensor] = None,
1272
+ cache_position: Optional[torch.LongTensor] = None,
1273
+ second_per_grid_ts: Optional[torch.Tensor] = None,
1274
+ **kwargs: Unpack[TransformersKwargs],
1275
+ ) -> Union[tuple, GroundAnythingVLMBaseModelOutputWithPast]:
1276
+ r"""
1277
+ image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
1278
+ The temporal, height and width of feature shape of each image in LLM.
1279
+ video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
1280
+ The temporal, height and width of feature shape of each video in LLM.
1281
+ patch_positions (`torch.LongTensor` of shape `(total_patches, 3)` or `(1, total_patches, 3)`, *optional*):
1282
+ Explicit per-patch `(t, h, w)` position indices used by the vision tower to compute 3D rotary
1283
+ position embeddings (and the optional absolute position embedding inside the patch merger).
1284
+ `total_patches` is the sum of `t * h * w` across all images and videos in the batch, matching
1285
+ the layout produced by the Qwen2VL-style image processor.
1286
+ second_per_grid_ts (`torch.Tensor` of shape `(num_videos)`, *optional*):
1287
+ The time interval (in seconds) for each grid along the temporal dimension in the 3D position IDs.
1288
+ cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
1289
+ Indices depicting the position of the input sequence tokens in the sequence. Contrarily to
1290
+ `position_ids`, this tensor is not affected by padding.
1291
+
1292
+ Note: see the top-level ``GroundAnythingVLMBaseForConditionalGeneration.forward``
1293
+ docstring; currently video flows in via the ``image_grid_thw`` / ``pixel_values``
1294
+ alias, so ``pixel_values_videos`` / ``video_grid_thw`` /
1295
+ ``second_per_grid_ts`` are unused at this layer.
1296
+ """
1297
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1298
+ output_hidden_states = (
1299
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1300
+ )
1301
+ return_dict = True if return_dict is None else return_dict
1302
+
1303
+ if inputs_embeds is None:
1304
+ inputs_embeds = self.get_input_embeddings()(input_ids)
1305
+
1306
+ image_embeds = None
1307
+
1308
+ if pixel_values is not None:
1309
+ image_embeds = self.get_image_features(pixel_values, image_grid_thw, patch_positions=patch_positions)
1310
+
1311
+ if image_embeds is not None:
1312
+ image_embeds = torch.cat(image_embeds, dim=0).to(inputs_embeds.device, inputs_embeds.dtype)
1313
+ image_mask, _ = self.get_placeholder_mask(
1314
+ input_ids, inputs_embeds=inputs_embeds, image_features=image_embeds
1315
+ )
1316
+ inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)
1317
+
1318
+ if pixel_values_videos is not None:
1319
+ video_embeds = self.get_video_features(
1320
+ pixel_values_videos, video_grid_thw, patch_positions=patch_positions
1321
+ )
1322
+ video_embeds = torch.cat(video_embeds, dim=0).to(inputs_embeds.device, inputs_embeds.dtype)
1323
+ _, video_mask = self.get_placeholder_mask(
1324
+ input_ids, inputs_embeds=inputs_embeds, video_features=video_embeds
1325
+ )
1326
+ inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)
1327
+
1328
+ # Use simple 1D position_ids
1329
+ if position_ids is None:
1330
+ batch_size, seq_length, _ = inputs_embeds.shape
1331
+ if attention_mask is not None:
1332
+ position_ids = attention_mask.long().cumsum(-1) - 1
1333
+ position_ids.masked_fill_(attention_mask == 0, 1)
1334
+ else:
1335
+ position_ids = (
1336
+ torch.arange(seq_length, device=inputs_embeds.device).unsqueeze(0).expand(batch_size, -1)
1337
+ )
1338
+
1339
+ # Handle cache_position for generation
1340
+ if cache_position is not None and cache_position[0] != 0:
1341
+ position_ids = position_ids + cache_position[0]
1342
+
1343
+ outputs = self.language_model(
1344
+ input_ids=None,
1345
+ position_ids=position_ids,
1346
+ attention_mask=attention_mask,
1347
+ past_key_values=past_key_values,
1348
+ inputs_embeds=inputs_embeds,
1349
+ use_cache=use_cache,
1350
+ output_attentions=output_attentions,
1351
+ output_hidden_states=output_hidden_states,
1352
+ return_dict=True,
1353
+ cache_position=cache_position,
1354
+ **kwargs,
1355
+ )
1356
+
1357
+ output = GroundAnythingVLMBaseModelOutputWithPast(
1358
+ last_hidden_state=outputs.last_hidden_state,
1359
+ past_key_values=outputs.past_key_values,
1360
+ hidden_states=outputs.hidden_states,
1361
+ attentions=outputs.attentions,
1362
+ )
1363
+ return output if return_dict else output.to_tuple()
1364
+
1365
+
1366
+ @auto_docstring
1367
+ class GroundAnythingVLMBaseForConditionalGeneration(GroundAnythingVLMPreTrainedModel, GenerationMixin):
1368
+ _tied_weights_keys = {"lm_head.weight": "model.language_model.embed_tokens.weight"}
1369
+ # Reference: fix gemma3 grad acc #37208
1370
+ accepts_loss_kwargs = False
1371
+
1372
+ def __init__(self, config):
1373
+ super().__init__(config)
1374
+ self.model = GroundAnythingVLMBaseModel(config)
1375
+ self.lm_head = nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False)
1376
+ self.post_init()
1377
+
1378
+ @classmethod
1379
+ def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
1380
+ model = super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
1381
+ model.model._streammind_model_path = str(pretrained_model_name_or_path)
1382
+ return model
1383
+
1384
+ def get_input_embeddings(self):
1385
+ return self.model.get_input_embeddings()
1386
+
1387
+ def set_input_embeddings(self, value):
1388
+ self.model.set_input_embeddings(value)
1389
+
1390
+ def set_decoder(self, decoder):
1391
+ self.model.set_decoder(decoder)
1392
+
1393
+ def get_decoder(self):
1394
+ return self.model.get_decoder()
1395
+
1396
+ def get_video_features(
1397
+ self,
1398
+ pixel_values_videos: torch.FloatTensor,
1399
+ video_grid_thw: Optional[torch.LongTensor] = None,
1400
+ patch_positions=None,
1401
+ ):
1402
+ return self.model.get_video_features(pixel_values_videos, video_grid_thw, patch_positions=patch_positions)
1403
+
1404
+ def get_image_features(self, pixel_values: torch.FloatTensor, image_grid_thw: Optional[torch.LongTensor] = None):
1405
+ return self.model.get_image_features(pixel_values, image_grid_thw)
1406
+
1407
+ # Make modules available through conditional class for BC
1408
+ @property
1409
+ def language_model(self):
1410
+ return self.model.language_model
1411
+
1412
+ @property
1413
+ def visual(self):
1414
+ return self.model.visual
1415
+
1416
+ def streammind_gate_forward(self, pixel_values, image_grid_thw, patch_positions=None):
1417
+ return self.model.streammind_gate_forward(
1418
+ pixel_values,
1419
+ image_grid_thw,
1420
+ patch_positions=patch_positions,
1421
+ )
1422
+
1423
+ def streammind_gate_forward_segments(self, segments):
1424
+ return self.model.streammind_gate_forward_segments(segments)
1425
+
1426
+ @can_return_tuple
1427
+ @auto_docstring
1428
+ def forward(
1429
+ self,
1430
+ input_ids: Optional[torch.LongTensor] = None,
1431
+ attention_mask: Optional[torch.Tensor] = None,
1432
+ position_ids: Optional[torch.LongTensor] = None,
1433
+ past_key_values: Optional[Cache] = None,
1434
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1435
+ labels: Optional[torch.LongTensor] = None,
1436
+ use_cache: Optional[bool] = None,
1437
+ output_attentions: Optional[bool] = None,
1438
+ output_hidden_states: Optional[bool] = None,
1439
+ pixel_values: Optional[torch.Tensor] = None,
1440
+ pixel_values_videos: Optional[torch.FloatTensor] = None,
1441
+ image_grid_thw: Optional[torch.LongTensor] = None,
1442
+ patch_positions: Optional[torch.LongTensor] = None,
1443
+ video_grid_thw: Optional[torch.LongTensor] = None,
1444
+ cache_position: Optional[torch.LongTensor] = None,
1445
+ second_per_grid_ts: Optional[torch.Tensor] = None,
1446
+ logits_to_keep: Union[int, torch.Tensor] = 0,
1447
+ **kwargs: Unpack[TransformersKwargs],
1448
+ ) -> Union[tuple, GroundAnythingVLMCausalLMOutputWithPast]:
1449
+ r"""
1450
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1451
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1452
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1453
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1454
+ image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
1455
+ The temporal, height and width of feature shape of each image in LLM.
1456
+ video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
1457
+ The temporal, height and width of feature shape of each video in LLM.
1458
+ patch_positions (`torch.LongTensor` of shape `(total_patches, 3)` or `(1, total_patches, 3)`, *optional*):
1459
+ Explicit per-patch `(t, h, w)` position indices used by the vision tower to compute 3D rotary
1460
+ position embeddings (and the optional absolute position embedding inside the patch merger).
1461
+ `total_patches` is the sum of `t * h * w` across all images and videos in the batch, matching
1462
+ the layout produced by the Qwen2VL-style image processor.
1463
+ second_per_grid_ts (`torch.Tensor` of shape `(num_videos)`, *optional*):
1464
+ The time interval (in seconds) for each grid along the temporal dimension in the 3D position IDs.
1465
+ cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
1466
+ Indices depicting the position of the input sequence tokens in the sequence. Contrarily to
1467
+ `position_ids`, this tensor is not affected by padding.
1468
+
1469
+ Note (native-video alias):
1470
+ The companion ``GroundAnythingVLMProcessor.__call__(videos=...)`` does NOT
1471
+ pass ``pixel_values_videos`` / ``video_grid_thw`` / ``second_per_grid_ts``
1472
+ to this forward. Instead it aliases the video patch tensor as
1473
+ ``pixel_values=`` and ``image_grid_thw=``, so video inputs share the
1474
+ same code path as multi-image inputs (vision is purely
1475
+ spatial; temporal information is carried by per-frame ``<X.X seconds>``
1476
+ text tags emitted by the processor). The ``*_videos`` and
1477
+ ``second_per_grid_ts`` kwargs are kept declared here only for API
1478
+ completeness and future use (e.g. 3D mRoPE / ``get_rope_index``); they
1479
+ are NOT consumed by the current vision encoder.
1480
+
1481
+ Example:
1482
+
1483
+ ```python
1484
+ >>> from PIL import Image
1485
+ >>> import requests
1486
+ >>> from transformers import AutoProcessor, GroundAnythingVLMBaseForConditionalGeneration
1487
+
1488
+ >>> model = GroundAnythingVLMBaseForConditionalGeneration.from_pretrained("/path/to/GroundAnything-VLM", trust_remote_code=True)
1489
+ >>> processor = AutoProcessor.from_pretrained("/path/to/GroundAnything-VLM", trust_remote_code=True)
1490
+
1491
+ >>> messages = [
1492
+ {
1493
+ "role": "user",
1494
+ "content": [
1495
+ {"type": "image"},
1496
+ {"type": "text", "text": "What is shown in this image?"},
1497
+ ],
1498
+ },
1499
+ ]
1500
+ >>> url = "https://www.ilankelman.org/stopsigns/australia.jpg"
1501
+ >>> image = Image.open(requests.get(url, stream=True).raw)
1502
+
1503
+ >>> text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
1504
+ >>> inputs = processor(text=[text], images=[image], return_tensors="pt")
1505
+
1506
+ >>> # Generate
1507
+ >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
1508
+ >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
1509
+ "The image shows a street scene with a red stop sign in the foreground. In the background, there is a large red gate with Chinese characters ..."
1510
+ ```"""
1511
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1512
+ output_hidden_states = (
1513
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1514
+ )
1515
+ outputs = self.model(
1516
+ input_ids=input_ids,
1517
+ pixel_values=pixel_values,
1518
+ pixel_values_videos=pixel_values_videos,
1519
+ image_grid_thw=image_grid_thw,
1520
+ patch_positions=patch_positions,
1521
+ video_grid_thw=video_grid_thw,
1522
+ second_per_grid_ts=second_per_grid_ts,
1523
+ position_ids=position_ids,
1524
+ attention_mask=attention_mask,
1525
+ past_key_values=past_key_values,
1526
+ inputs_embeds=inputs_embeds,
1527
+ use_cache=use_cache,
1528
+ output_attentions=output_attentions,
1529
+ output_hidden_states=output_hidden_states,
1530
+ return_dict=True,
1531
+ cache_position=cache_position,
1532
+ **kwargs,
1533
+ )
1534
+
1535
+ hidden_states = outputs[0]
1536
+
1537
+ # Only compute necessary logits, and do not upcast them to float if we are not computing the loss
1538
+ slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
1539
+ logits = self.lm_head(hidden_states[:, slice_indices, :])
1540
+
1541
+ loss = None
1542
+ if labels is not None:
1543
+ loss = self.loss_function(
1544
+ logits=logits, labels=labels, vocab_size=self.config.text_config.vocab_size, **kwargs
1545
+ )
1546
+
1547
+ # A packed batch can be text-only on one distributed rank. Keep every
1548
+ # trainable projector tensor in that rank's graph so ZeRO launches the
1549
+ # same gradient collectives as ranks that received visual samples.
1550
+ projector_zero_anchor = None
1551
+ for parameter in self.model.visual.projector.parameters():
1552
+ if parameter.requires_grad:
1553
+ term = parameter.reshape(-1)[0] * 0.0
1554
+ projector_zero_anchor = term if projector_zero_anchor is None else projector_zero_anchor + term
1555
+ if projector_zero_anchor is not None:
1556
+ loss = loss + projector_zero_anchor
1557
+
1558
+ return GroundAnythingVLMCausalLMOutputWithPast(
1559
+ loss=loss,
1560
+ logits=logits,
1561
+ past_key_values=outputs.past_key_values,
1562
+ hidden_states=outputs.hidden_states,
1563
+ attentions=outputs.attentions,
1564
+ )
1565
+
1566
+
1567
+
1568
+ def prepare_inputs_for_generation(
1569
+ self,
1570
+ input_ids,
1571
+ past_key_values=None,
1572
+ attention_mask=None,
1573
+ inputs_embeds=None,
1574
+ cache_position=None,
1575
+ position_ids=None,
1576
+ use_cache=True,
1577
+ pixel_values=None,
1578
+ pixel_values_videos=None,
1579
+ image_grid_thw=None,
1580
+ patch_positions=None,
1581
+ video_grid_thw=None,
1582
+ second_per_grid_ts=None,
1583
+ is_first_iteration=False,
1584
+ **kwargs,
1585
+ ):
1586
+ # Overwritten -- in specific circumstances we don't want to forward image inputs to the model
1587
+ model_inputs = super().prepare_inputs_for_generation(
1588
+ input_ids,
1589
+ past_key_values=past_key_values,
1590
+ attention_mask=attention_mask,
1591
+ inputs_embeds=inputs_embeds,
1592
+ cache_position=cache_position,
1593
+ position_ids=position_ids,
1594
+ pixel_values=pixel_values,
1595
+ pixel_values_videos=pixel_values_videos,
1596
+ image_grid_thw=image_grid_thw,
1597
+ video_grid_thw=video_grid_thw,
1598
+ second_per_grid_ts=second_per_grid_ts,
1599
+ patch_positions=patch_positions,
1600
+ use_cache=use_cache,
1601
+ is_first_iteration=is_first_iteration,
1602
+ **kwargs,
1603
+ )
1604
+
1605
+ # After the prefill iteration, drop image inputs so the vision tower
1606
+ # isn't re-run on decode steps. Gating on `is_first_iteration` (the
1607
+ # Qwen3-VL convention) is the only reliable signal in transformers
1608
+ # 5.x: `past_key_values` is non-None even on the first call (an empty
1609
+ # DynamicCache is created up-front by `generate`), and `cache_position`
1610
+ # may be `None` for remote-code models.
1611
+ if not is_first_iteration and use_cache:
1612
+ model_inputs["pixel_values"] = None
1613
+ model_inputs["pixel_values_videos"] = None
1614
+
1615
+ return model_inputs
1616
+
1617
+ def _get_image_nums_and_video_nums(
1618
+ self,
1619
+ input_ids: Optional[torch.LongTensor],
1620
+ inputs_embeds: Optional[torch.Tensor] = None,
1621
+ ) -> tuple[torch.Tensor, torch.Tensor]:
1622
+ """
1623
+ Get the number of images and videos for each sample to calculate the separation length of the sample tensor.
1624
+ These parameters are not passed through the processor to avoid unpredictable impacts from interface modifications.
1625
+
1626
+ Args:
1627
+ input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
1628
+ Indices of input sequence tokens in the vocabulary.
1629
+
1630
+ Returns:
1631
+ image_nums (`torch.LongTensor` of shape `(batch_size, num_images_sample)`)
1632
+ video_nums (`torch.LongTensor` of shape `(batch_size, num_videos_sample)`)
1633
+ """
1634
+ image_token_id = self.config.image_token_id
1635
+ video_token_id = self.config.video_token_id
1636
+ vision_start_token_id = self.config.vision_start_token_id
1637
+
1638
+ if inputs_embeds is not None:
1639
+ vision_start_mask = (
1640
+ inputs_embeds
1641
+ == self.get_input_embeddings()(
1642
+ torch.tensor(vision_start_token_id, dtype=torch.long, device=inputs_embeds.device)
1643
+ )
1644
+ )[..., 0]
1645
+ image_mask = (
1646
+ inputs_embeds
1647
+ == self.get_input_embeddings()(
1648
+ torch.tensor(image_token_id, dtype=torch.long, device=inputs_embeds.device)
1649
+ )
1650
+ )[..., 0]
1651
+ video_mask = (
1652
+ inputs_embeds
1653
+ == self.get_input_embeddings()(
1654
+ torch.tensor(video_token_id, dtype=torch.long, device=inputs_embeds.device)
1655
+ )
1656
+ )[..., 0]
1657
+ else:
1658
+ vision_start_mask = input_ids == vision_start_token_id
1659
+ image_mask = input_ids == image_token_id
1660
+ video_mask = input_ids == video_token_id
1661
+
1662
+ vision_first_mask = torch.roll(vision_start_mask, shifts=1, dims=1)
1663
+ image_nums = torch.sum(vision_first_mask & image_mask, dim=1)
1664
+ video_nums = torch.sum(vision_first_mask & video_mask, dim=1)
1665
+
1666
+ return image_nums, video_nums
1667
+
1668
+ def _expand_inputs_for_generation(
1669
+ self,
1670
+ expand_size: int = 1,
1671
+ is_encoder_decoder: bool = False,
1672
+ input_ids: Optional[torch.LongTensor] = None,
1673
+ **model_kwargs,
1674
+ ) -> tuple[torch.LongTensor, dict[str, Any]]:
1675
+ # Overwritten -- Support for expanding tensors without a batch size dimension
1676
+ # e.g., pixel_values, image_grid_thw, pixel_values_videos, video_grid_thw, second_per_grid_t
1677
+ # pixel_values.shape[0] is sum(seqlen_images for samples)
1678
+ # image_grid_thw.shape[0] is sum(num_images for samples)
1679
+
1680
+ if expand_size == 1:
1681
+ return input_ids, model_kwargs
1682
+
1683
+ visual_keys = [
1684
+ "pixel_values",
1685
+ "image_grid_thw",
1686
+ "pixel_values_videos",
1687
+ "video_grid_thw",
1688
+ "second_per_grid_ts",
1689
+ "patch_positions",
1690
+ ]
1691
+
1692
+ def _expand_dict_for_generation_visual(dict_to_expand):
1693
+ image_grid_thw = model_kwargs.get("image_grid_thw", None)
1694
+ video_grid_thw = model_kwargs.get("video_grid_thw", None)
1695
+ image_nums, video_nums = self._get_image_nums_and_video_nums(
1696
+ input_ids, inputs_embeds=model_kwargs.get("inputs_embeds", None)
1697
+ )
1698
+
1699
+ def _repeat_interleave_samples(x, lengths, repeat_times):
1700
+ samples = torch.split(x, lengths)
1701
+ repeat_args = [repeat_times] + [1] * (x.dim() - 1)
1702
+ result = torch.cat([sample.repeat(*repeat_args) for sample in samples], dim=0)
1703
+ return result
1704
+
1705
+ for key in dict_to_expand:
1706
+ if key == "pixel_values":
1707
+ # split images into samples
1708
+ samples = torch.split(image_grid_thw, list(image_nums))
1709
+ # compute the sequence length of images for each sample
1710
+ lengths = [torch.prod(sample, dim=1).sum() for sample in samples]
1711
+ dict_to_expand[key] = _repeat_interleave_samples(
1712
+ dict_to_expand[key], lengths=lengths, repeat_times=expand_size
1713
+ )
1714
+ elif key == "image_grid_thw":
1715
+ # get the num of images for each sample
1716
+ lengths = list(image_nums)
1717
+ dict_to_expand[key] = _repeat_interleave_samples(
1718
+ dict_to_expand[key], lengths=lengths, repeat_times=expand_size
1719
+ )
1720
+ elif key == "pixel_values_videos":
1721
+ samples = torch.split(video_grid_thw, list(video_nums))
1722
+ lengths = [torch.prod(sample, dim=1).sum() for sample in samples]
1723
+ dict_to_expand[key] = _repeat_interleave_samples(
1724
+ dict_to_expand[key], lengths=lengths, repeat_times=expand_size
1725
+ )
1726
+ elif key == "video_grid_thw":
1727
+ lengths = list(video_nums)
1728
+ dict_to_expand[key] = _repeat_interleave_samples(
1729
+ dict_to_expand[key], lengths=lengths, repeat_times=expand_size
1730
+ )
1731
+ elif key == "second_per_grid_ts":
1732
+ dict_to_expand[key] = _repeat_interleave_samples(
1733
+ dict_to_expand[key], lengths=list(video_nums), repeat_times=expand_size
1734
+ )
1735
+ elif key == "patch_positions":
1736
+ if image_grid_thw is not None and image_grid_thw.numel() > 0 and image_nums.sum() > 0:
1737
+ samples = torch.split(image_grid_thw, list(image_nums))
1738
+ lengths = [torch.prod(sample, dim=1).sum() for sample in samples]
1739
+ elif video_grid_thw is not None and video_grid_thw.numel() > 0 and video_nums.sum() > 0:
1740
+ samples = torch.split(video_grid_thw, list(video_nums))
1741
+ lengths = [torch.prod(sample, dim=1).sum() for sample in samples]
1742
+ else:
1743
+ continue
1744
+ dict_to_expand[key] = _repeat_interleave_samples(
1745
+ dict_to_expand[key], lengths=lengths, repeat_times=expand_size
1746
+ )
1747
+ return dict_to_expand
1748
+
1749
+ def _expand_dict_for_generation(dict_to_expand):
1750
+ for key in dict_to_expand:
1751
+ if (
1752
+ key != "cache_position"
1753
+ and dict_to_expand[key] is not None
1754
+ and isinstance(dict_to_expand[key], torch.Tensor)
1755
+ and key not in visual_keys
1756
+ ):
1757
+ dict_to_expand[key] = dict_to_expand[key].repeat_interleave(expand_size, dim=0)
1758
+ return dict_to_expand
1759
+
1760
+ model_kwargs = _expand_dict_for_generation_visual(model_kwargs)
1761
+
1762
+ if input_ids is not None:
1763
+ input_ids = input_ids.repeat_interleave(expand_size, dim=0)
1764
+
1765
+ model_kwargs = _expand_dict_for_generation(model_kwargs)
1766
+
1767
+ if is_encoder_decoder:
1768
+ if model_kwargs.get("encoder_outputs") is None:
1769
+ raise ValueError("If `is_encoder_decoder` is True, make sure that `encoder_outputs` is defined.")
1770
+ model_kwargs["encoder_outputs"] = _expand_dict_for_generation(model_kwargs["encoder_outputs"])
1771
+
1772
+ return input_ids, model_kwargs
1773
+
1774
+
1775
+ class GroundAnythingVLMModel(GroundAnythingVLMBaseModel):
1776
+ """Named GroundAnything-VLM base-model entry point."""
1777
+
1778
+
1779
+ class GroundAnythingVLMForConditionalGeneration(GroundAnythingVLMBaseForConditionalGeneration):
1780
+ """GroundAnything Qwen3-4B with the Kimi-K3 MoonViT3D visual encoder."""
1781
+
1782
+
1783
+ __all__ = [
1784
+ "GroundAnythingVLMForConditionalGeneration",
1785
+ "GroundAnythingVLMModel",
1786
+ "GroundAnythingVLMBaseForConditionalGeneration",
1787
+ "GroundAnythingVLMBaseModel",
1788
+ "GroundAnythingVLMPreTrainedModel",
1789
+ ]
1790
+
1791
+
1792
+ class GroundAnythingModel(GroundAnythingVLMModel):
1793
+ """DLM release model identity; SGLang loads wrapper weights separately."""
1794
+
1795
+ config_class = GroundAnythingConfig
1796
+
1797
+
1798
+ class GroundAnythingForConditionalGeneration(GroundAnythingVLMForConditionalGeneration):
1799
+ config_class = GroundAnythingConfig
modeling_groundinganything_vision.py ADDED
@@ -0,0 +1,1345 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2025-2026 The Moonshot AI Team and HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # The code is based on llava (llava/modeling_llava.py), but modified for Kimi-K3.
5
+ #
6
+ # Licensing Information:
7
+ # - Code derived from llava (llava/modeling_llava.py) is licensed under the Apache License, Version 2.0.
8
+ # - Other parts of the code are licensed under the Kimi K3 License (see the LICENSE file in this repository).
9
+ #
10
+ # Apache License, Version 2.0:
11
+ # Licensed under the Apache License, Version 2.0 (the "License");
12
+ # you may not use this file except in compliance with the License.
13
+ # You may obtain a copy of the License at
14
+ #
15
+ # http://www.apache.org/licenses/LICENSE-2.0
16
+ #
17
+ # Unless required by applicable law or agreed to in writing, software
18
+ # distributed under the License is distributed on an "AS IS" BASIS,
19
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
20
+ # See the License for the specific language governing permissions and
21
+ # limitations under the License.
22
+
23
+
24
+ # NOTE: Reference implementation for model architecture; see the model card for production deployment.
25
+ import math
26
+ from collections.abc import Sequence
27
+ from copy import deepcopy
28
+ from typing import Optional
29
+
30
+ import numpy as np
31
+ import torch
32
+ import torch.nn as nn
33
+ import torch.nn.functional as F
34
+ from transformers import activations
35
+
36
+ try:
37
+ from transformers.activations import PytorchGELUTanh
38
+ except ImportError:
39
+ from transformers.activations import GELUTanh
40
+ activations.PytorchGELUTanh = GELUTanh
41
+ PytorchGELUTanh = GELUTanh
42
+ from transformers.activations import PytorchGELUTanh
43
+ from transformers.configuration_utils import PretrainedConfig
44
+ from transformers.modeling_utils import PreTrainedModel
45
+ from transformers.models.llava.modeling_llava import \
46
+ LlavaCausalLMOutputWithPast
47
+ from transformers.utils import is_flash_attn_2_available
48
+
49
+ GroundAnythingBackboneConfig = PretrainedConfig
50
+ # GroundAnything-VLM imports only MoonViT3D from this reference module. Keep the
51
+ # Kimi text implementation optional so loading the vision tower does not
52
+ # require fla-core.
53
+ GroundAnythingBackboneLinearForCausalLM = None
54
+
55
+ # Flash attention imports
56
+ if is_flash_attn_2_available():
57
+ from flash_attn import flash_attn_varlen_func
58
+ else:
59
+ flash_attn_varlen_func = None
60
+
61
+
62
+ def multihead_attention(
63
+ q: torch.Tensor,
64
+ k: torch.Tensor,
65
+ v: torch.Tensor,
66
+ q_cu_seqlens: torch.Tensor | None = None,
67
+ k_cu_seqlens: torch.Tensor | None = None,
68
+ max_seqlen_q: int | None = None,
69
+ max_seqlen_k: int | None = None,
70
+ deterministic: bool = False,
71
+ ):
72
+ """Multi-head attention using flash attention 2.
73
+
74
+ Args:
75
+ q, k, v: tensor of shape (batch_size, seqlen, num_heads, head_dim),
76
+ or (tot_seqlens, num_heads, head_dim) if packing.
77
+ q_cu_seqlens (torch.Tensor): cumulative sequence lengths of q.
78
+ The first element should be 0 and the last element should be q.shape[0].
79
+ k_cu_seqlens (torch.Tensor): cumulative sequence lengths of k.
80
+ The first element should be 0 and the last element should be k.shape[0].
81
+
82
+ Returns:
83
+ output: shape (batch_size, seqlen, dim) or (tot_seqlens, dim) if packing,
84
+ where dim = num_heads * head_dim
85
+ """
86
+ attn_out = flash_attn_varlen_func(
87
+ q,
88
+ k,
89
+ v,
90
+ q_cu_seqlens,
91
+ k_cu_seqlens,
92
+ max_seqlen_q,
93
+ max_seqlen_k,
94
+ causal=False,
95
+ deterministic=deterministic,
96
+ )
97
+ if isinstance(attn_out, tuple):
98
+ attn_out = attn_out[0]
99
+
100
+ attn_out = attn_out.flatten(start_dim=-2)
101
+
102
+ return attn_out
103
+
104
+
105
+ def eager_attention(
106
+ q: torch.Tensor,
107
+ k: torch.Tensor,
108
+ v: torch.Tensor,
109
+ q_cu_seqlens: Optional[torch.Tensor] = None,
110
+ k_cu_seqlens: Optional[torch.Tensor] = None,
111
+ **kwargs,
112
+ ) -> torch.Tensor:
113
+ seq_length = q.shape[0]
114
+ attention_mask = torch.zeros([1, seq_length, seq_length],
115
+ device=q.device,
116
+ dtype=torch.bool)
117
+ for i in range(1, len(q_cu_seqlens)):
118
+ attention_mask[
119
+ ...,
120
+ q_cu_seqlens[i - 1]:q_cu_seqlens[i],
121
+ q_cu_seqlens[i - 1]:q_cu_seqlens[i],
122
+ ] = True
123
+ q = q.transpose(0, 1)
124
+ k = k.transpose(0, 1)
125
+ v = v.transpose(0, 1)
126
+
127
+ attn_weight = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1])
128
+ attn_weight = attn_weight.masked_fill(
129
+ ~attention_mask, torch.finfo(attn_weight.dtype).min)
130
+ attn_weight = torch.softmax(attn_weight, dim=-1,
131
+ dtype=torch.float32).to(q.dtype)
132
+
133
+ attn_output = attn_weight @ v
134
+ attn_output = attn_output.transpose(0, 1)
135
+ attn_output = attn_output.reshape(seq_length, -1)
136
+ return attn_output
137
+
138
+
139
+ def sdpa_attention(
140
+ q: torch.Tensor,
141
+ k: torch.Tensor,
142
+ v: torch.Tensor,
143
+ q_cu_seqlens: torch.Tensor | None = None,
144
+ k_cu_seqlens: torch.Tensor | None = None,
145
+ **kwargs,
146
+ ) -> torch.Tensor:
147
+ del k_cu_seqlens, kwargs
148
+ outputs = []
149
+ for index in range(1, len(q_cu_seqlens)):
150
+ start = int(q_cu_seqlens[index - 1])
151
+ end = int(q_cu_seqlens[index])
152
+ query = q[start:end].transpose(0, 1)
153
+ key = k[start:end].transpose(0, 1)
154
+ value = v[start:end].transpose(0, 1)
155
+ output = F.scaled_dot_product_attention(
156
+ query, key, value, dropout_p=0.0, is_causal=False
157
+ )
158
+ outputs.append(output.transpose(0, 1))
159
+ return torch.cat(outputs, dim=0).flatten(start_dim=-2)
160
+
161
+
162
+ VL_VISION_ATTENTION_FUNCTIONS = {
163
+ "flash_attention_2": multihead_attention,
164
+ "eager": eager_attention,
165
+ "sdpa": sdpa_attention,
166
+ }
167
+
168
+
169
+ def _apply_rope_input_validation(x, freqs_cis):
170
+ assert x.ndim == freqs_cis.ndim + 1, (x.shape, freqs_cis.shape)
171
+ assert x.shape[:-2] == freqs_cis.shape[:-1], (x.shape, freqs_cis.shape)
172
+ assert x.shape[-1] == 2 * freqs_cis.shape[-1], (x.shape, freqs_cis.shape)
173
+ assert freqs_cis.dtype == torch.complex64, freqs_cis.dtype
174
+
175
+
176
+ def get_rope_shape_decorate(func):
177
+ _get_rope_shape_first_call_flag = set()
178
+
179
+ def wrapper(org, interpolation_mode, shape):
180
+ key = (org.requires_grad, torch.is_grad_enabled(), interpolation_mode)
181
+ if key not in _get_rope_shape_first_call_flag:
182
+ _get_rope_shape_first_call_flag.add(key)
183
+ _ = func(org, interpolation_mode, shape=(64, 64))
184
+ return func(org, interpolation_mode, shape)
185
+
186
+ return wrapper
187
+
188
+
189
+ @get_rope_shape_decorate
190
+ @torch.compile(dynamic=True)
191
+ def get_rope_shape(org, interpolation_mode, shape):
192
+ return (F.interpolate(
193
+ org.permute((2, 0, 1)).unsqueeze(0),
194
+ size=shape,
195
+ mode=interpolation_mode,
196
+ ).squeeze(0).permute((1, 2, 0)).flatten(end_dim=1))
197
+
198
+
199
+ def apply_rope(xq: torch.Tensor, xk: torch.Tensor,
200
+ freqs_cis: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
201
+ """
202
+ Args: (The leading dimensions of all inputs should be the same)
203
+ xq: query, tensor of shape (..., num_heads, head_dim)
204
+ xk: key, tensor of shape (..., num_heads, head_dim)
205
+ freqs_cis: tensor of shape (..., head_dim/2), dtype=torch.complex64. It contains the precomputed cis(freqs) for each position in the 2D grid.
206
+ Returns:
207
+ xq_out, xk_out: tensors of shape (..., num_heads, head_dim)
208
+ """
209
+ _apply_rope_input_validation(xq, freqs_cis)
210
+ _apply_rope_input_validation(xk, freqs_cis)
211
+
212
+ freqs_cis = freqs_cis.unsqueeze(-2) # ..., 1, head_dim/2
213
+ # ..., num_heads, head_dim/2
214
+ xq_ = torch.view_as_complex(xq.float().view(*xq.shape[:-1], -1, 2))
215
+ xk_ = torch.view_as_complex(xk.float().view(*xq.shape[:-1], -1, 2))
216
+ xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(
217
+ -2) # ..., num_heads, head_dim
218
+ xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(
219
+ -2) # ..., num_heads, head_dim
220
+ return xq_out.type_as(xq), xk_out.type_as(xk)
221
+
222
+
223
+ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
224
+ """
225
+ From:
226
+ https://github.com/OpenGVLab/InternVideo/blob/421f6d2361fc8f61a3394244571f2601a4e99e29/InternVideo2/multi_modality/models/backbones/internvideo2/pos_embed.py#L86
227
+ embed_dim: output dimension for each position
228
+ pos: a list of positions to be encoded: size (M,)
229
+ out: (M, D)
230
+ """
231
+ assert embed_dim % 2 == 0
232
+ omega = np.arange(embed_dim // 2, dtype=np.float32)
233
+ omega /= embed_dim / 2.0
234
+ omega = 1.0 / 10000**omega # (D/2,)
235
+
236
+ pos = pos.reshape(-1) # (M,)
237
+ out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
238
+
239
+ emb_sin = np.sin(out) # (M, D/2)
240
+ emb_cos = np.cos(out) # (M, D/2)
241
+
242
+ emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
243
+ return emb
244
+
245
+
246
+ def get_1d_sincos_pos_embed(embed_dim, t_size, cls_token=False):
247
+ """
248
+ t_size: int of the temporal size
249
+ return:
250
+ pos_embed: [t_size, embed_dim] or [1+t_size, embed_dim] (w/ or w/o cls_token)
251
+ """
252
+ grid_t = np.arange(t_size, dtype=np.float32)
253
+ pos_embed = get_1d_sincos_pos_embed_from_grid(embed_dim, grid_t)
254
+ if cls_token:
255
+ pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed],
256
+ axis=0)
257
+ return pos_embed
258
+
259
+
260
+ class Learnable2DInterpPosEmbDivided_fixed(nn.Module):
261
+
262
+ def __init__(self,
263
+ height: int,
264
+ width: int,
265
+ num_frames: int,
266
+ dim: int,
267
+ interpolation_mode: str = 'bicubic') -> None:
268
+ super().__init__()
269
+ self.height = height
270
+ self.width = width
271
+ self.num_frames = num_frames
272
+ self.dim = dim
273
+ self.interpolation_mode = interpolation_mode
274
+ self.weight = nn.Parameter(torch.empty(height, width, dim))
275
+ self.register_buffer('time_weight',
276
+ torch.from_numpy(
277
+ get_1d_sincos_pos_embed(
278
+ self.dim,
279
+ self.num_frames)).float().unsqueeze(1),
280
+ persistent=False)
281
+
282
+ self.reset_parameters()
283
+
284
+ def reset_parameters(self):
285
+ nn.init.normal_(self.weight)
286
+
287
+ def forward(self, x: torch.Tensor,
288
+ grid_thws: torch.Tensor) -> torch.Tensor:
289
+ pos_embs = []
290
+ for t, h, w in grid_thws.tolist():
291
+ assert t <= self.num_frames, f't:{t} > self.num_frames:{self.num_frames}'
292
+ if (h, w) == self.weight.shape[:-1]:
293
+ pos_emb_2d = self.weight.flatten(end_dim=1)
294
+ else:
295
+ pos_emb_2d = get_rope_shape(
296
+ self.weight,
297
+ interpolation_mode=self.interpolation_mode,
298
+ shape=(h, w),
299
+ )
300
+
301
+ if t == 1:
302
+ pos_emb_3d = pos_emb_2d
303
+ else:
304
+ pos_emb_3d = pos_emb_2d.unsqueeze(0).repeat(
305
+ t, 1, 1) + self.time_weight[0:t]
306
+
307
+ pos_embs.append(pos_emb_3d.reshape(-1, pos_emb_3d.shape[-1]))
308
+
309
+ out = x + torch.cat(pos_embs)
310
+ return out
311
+
312
+
313
+ class MoonVision3dPatchEmbed(nn.Module):
314
+
315
+ def __init__(self,
316
+ out_dim: int,
317
+ in_dim: int = 3,
318
+ patch_size: int | tuple[int, int] = (14, 14),
319
+ pos_emb_height: int = 14,
320
+ pos_emb_width: int = 14,
321
+ pos_emb_time: int = 4,
322
+ pos_emb_type: str = 'divided_fixed',
323
+ patch_embed_proj_bias: bool = True,
324
+ pos_emb_interpolation_mode: str = 'bicubic'):
325
+ super().__init__()
326
+ assert isinstance(
327
+ patch_size,
328
+ int | Sequence), f'Invalid patch_size type: {type(patch_size)}'
329
+ if isinstance(patch_size, int):
330
+ patch_size = (patch_size, patch_size)
331
+ assert (len(patch_size) == 2
332
+ ), f'Expected patch_size to be a tuple of 2, got {patch_size}'
333
+ self.patch_size = patch_size
334
+
335
+ self.proj = nn.Conv2d(in_dim,
336
+ out_dim,
337
+ kernel_size=patch_size,
338
+ stride=patch_size,
339
+ bias=patch_embed_proj_bias)
340
+
341
+ if pos_emb_type == 'divided_fixed':
342
+ self.pos_emb = Learnable2DInterpPosEmbDivided_fixed(
343
+ height=pos_emb_height,
344
+ width=pos_emb_width,
345
+ num_frames=pos_emb_time,
346
+ dim=out_dim,
347
+ interpolation_mode=pos_emb_interpolation_mode)
348
+ else:
349
+ raise NotImplementedError(
350
+ f'Not support pos_emb_type: {pos_emb_type}')
351
+
352
+ def forward(self, x: torch.Tensor,
353
+ grid_thws: torch.Tensor) -> torch.Tensor:
354
+ """
355
+ Args:
356
+ x (L, Channels): input tensor
357
+ grid_hws (N, 3): temporal, height and width
358
+
359
+ Returns:
360
+ (L, Cout) tensor
361
+ """
362
+ x = self.proj(x).view(x.size(0), -1)
363
+ # apply positional embedding
364
+ x = self.pos_emb(x, grid_thws)
365
+ return x
366
+
367
+
368
+ class Rope2DPosEmbRepeated(nn.Module):
369
+ """2D rotary position embedding with multi-resolution support.
370
+
371
+ This class is intended to be used in the following way:
372
+ 1. Before training, create an instance of Rope2DPosEmb. This instance will hold the precomputed cis.
373
+ 2. Before each forward pass, call `get_freqs_cis_by_*` to get the `freqs_cis` tensor for this iteration.
374
+ 3. During the forward pass, pass the `freqs_cis` tensor to each attention layer, and call `apply` just before each attention operation.
375
+ The rope is shared across all attention layers and all heads.
376
+
377
+ Refs:
378
+ - RoFormer: https://arxiv.org/abs/2104.09864
379
+ - VisionLLaMA: https://arxiv.org/abs/2403.00522
380
+ - https://github.com/Meituan-AutoML/VisionLLaMA/blob/main/dit/models.py
381
+
382
+ Args:
383
+ dim (int): usually the multi-head attention dimension, should be divisible by 4 (TODO: relax this constraint if needed)
384
+ max_height (int): the maximum height of the 2D grid
385
+ max_width (int): the maximum width of the 2D grid
386
+ theta_base (float): the base of the theta
387
+ device (str): the device to store the precomputed cis
388
+ """
389
+
390
+ def __init__(self,
391
+ dim: int,
392
+ max_height: int,
393
+ max_width: int,
394
+ theta_base=10000):
395
+ super().__init__()
396
+ self.dim = dim
397
+ assert self.dim % 4 == 0, 'dim must be divisible by 4'
398
+ self.max_height = max_height
399
+ self.max_width = max_width
400
+ self.theta_base = theta_base
401
+
402
+ def extra_repr(self):
403
+ return f'dim={self.dim}, max_height={self.max_height}, max_width={self.max_width}, theta_base={self.theta_base}'
404
+
405
+ def _precompute_freqs_cis(self, device: torch.device) -> torch.Tensor:
406
+ """Calculate the cis(freqs) for each position in the 2D grid.
407
+
408
+ Return: complex tensor of shape (max_height, max_width, dim//2) and value:
409
+ height axis: ret[h, w, 2*i] = cis(h * theta_base**(-4*i/dim))
410
+ weight axis: ret[h, w, 2*i+1] = cis(w * theta_base**(-4*i/dim)) with (i in [0, dim//4))
411
+ note: `cis` is a mathematical notation defined by cis x = cos x + i sin x,
412
+ """
413
+ N = self.max_height * self.max_width
414
+ flat_pos = torch.arange(0, N).float().to(device)
415
+ x_pos = flat_pos % self.max_width
416
+ y_pos = flat_pos // self.max_width
417
+ dim_range = (torch.arange(0, self.dim,
418
+ 4)[:(self.dim // 4)].float().to(device)
419
+ ) # C/4
420
+ freqs = 1.0 / (self.theta_base**(dim_range / self.dim))
421
+ x_freqs = torch.outer(x_pos, freqs).float() # N, C/4
422
+ y_freqs = torch.outer(y_pos, freqs).float() # N, C/4
423
+ x_cis = torch.polar(torch.ones_like(x_freqs), x_freqs) # N, C/4
424
+ y_cis = torch.polar(torch.ones_like(y_freqs), y_freqs) # N, C/4
425
+ # N, C/4, 2
426
+ freqs_cis = torch.cat(
427
+ [x_cis.unsqueeze(dim=-1),
428
+ y_cis.unsqueeze(dim=-1)], dim=-1)
429
+ # max_height, max_width, C/2
430
+ freqs_cis = freqs_cis.reshape(self.max_height, self.max_width, -1)
431
+ return freqs_cis
432
+
433
+ def get_freqs_cis(self, grid_thws: torch.Tensor,
434
+ device: torch.device) -> torch.Tensor:
435
+ """
436
+ Args:
437
+ grid_thws (torch.Tensor): grid time, height and width
438
+
439
+ Returns:
440
+ freqs_cis: tensor of shape (sum(t * height * width), dim//2)
441
+ """
442
+ if not hasattr(self, 'freqs_cis'):
443
+ self.register_buffer('freqs_cis',
444
+ self._precompute_freqs_cis(device),
445
+ persistent=False)
446
+
447
+ shapes = grid_thws.tolist()
448
+ assert all(1 <= h <= self.max_height and 1 <= w <= self.max_width
449
+ for t, h, w in shapes), (
450
+ shapes,
451
+ self.max_height,
452
+ self.max_width,
453
+ )
454
+ freqs_cis = torch.cat(
455
+ [
456
+ self.freqs_cis[:h, :w].reshape(-1, self.dim // 2).repeat(t, 1)
457
+ for t, h, w in shapes
458
+ ],
459
+ dim=0,
460
+ )
461
+ return freqs_cis
462
+
463
+
464
+ class MLP2(nn.Module):
465
+ """
466
+ Args:
467
+ dims: [in_dim, hidden_dim, out_dim]
468
+ bias: whether to use bias in linear layer.
469
+ """
470
+
471
+ def __init__(self, dims: list[int], activation, bias=True):
472
+ super().__init__()
473
+ assert len(dims) == 3
474
+ self.fc0 = nn.Linear(dims[0], dims[1], bias=bias)
475
+ self.fc1 = nn.Linear(dims[1], dims[2], bias=bias)
476
+ self.activation = activation
477
+ for m in [self.fc0, self.fc1]:
478
+ nn.init.trunc_normal_(m.weight, std=math.sqrt(2 / m.in_features))
479
+ if m.bias is not None:
480
+ nn.init.zeros_(m.bias)
481
+
482
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
483
+ x = self.fc0(x)
484
+ x = self.activation(x)
485
+ return self.fc1(x)
486
+
487
+
488
+ class MoonViTEncoderLayer(nn.Module):
489
+
490
+ def __init__(
491
+ self,
492
+ num_heads: int,
493
+ hidden_dim: int,
494
+ mlp_dim: int,
495
+ qkv_hidden_size: int | None = None,
496
+ norm_type: str = 'layernorm',
497
+ mlp_type: str = 'mlp2',
498
+ *,
499
+ attn_implementation: str = 'flash_attention_2',
500
+ activation=F.gelu,
501
+ attn_bias: bool = False,
502
+ linear_bias: bool = True,
503
+ use_deterministic_attn: bool = False,
504
+ ):
505
+ super().__init__()
506
+ self.num_heads = num_heads
507
+ self.hidden_dim = hidden_dim
508
+ self.qkv_hidden_size = hidden_dim if qkv_hidden_size is None else qkv_hidden_size
509
+ self.hidden_size_per_attention_head = self.qkv_hidden_size // self.num_heads
510
+ self.attn_implementation = attn_implementation
511
+ self.use_deterministic_attn = use_deterministic_attn
512
+
513
+ if norm_type == "layernorm":
514
+ self.norm0 = nn.LayerNorm(hidden_dim)
515
+ self.norm1 = nn.LayerNorm(hidden_dim)
516
+ elif norm_type == "rmsnorm":
517
+ self.norm0 = nn.RMSNorm(hidden_dim, eps=1e-6)
518
+ self.norm1 = nn.RMSNorm(hidden_dim, eps=1e-6)
519
+ else:
520
+ raise NotImplementedError(f"Not support norm_type: {norm_type}")
521
+
522
+ if mlp_type == "mlp2":
523
+ self.mlp = MLP2([hidden_dim, mlp_dim, hidden_dim],
524
+ activation,
525
+ bias=linear_bias)
526
+ else:
527
+ raise NotImplementedError(f"Not support mlp_type: {mlp_type}")
528
+
529
+ self.wqkv = nn.Linear(hidden_dim,
530
+ self.qkv_hidden_size * 3,
531
+ bias=attn_bias)
532
+ self.wo = nn.Linear(self.qkv_hidden_size, hidden_dim, bias=attn_bias)
533
+
534
+ def attention_qkvpacked(
535
+ self,
536
+ x: torch.Tensor,
537
+ cu_seqlens: torch.Tensor,
538
+ max_seqlen: torch.Tensor,
539
+ rope_freqs_cis: torch.Tensor | None = None,
540
+ ):
541
+ """
542
+ Args:
543
+ x (torch.Tensor): (batch_size, seqlen, hidden_dim)
544
+ cu_seqlens (torch.Tensor):
545
+ """
546
+ xqkv = self.wqkv(x)
547
+
548
+ qkv_shape = xqkv.size()[:-1] + (
549
+ 3,
550
+ self.num_heads,
551
+ self.hidden_size_per_attention_head,
552
+ )
553
+ # xqkv: (batch_size, seqlen, 3, nheads, headdim)
554
+ xqkv = xqkv.view(*qkv_shape)
555
+ xq, xk, xv = torch.unbind(xqkv, dim=-3)
556
+
557
+ xq, xk = apply_rope(xq, xk, rope_freqs_cis)
558
+
559
+ attn_func = VL_VISION_ATTENTION_FUNCTIONS[self.attn_implementation]
560
+ attn_out = attn_func(xq,
561
+ xk,
562
+ xv,
563
+ q_cu_seqlens=cu_seqlens,
564
+ k_cu_seqlens=cu_seqlens,
565
+ max_seqlen_k=max_seqlen,
566
+ max_seqlen_q=max_seqlen,
567
+ deterministic=self.use_deterministic_attn)
568
+
569
+ attn_out = self.wo(attn_out)
570
+ return attn_out
571
+
572
+ def forward(
573
+ self,
574
+ hidden_states: torch.Tensor,
575
+ cu_seqlens: torch.Tensor,
576
+ max_seqlen: int,
577
+ rope_freqs_cis: torch.Tensor | None = None,
578
+ ):
579
+ residual = hidden_states
580
+ hidden_states = self.norm0(hidden_states)
581
+
582
+ hidden_states = self.attention_qkvpacked(hidden_states, cu_seqlens,
583
+ max_seqlen, rope_freqs_cis)
584
+ hidden_states = residual + hidden_states
585
+
586
+ residual = hidden_states
587
+ hidden_states = self.norm1(hidden_states)
588
+ hidden_states = self.mlp(hidden_states)
589
+ hidden_states = residual + hidden_states
590
+
591
+ return hidden_states
592
+
593
+
594
+ class MoonViT3dEncoder(nn.Module):
595
+
596
+ def __init__(self,
597
+ hidden_dim: int,
598
+ num_layers: int,
599
+ block_cfg: dict,
600
+ use_deterministic_attn: bool = False) -> None:
601
+ super().__init__()
602
+ self.use_deterministic_attn = use_deterministic_attn
603
+
604
+ qkv_hidden_size = block_cfg['hidden_dim'] if block_cfg.get(
605
+ 'qkv_hidden_size') is None else block_cfg['qkv_hidden_size']
606
+ self.rope_2d = Rope2DPosEmbRepeated(
607
+ qkv_hidden_size // block_cfg['num_heads'], 512, 512)
608
+ self.blocks = nn.ModuleList([
609
+ MoonViTEncoderLayer(
610
+ **block_cfg,
611
+ use_deterministic_attn=self.use_deterministic_attn)
612
+ for _ in range(num_layers)
613
+ ])
614
+ norm_type = block_cfg.get('norm_type', 'layernorm')
615
+ if norm_type == "layernorm":
616
+ self.final_layernorm = nn.LayerNorm(hidden_dim)
617
+ elif norm_type == "rmsnorm":
618
+ self.final_layernorm = nn.RMSNorm(hidden_dim, eps=1e-6)
619
+ else:
620
+ raise NotImplementedError(f"Not support norm_type: {norm_type}")
621
+
622
+ def forward(
623
+ self,
624
+ hidden_states: torch.Tensor,
625
+ grid_thws: torch.Tensor,
626
+ ) -> torch.Tensor:
627
+ rope_freqs_cis = self.rope_2d.get_freqs_cis(
628
+ grid_thws=grid_thws, device=hidden_states.device)
629
+
630
+ lengths = torch.cat((
631
+ torch.zeros(1, dtype=grid_thws.dtype, device=grid_thws.device),
632
+ grid_thws[:, 0] * grid_thws[:, 1] * grid_thws[:, 2],
633
+ ))
634
+
635
+ max_seqlen = lengths.max()
636
+ cu_seqlens = lengths.to(hidden_states.device).cumsum(dim=0,
637
+ dtype=torch.int32)
638
+ for block in self.blocks:
639
+ hidden_states = block(hidden_states,
640
+ cu_seqlens,
641
+ max_seqlen,
642
+ rope_freqs_cis=rope_freqs_cis)
643
+
644
+ hidden_states = self.final_layernorm(hidden_states)
645
+ return hidden_states
646
+
647
+
648
+ def tpool_patch_merger(
649
+ x: torch.Tensor,
650
+ grid_thws: torch.Tensor,
651
+ merge_kernel_size: tuple[int, int] = (2, 2),
652
+ ) -> list[torch.Tensor]:
653
+ d_model = x.size(-1)
654
+
655
+ outputs = []
656
+ pre_sum = 0
657
+ for t, h, w in grid_thws.tolist():
658
+ # Get the current sequence
659
+ seq = x[pre_sum:pre_sum + t * h * w]
660
+ # Reshape along self.merge_kernel_size and concat to the last dimension
661
+ kernel_height, kernel_width = merge_kernel_size
662
+ new_height, new_width = h // kernel_height, w // kernel_width
663
+ reshaped_seq = seq.view(t, new_height, kernel_height, new_width,
664
+ kernel_width, d_model)
665
+ reshaped_seq = reshaped_seq.permute(0, 1,
666
+ 3, 2, 4, 5).contiguous().mean(
667
+ dim=0) # temporal pooling
668
+ padded_seq = reshaped_seq.view(new_height * new_width,
669
+ kernel_height * kernel_width, -1)
670
+ outputs.append(padded_seq)
671
+ pre_sum += t * h * w
672
+
673
+ return outputs
674
+
675
+
676
+ class MoonViT3dPretrainedModel(PreTrainedModel):
677
+ config_class = None
678
+ model_type = 'moonvit3d'
679
+ _no_split_modules = ['MoonViTEncoderLayer']
680
+ _supports_flash_attn = True
681
+ _supports_flash_attn_2 = True
682
+ _supports_sdpa = True
683
+
684
+ def __init__(self, config, *inputs, **kwargs):
685
+ super().__init__(config, *inputs, **kwargs)
686
+ config = deepcopy(config)
687
+ self.merge_kernel_size = config.merge_kernel_size
688
+ self.patch_size = config.patch_size
689
+ self.merge_type = config.merge_type
690
+
691
+ self.patch_embed = MoonVision3dPatchEmbed(
692
+ out_dim=config.hidden_size,
693
+ patch_size=config.patch_size,
694
+ pos_emb_height=config.init_pos_emb_height,
695
+ pos_emb_width=config.init_pos_emb_width,
696
+ pos_emb_time=config.init_pos_emb_time,
697
+ pos_emb_type=config.pos_emb_type,
698
+ patch_embed_proj_bias=getattr(config, 'patch_embed_proj_bias',
699
+ True),
700
+ pos_emb_interpolation_mode=getattr(
701
+ config, 'pos_emb_interpolation_mode', 'bicubic'),
702
+ )
703
+
704
+ self.encoder = MoonViT3dEncoder(
705
+ hidden_dim=config.hidden_size,
706
+ num_layers=config.num_hidden_layers,
707
+ block_cfg={
708
+ 'num_heads': config.num_attention_heads,
709
+ 'hidden_dim': config.hidden_size,
710
+ 'qkv_hidden_size': getattr(config, 'qkv_hidden_size', None),
711
+ 'mlp_dim': config.intermediate_size,
712
+ 'norm_type': getattr(config, 'norm_type', 'layernorm'),
713
+ 'mlp_type': getattr(config, 'mlp_type', 'mlp2'),
714
+ 'activation': PytorchGELUTanh(),
715
+ 'attn_bias': getattr(config, 'attn_bias', True),
716
+ 'linear_bias': getattr(config, 'linear_bias', True),
717
+ 'attn_implementation': config._attn_implementation,
718
+ },
719
+ use_deterministic_attn=getattr(self, 'use_deterministic_attn',
720
+ False))
721
+
722
+ def forward(self, pixel_values: torch.Tensor,
723
+ grid_thws: torch.Tensor) -> torch.Tensor:
724
+ """
725
+ Args:
726
+ pixel_values (torch.Tensor): The input pixel values.
727
+ grid_thws (torch.Tensor): Temporal, height and width.
728
+
729
+ Returns:
730
+ torch.Tensor: The output tokens.
731
+ """
732
+ # grid_thws = grid_thws.to('cpu')
733
+ assert grid_thws.ndim == 2, f'grid_thws should be 2D, got {grid_thws.ndim}'
734
+ assert grid_thws.size(1) == 3, f'No support for thw: {grid_thws}'
735
+ hidden_states = self.patch_embed(pixel_values, grid_thws)
736
+ hidden_states = self.encoder(hidden_states, grid_thws)
737
+ if self.merge_type == 'sd2_tpool': # spatial downsampling 2x with temporal pooling all
738
+ hidden_states = tpool_patch_merger(
739
+ hidden_states,
740
+ grid_thws,
741
+ merge_kernel_size=self.merge_kernel_size)
742
+ else:
743
+ raise NotImplementedError(f'Not support {self.merge_type}')
744
+
745
+ return hidden_states
746
+
747
+
748
+ # ============================================================================
749
+ # MM Projector Helper Classes (from mm_projector/modeling_mm_projectors.py)
750
+ # ============================================================================
751
+
752
+
753
+ class IdentityMap(nn.Module):
754
+
755
+ def __init__(self):
756
+ super().__init__()
757
+
758
+ def forward(self, x, *args, **kwargs):
759
+ return x
760
+
761
+
762
+ class MLP(nn.Module):
763
+
764
+ def __init__(self, config):
765
+ super().__init__()
766
+ # TODO, use faster LayerNorm
767
+ self.pre_norm = nn.LayerNorm(config.mm_hidden_size)
768
+ self.proj = nn.Sequential(
769
+ nn.Linear(config.mm_hidden_size, config.hidden_size), nn.GELU(),
770
+ nn.Linear(config.hidden_size, config.hidden_size))
771
+
772
+ def forward(self, x, *args, **kwargs):
773
+ assert isinstance(x,
774
+ list | tuple), f'x is not a list or tuple: {type(x)}'
775
+ lengths = [item.shape[0] for item in x]
776
+ x = torch.cat(x, dim=0)
777
+ x = self.pre_norm(x)
778
+ x = self.proj(x)
779
+ x = torch.split(x, lengths, dim=0)
780
+
781
+ return x
782
+
783
+
784
+ class PatchMergerMLP(nn.Module):
785
+
786
+ def __init__(self, config):
787
+ super().__init__()
788
+ eps = config.projector_ln_eps
789
+ self.hidden_size = config.mm_hidden_size * (
790
+ config.merge_kernel_size[0] * config.merge_kernel_size[1])
791
+ self.pre_norm = nn.LayerNorm(config.mm_hidden_size, eps=eps)
792
+ self.proj = nn.Sequential(
793
+ nn.Linear(self.hidden_size, self.hidden_size),
794
+ nn.GELU(),
795
+ nn.Linear(self.hidden_size, config.hidden_size),
796
+ )
797
+
798
+ def forward(self, x, *args, **kwargs):
799
+ if isinstance(x, list) or isinstance(x, tuple):
800
+ x = [
801
+ self.proj(self.pre_norm(item).view(item.shape[0], -1))
802
+ for item in x
803
+ ]
804
+ else:
805
+ # B, N, N_k, C = x.shape
806
+ B = x.shape[0]
807
+ x = self.proj(self.pre_norm(x).view(B, -1, self.hidden_size))
808
+ return x
809
+
810
+
811
+ class PatchMergerMLPV2(nn.Module):
812
+
813
+ def __init__(self, config):
814
+ super().__init__()
815
+ eps = config.projector_ln_eps
816
+ self.hidden_size = config.mm_hidden_size * (
817
+ config.merge_kernel_size[0] * config.merge_kernel_size[1])
818
+ self.proj = nn.Sequential(
819
+ nn.Linear(self.hidden_size, self.hidden_size, bias=False),
820
+ nn.GELU(),
821
+ nn.Linear(self.hidden_size, config.hidden_size, bias=False),
822
+ )
823
+ self.post_norm = nn.RMSNorm(config.hidden_size, eps=eps)
824
+ for m in self.proj.modules():
825
+ if isinstance(m, nn.Linear):
826
+ nn.init.trunc_normal_(m.weight,
827
+ std=math.sqrt(2 / m.in_features))
828
+ if m.bias is not None:
829
+ nn.init.zeros_(m.bias)
830
+
831
+ def forward(self, x, *args, **kwargs):
832
+ if isinstance(x, list) or isinstance(x, tuple):
833
+ lengths = [item.shape[0] for item in x]
834
+ x = torch.concat([item.view(item.shape[0], -1) for item in x],
835
+ dim=0)
836
+ x = self.post_norm(self.proj(x))
837
+ x = torch.split(x, lengths, dim=0)
838
+ else:
839
+ # B, N, N_k, C = x.shape
840
+ B = x.shape[0]
841
+ x = self.proj(x.view(B, -1, self.hidden_size))
842
+ x = self.post_norm(x)
843
+ return x
844
+
845
+
846
+ class GroundAnythingBackbonePreTrainedModel(PreTrainedModel):
847
+ config_class = GroundAnythingBackboneConfig
848
+ base_model_prefix = "model"
849
+ _no_split_modules = [
850
+ "MoonViT3dPretrainedModel",
851
+ "MoonViTEncoderLayer",
852
+ "KimiDecoderLayer",
853
+ "PatchMergerMLP",
854
+ "PatchMergerMLPV2",
855
+ ]
856
+ _skip_keys_device_placement = "past_key_values"
857
+ _supports_flash_attn_2 = True
858
+ _supports_sdpa = False
859
+
860
+ def _init_weights(self, module):
861
+ # important: this ported version of Llava isn't meant for training from scratch - only
862
+ # inference and fine-tuning - so the proper init weights code has been removed - the original codebase
863
+ # https://github.com/haotian-liu/LLaVA/tree/main/llava should serve for that purpose
864
+ std = (self.config.initializer_range if hasattr(
865
+ self.config, "initializer_range") else
866
+ self.config.text_config.initializer_range)
867
+
868
+ if hasattr(module, "class_embedding"):
869
+ module.class_embedding.data.normal_(mean=0.0, std=std)
870
+
871
+ if isinstance(module, (nn.Linear, nn.Conv2d)):
872
+ module.weight.data.normal_(mean=0.0, std=std)
873
+ if module.bias is not None:
874
+ module.bias.data.zero_()
875
+ elif isinstance(module, nn.Embedding):
876
+ module.weight.data.normal_(mean=0.0, std=std)
877
+ if module.padding_idx is not None:
878
+ module.weight.data[module.padding_idx].zero_()
879
+
880
+
881
+ class VisionTowerConfig(PretrainedConfig):
882
+ model_type = 'moonvit3d'
883
+
884
+ def __init__(self, config: GroundAnythingBackboneConfig, **kwargs):
885
+ super().__init__(**kwargs)
886
+ self.patch_size = config.patch_size
887
+ self.init_pos_emb_height = config.init_pos_emb_height
888
+ self.init_pos_emb_width = config.init_pos_emb_width
889
+ self.init_pos_emb_time = config.init_pos_emb_time
890
+ self.pos_emb_type = config.pos_emb_type
891
+ self.num_attention_heads = config.vt_num_attention_heads
892
+ self.num_hidden_layers = config.vt_num_hidden_layers
893
+ self.hidden_size = config.vt_hidden_size
894
+ self.intermediate_size = config.vt_intermediate_size
895
+ self.merge_kernel_size = config.merge_kernel_size
896
+ self.merge_type = config.merge_type
897
+ self._attn_implementation = config._attn_implementation
898
+ self.qkv_hidden_size = getattr(config, 'qkv_hidden_size', None)
899
+ self.norm_type = getattr(config, 'norm_type', 'layernorm')
900
+ self.attn_bias = getattr(config, 'attn_bias', True)
901
+ self.patch_embed_proj_bias = getattr(config, 'patch_embed_proj_bias',
902
+ True)
903
+ self.mlp_type = getattr(config, 'mlp_type', 'mlp2')
904
+ self.linear_bias = getattr(config, 'linear_bias', True)
905
+ self.pos_emb_interpolation_mode = getattr(
906
+ config, 'pos_emb_interpolation_mode', 'bilinear')
907
+
908
+
909
+ class ProjectorConfig:
910
+
911
+ def __init__(self, config: GroundAnythingBackboneConfig):
912
+ self.mm_projector_type = config.mm_projector_type
913
+ self.mm_hidden_size = config.mm_hidden_size
914
+ self.hidden_size = config.text_hidden_size
915
+ self.merge_kernel_size = config.merge_kernel_size
916
+ self.projector_hidden_act = config.projector_hidden_act
917
+ self.projector_ln_eps = config.projector_ln_eps
918
+
919
+
920
+ # ref https://github.com/huggingface/transformers/blob/78b2929c0554b79e0489b451ce4ece14d265ead2/src/transformers/models/llava/modeling_llava.py#L240
921
+ class GroundAnythingBackboneForConditionalGeneration(GroundAnythingBackbonePreTrainedModel):
922
+
923
+ @classmethod
924
+ def _supports_default_dynamic_cache(cls) -> bool:
925
+ return False
926
+
927
+ def __init__(self, config: GroundAnythingBackboneConfig):
928
+ super().__init__(config)
929
+
930
+ vt_config = VisionTowerConfig(config.vision_config)
931
+ self.vision_tower = MoonViT3dPretrainedModel(vt_config)
932
+
933
+ proj_config = ProjectorConfig(config.vision_config)
934
+ if proj_config.mm_projector_type == 'identity':
935
+ self.mm_projector = IdentityMap()
936
+ elif proj_config.mm_projector_type == 'mlp':
937
+ self.mm_projector = MLP(proj_config)
938
+ elif proj_config.mm_projector_type == 'patchmerger':
939
+ self.mm_projector = PatchMergerMLP(proj_config)
940
+ elif proj_config.mm_projector_type == 'patchmergerv2':
941
+ self.mm_projector = PatchMergerMLPV2(proj_config)
942
+ else:
943
+ raise ValueError(
944
+ f"Unsupported mm_projector_type: {proj_config.mm_projector_type}"
945
+ )
946
+
947
+ self.language_model = GroundAnythingBackboneLinearForCausalLM(config.text_config)
948
+ self.post_init()
949
+
950
+ if hasattr(self.language_model, 'dtype'):
951
+ target_dtype = self.language_model.dtype
952
+ self.vision_tower = self.vision_tower.to(dtype=target_dtype)
953
+ self.mm_projector = self.mm_projector.to(dtype=target_dtype)
954
+
955
+ def get_input_embeddings(self):
956
+ return self.language_model.get_input_embeddings()
957
+
958
+ def set_input_embeddings(self, value):
959
+ self.language_model.set_input_embeddings(value)
960
+
961
+ def get_output_embeddings(self):
962
+ return self.language_model.get_output_embeddings()
963
+
964
+ def set_output_embeddings(self, new_embeddings):
965
+ self.language_model.set_output_embeddings(new_embeddings)
966
+
967
+ def set_decoder(self, decoder):
968
+ self.language_model.set_decoder(decoder)
969
+
970
+ def get_decoder(self):
971
+ return self.language_model.get_decoder()
972
+
973
+ def tie_weights(self):
974
+ return self.language_model.tie_weights()
975
+
976
+ def resize_token_embeddings(self,
977
+ new_num_tokens: int | None = None,
978
+ pad_to_multiple_of=None) -> nn.Embedding:
979
+ model_embeds = self.language_model.resize_token_embeddings(
980
+ new_num_tokens, pad_to_multiple_of)
981
+ # update vocab size
982
+ self.config.text_config.vocab_size = model_embeds.num_embeddings
983
+ self.vocab_size = model_embeds.num_embeddings
984
+ return model_embeds
985
+
986
+ def _merge_input_ids_with_image_features(
987
+ self,
988
+ image_features: list[torch.Tensor],
989
+ inputs_embeds: torch.Tensor,
990
+ input_ids: torch.Tensor,
991
+ attention_mask: torch.Tensor,
992
+ labels: torch.Tensor | None = None,
993
+ ):
994
+ """
995
+ Args:
996
+ image_features (:obj:`torch.Tensor` of shape :obj:`(num_image_tokens, embed_dim)`):
997
+ The image features to merge with the input embeddings.
998
+ inputs_embeds (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length, embed_dim)`):
999
+ The input embeddings.
1000
+ input_ids (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`):
1001
+ The input ids.
1002
+ attention_mask (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`):
1003
+ The attention mask.
1004
+ labels (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`, *optional*):
1005
+ The labels.
1006
+ """
1007
+ _, embed_dim = image_features[0].shape
1008
+ feature_lengths = [x.shape[0] for x in image_features]
1009
+ image_features = torch.cat(image_features, dim=0)
1010
+
1011
+ image_token_index: int = self.config.media_placeholder_token_id
1012
+ pad_token_id: int = self.config.pad_token_id
1013
+ ignore_index: int = self.config.ignore_index
1014
+
1015
+ batch_size, sequence_length = input_ids.shape
1016
+ left_padding = not torch.sum(
1017
+ input_ids[:, -1] == torch.tensor(pad_token_id))
1018
+
1019
+ # 1. Create a mask to know where special image tokens are
1020
+ _token_occupation_table = torch.ones_like(input_ids.flatten())
1021
+ _token_occupation_table[input_ids.flatten() ==
1022
+ image_token_index] = torch.tensor(
1023
+ feature_lengths,
1024
+ dtype=torch.long,
1025
+ device=input_ids.device)
1026
+ _token_occupation_table = _token_occupation_table.reshape(
1027
+ input_ids.shape)
1028
+
1029
+ max_embed_dim = _token_occupation_table.sum(-1).max().item()
1030
+ assert (
1031
+ max_embed_dim >= sequence_length
1032
+ ), f"The maximum embedding dimension ({max_embed_dim}) is less than the sequence length ({sequence_length})"
1033
+ batch_indices, non_image_indices = torch.where(
1034
+ input_ids != image_token_index)
1035
+
1036
+ # 2. Compute the positions where text should be written
1037
+ # Calculate new positions for text tokens in merged image-text sequence.
1038
+ new_token_positions = torch.cumsum(_token_occupation_table, -1) - 1
1039
+ nb_image_pad = max_embed_dim - 1 - new_token_positions[:, -1]
1040
+ if left_padding:
1041
+ new_token_positions += nb_image_pad[:,
1042
+ None] # offset for left padding
1043
+ text_to_overwrite = new_token_positions[batch_indices,
1044
+ non_image_indices]
1045
+
1046
+ # 3. Create the full embedding, already padded to the maximum position
1047
+ final_embedding = torch.zeros(
1048
+ batch_size,
1049
+ max_embed_dim,
1050
+ embed_dim,
1051
+ dtype=inputs_embeds.dtype,
1052
+ device=inputs_embeds.device,
1053
+ )
1054
+ final_attention_mask = torch.zeros(batch_size,
1055
+ max_embed_dim,
1056
+ dtype=attention_mask.dtype,
1057
+ device=inputs_embeds.device)
1058
+ if labels is not None:
1059
+ final_labels = torch.full(
1060
+ (batch_size, max_embed_dim),
1061
+ ignore_index,
1062
+ dtype=input_ids.dtype,
1063
+ device=input_ids.device,
1064
+ )
1065
+ # In case the Vision model or the Language model has been offloaded to CPU, we need to manually
1066
+ # set the corresponding tensors into their correct target device.
1067
+ target_device = inputs_embeds.device
1068
+ batch_indices, non_image_indices, text_to_overwrite = (
1069
+ batch_indices.to(target_device),
1070
+ non_image_indices.to(target_device),
1071
+ text_to_overwrite.to(target_device),
1072
+ )
1073
+ attention_mask = attention_mask.to(target_device)
1074
+
1075
+ # 4. Fill the embeddings based on the mask.
1076
+ final_embedding[batch_indices,
1077
+ text_to_overwrite] = inputs_embeds[batch_indices,
1078
+ non_image_indices]
1079
+ final_attention_mask[batch_indices,
1080
+ text_to_overwrite] = attention_mask[
1081
+ batch_indices, non_image_indices]
1082
+ if labels is not None:
1083
+ final_labels[batch_indices,
1084
+ text_to_overwrite] = labels[batch_indices,
1085
+ non_image_indices]
1086
+
1087
+ # 5. Fill the embeddings corresponding to the images. Anything that is not `text_positions` needs filling (#29835)
1088
+ image_to_overwrite = torch.full((batch_size, max_embed_dim),
1089
+ True,
1090
+ dtype=torch.bool,
1091
+ device=inputs_embeds.device)
1092
+ image_to_overwrite[batch_indices, text_to_overwrite] = False
1093
+ image_to_overwrite &= image_to_overwrite.cumsum(
1094
+ -1) - 1 >= nb_image_pad[:, None].to(target_device)
1095
+
1096
+ if image_to_overwrite.sum() != image_features.shape[:-1].numel():
1097
+ raise ValueError(
1098
+ f"The input provided to the model are wrong. The number of image tokens is {image_to_overwrite.sum()} while"
1099
+ f" the number of image features given to the model is {image_features.shape[:-1].numel()}. "
1100
+ "This prevents correct indexing and breaks batch generation.")
1101
+
1102
+ final_embedding[image_to_overwrite] = (
1103
+ image_features.contiguous().reshape(-1,
1104
+ embed_dim).to(target_device))
1105
+ final_attention_mask |= image_to_overwrite
1106
+ position_ids = (final_attention_mask.cumsum(-1) - 1).masked_fill_(
1107
+ (final_attention_mask == 0), 1)
1108
+
1109
+ # 6. Mask out the embedding at padding positions, as we later use the past_key_value value to determine the non-attended tokens.
1110
+ batch_indices, pad_indices = torch.where(input_ids == pad_token_id)
1111
+ indices_to_mask = new_token_positions[batch_indices, pad_indices]
1112
+
1113
+ final_embedding[batch_indices, indices_to_mask] = 0
1114
+
1115
+ if labels is None:
1116
+ final_labels = None
1117
+
1118
+ return final_embedding, final_attention_mask, final_labels, position_ids
1119
+
1120
+ def _extract_image_features(self, pixel_values: torch.Tensor,
1121
+ grid_thws: torch.Tensor) -> list[torch.Tensor]:
1122
+ """
1123
+ Args:
1124
+ pixel_values (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, num_channels, height, width)`):
1125
+ The pixel values of the images processed by image processor.
1126
+ grid_thws (:obj:`torch.Tensor` of shape :obj:`(batch_size, 3)`):
1127
+ The grid, height, width of the images.
1128
+
1129
+ Returns:
1130
+ selected_image_feature (:obj:`torch.FloatTensor` of shape :obj:`(num_image_tokens, embed_dim)`):
1131
+ The selected image features to use as input to the projector head.
1132
+
1133
+ """
1134
+
1135
+ target_dtype = self.vision_tower.patch_embed.proj.weight.dtype
1136
+ pixel_values = pixel_values.to(target_dtype)
1137
+
1138
+ image_features = self.vision_tower(pixel_values, grid_thws)
1139
+ return image_features
1140
+
1141
+ def forward(
1142
+ self,
1143
+ input_ids: torch.LongTensor | None = None,
1144
+ pixel_values: torch.FloatTensor | list[torch.FloatTensor]
1145
+ | None = None,
1146
+ grid_thws: torch.Tensor | None = None,
1147
+ attention_mask: torch.Tensor | None = None,
1148
+ position_ids: torch.LongTensor | None = None,
1149
+ past_key_values: list[torch.FloatTensor] | None = None,
1150
+ inputs_embeds: torch.FloatTensor | None = None,
1151
+ labels: torch.LongTensor | None = None,
1152
+ use_cache: bool | None = None,
1153
+ output_attentions: bool | None = None,
1154
+ output_hidden_states: bool | None = None,
1155
+ return_dict: bool | None = None,
1156
+ ) -> tuple | LlavaCausalLMOutputWithPast:
1157
+ r"""
1158
+ Args:
1159
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1160
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1161
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1162
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1163
+
1164
+ ```"""
1165
+ assert self.vision_tower is not None, "vision_tower is not loaded"
1166
+ output_attentions = (output_attentions if output_attentions is not None
1167
+ else self.config.output_attentions)
1168
+ output_hidden_states = (output_hidden_states
1169
+ if output_hidden_states is not None else
1170
+ self.config.output_hidden_states)
1171
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1172
+
1173
+ if inputs_embeds is None:
1174
+ # 1. Extra the input embeddings
1175
+ inputs_embeds = self.get_input_embeddings()(input_ids)
1176
+
1177
+ # 2. Merge text and images
1178
+ if pixel_values is not None and len(
1179
+ pixel_values) > 0 and input_ids.shape[1] != 1:
1180
+ image_features = self._extract_image_features(
1181
+ pixel_values, grid_thws)
1182
+ if self.mm_projector:
1183
+ image_features = self.mm_projector(image_features)
1184
+
1185
+ inputs_embeds = inputs_embeds.to(
1186
+ image_features[0].dtype) # num_tokens, embed_dim
1187
+ inputs_embeds, attention_mask, labels, position_ids = (
1188
+ self._merge_input_ids_with_image_features(
1189
+ image_features,
1190
+ inputs_embeds,
1191
+ input_ids,
1192
+ attention_mask,
1193
+ labels,
1194
+ ))
1195
+
1196
+ # In case input_ids.shape[1] == 1 & pixel_values==None & past_key_values != None, we are in the case of
1197
+ # generation with cache
1198
+ elif (past_key_values is not None and pixel_values is not None
1199
+ and input_ids.shape[1] == 1):
1200
+ # Retrieve the first layer to inspect the logits and mask out the hidden states
1201
+ # that are set to 0
1202
+ first_layer_past_key_value = past_key_values[0][0][:, :, :, 0]
1203
+
1204
+ # Sum all dimensions of head_dim (-2) to avoid random errors such as: https://github.com/huggingface/transformers/pull/28032#issuecomment-1863691941
1205
+ batch_index, non_attended_tokens = torch.where(
1206
+ first_layer_past_key_value.float().sum(-2) == 0)
1207
+
1208
+ # Get the target length
1209
+ target_length = input_ids.shape[1]
1210
+ past_length = first_layer_past_key_value.shape[-1]
1211
+
1212
+ extended_attention_mask = torch.ones(
1213
+ (attention_mask.shape[0], past_length),
1214
+ dtype=attention_mask.dtype,
1215
+ device=attention_mask.device,
1216
+ )
1217
+
1218
+ # Filter out only the tokens that can be un-attended, this can happen
1219
+ # if one uses Llava + Fused modules where the cache on the
1220
+ # first iteration is already big enough, or if one passes custom cache
1221
+ valid_indices = non_attended_tokens < extended_attention_mask.size(
1222
+ -1)
1223
+ new_batch_index = batch_index[valid_indices]
1224
+ new_non_attended_tokens = non_attended_tokens[valid_indices]
1225
+
1226
+ # Zero-out the places where we don't need to attend
1227
+ extended_attention_mask[new_batch_index,
1228
+ new_non_attended_tokens] = 0
1229
+
1230
+ attention_mask = torch.cat(
1231
+ (extended_attention_mask, attention_mask[:,
1232
+ -target_length:]),
1233
+ dim=1)
1234
+ position_ids = torch.sum(attention_mask,
1235
+ dim=1).unsqueeze(-1) - 1
1236
+
1237
+ outputs = self.language_model(
1238
+ attention_mask=attention_mask,
1239
+ position_ids=position_ids,
1240
+ past_key_values=past_key_values,
1241
+ inputs_embeds=inputs_embeds,
1242
+ use_cache=use_cache,
1243
+ output_attentions=output_attentions,
1244
+ output_hidden_states=output_hidden_states,
1245
+ return_dict=return_dict,
1246
+ )
1247
+
1248
+ logits = outputs[0]
1249
+
1250
+ loss = None
1251
+ if labels is not None:
1252
+ # Shift so that tokens < n predict n
1253
+ if attention_mask is not None:
1254
+ shift_attention_mask = attention_mask[..., 1:]
1255
+ shift_logits = logits[..., :-1, :][shift_attention_mask.to(
1256
+ logits.device) != 0].contiguous()
1257
+ shift_labels = labels[..., 1:][shift_attention_mask.to(
1258
+ labels.device) != 0].contiguous()
1259
+ else:
1260
+ shift_logits = logits[..., :-1, :].contiguous()
1261
+ shift_labels = labels[..., 1:].contiguous()
1262
+ # Flatten the tokens
1263
+ loss_fct = nn.CrossEntropyLoss()
1264
+ loss = loss_fct(
1265
+ shift_logits.view(-1, shift_logits.size(-1)),
1266
+ shift_labels.view(-1).to(shift_logits.device),
1267
+ )
1268
+
1269
+ if not return_dict:
1270
+ output = (logits, ) + outputs[1:]
1271
+ return (loss, ) + output if loss is not None else output
1272
+
1273
+ return LlavaCausalLMOutputWithPast(
1274
+ loss=loss,
1275
+ logits=logits,
1276
+ past_key_values=outputs.past_key_values,
1277
+ hidden_states=outputs.hidden_states,
1278
+ attentions=outputs.attentions,
1279
+ )
1280
+
1281
+ def prepare_inputs_for_generation(
1282
+ self,
1283
+ input_ids,
1284
+ past_key_values=None,
1285
+ inputs_embeds=None,
1286
+ pixel_values=None,
1287
+ grid_thws=None,
1288
+ attention_mask=None,
1289
+ **kwargs,
1290
+ ):
1291
+ if past_key_values is not None:
1292
+ if hasattr(past_key_values, "get_seq_length"):
1293
+ cache_length = past_key_values.get_seq_length()
1294
+ past_length = getattr(past_key_values, 'seen_tokens',
1295
+ cache_length)
1296
+ else:
1297
+ cache_length = past_length = past_key_values[0][0].shape[2]
1298
+
1299
+ # Keep only the unprocessed tokens:
1300
+ # 1 - If the length of the attention_mask exceeds the length of input_ids, then we are in a setting where
1301
+ # some of the inputs are exclusively passed as part of the cache (e.g. when passing input_embeds as
1302
+ # input)
1303
+ if attention_mask is not None and attention_mask.shape[
1304
+ 1] > input_ids.shape[1]:
1305
+ input_ids = input_ids[:, -(attention_mask.shape[1] -
1306
+ past_length):]
1307
+ # 2 - If the past_length is smaller than input_ids', then input_ids holds all input tokens. We can discard
1308
+ # input_ids based on the past_length.
1309
+ elif past_length < input_ids.shape[1]:
1310
+ input_ids = input_ids[:, past_length:]
1311
+ # 3 - Otherwise (past_length >= input_ids.shape[1]), let's assume input_ids only has unprocessed tokens.
1312
+ elif self.config.media_placeholder_token_id in input_ids:
1313
+ input_ids = input_ids[:, input_ids.shape[1] - 1:]
1314
+ # If the cache has seen more tokens than it can hold, then the cache has a size limit. Let's discard the
1315
+ # older attention values, as their corresponding values are not part of the input.
1316
+ if cache_length < past_length and attention_mask is not None:
1317
+ attention_mask = attention_mask[:, -(cache_length +
1318
+ input_ids.shape[1]):]
1319
+
1320
+ position_ids = kwargs.get("position_ids", None)
1321
+ if attention_mask is not None and position_ids is None:
1322
+ # create position_ids on the fly for batch generation
1323
+ position_ids = attention_mask.long().cumsum(-1) - 1
1324
+ position_ids.masked_fill_(attention_mask == 0, 1)
1325
+ if past_key_values:
1326
+ position_ids = position_ids[:, -input_ids.shape[1]:]
1327
+
1328
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1329
+ if inputs_embeds is not None and past_key_values is None:
1330
+ model_inputs = {"inputs_embeds": inputs_embeds}
1331
+ else:
1332
+ model_inputs = {"input_ids": input_ids}
1333
+
1334
+ model_inputs.update({
1335
+ "position_ids": position_ids,
1336
+ "past_key_values": past_key_values,
1337
+ "use_cache": kwargs.get("use_cache"),
1338
+ "attention_mask": attention_mask,
1339
+ "pixel_values": pixel_values,
1340
+ "grid_thws": grid_thws,
1341
+ })
1342
+ return model_inputs
1343
+
1344
+ def _reorder_cache(self, *args, **kwargs):
1345
+ return self.language_model._reorder_cache(*args, **kwargs)
preprocessor_config.json ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "media_proc_cfg": {
3
+ "in_patch_limit": 4096,
4
+ "patch_size": 14,
5
+ "image_mean": [
6
+ 0.5,
7
+ 0.5,
8
+ 0.5
9
+ ],
10
+ "image_std": [
11
+ 0.5,
12
+ 0.5,
13
+ 0.5
14
+ ],
15
+ "merge_kernel_size": 2,
16
+ "fixed_output_tokens": null,
17
+ "patch_limit_on_one_side": 512,
18
+ "in_patch_limit_each_frame": 16384,
19
+ "in_patch_limit_video": 655360,
20
+ "sample_fps": 8.0,
21
+ "max_num_frames_each_video": null,
22
+ "temporal_merge_kernel_size": 4,
23
+ "timestamp_mode": "hh:mm:ss.fff",
24
+ "transparent_bg_config": {
25
+ "pattern": "chessboard",
26
+ "chessboard_square_size": 8,
27
+ "chessboard_square_on_top_left": true,
28
+ "chessboard_white_value": 255,
29
+ "chessboard_gray_value": 180
30
+ },
31
+ "transparent_bg_fill_stage": "after_resize",
32
+ "config_type": "media_proc.processors.moonvit.MoonViTMediaProcessorConfig"
33
+ },
34
+ "auto_map": {
35
+ "AutoProcessor": "processing_groundinganything.GroundAnythingProcessor"
36
+ }
37
+ }
processing_groundinganything.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Processor glue for GroundAnything-VLM with Kimi-K3 MoonViT preprocessing."""
2
+
3
+ from transformers.feature_extraction_utils import BatchFeature
4
+ from transformers.processing_utils import ProcessorMixin
5
+
6
+ from .media_utils import MediaInput
7
+ from .image_processing_groundinganything import GroundAnythingVLMImageProcessor
8
+
9
+
10
+ class GroundAnythingVLMProcessor(ProcessorMixin):
11
+ attributes = ["image_processor", "tokenizer"]
12
+ image_processor_class = "AutoImageProcessor"
13
+ tokenizer_class = "AutoTokenizer"
14
+
15
+ def __init__(self, image_processor=None, tokenizer=None, chat_template=None, **kwargs):
16
+ del kwargs
17
+ super().__init__(
18
+ image_processor=image_processor,
19
+ tokenizer=tokenizer,
20
+ chat_template=chat_template or getattr(tokenizer, "chat_template", None),
21
+ )
22
+
23
+ @property
24
+ def image_token(self):
25
+ return "<|image_pad|>"
26
+
27
+ @property
28
+ def image_token_id(self):
29
+ return self.tokenizer.convert_tokens_to_ids(self.image_token)
30
+
31
+ def _get_num_multimodal_tokens(self, image_sizes=None, **kwargs):
32
+ del kwargs
33
+ num_image_tokens = []
34
+ num_image_patches = []
35
+ for height, width in image_sizes or ():
36
+ image_stub = type("ImageSize", (), {"size": (width, height)})()
37
+ resize = self.image_processor.get_resize_config(
38
+ {"type": "image", "image": image_stub}
39
+ )
40
+ tokens = int(resize["num_tokens"])
41
+ num_image_tokens.append(tokens)
42
+ num_image_patches.append(tokens * self.image_processor.merge_size**2)
43
+ return {
44
+ "num_image_tokens": num_image_tokens,
45
+ "num_image_patches": num_image_patches,
46
+ }
47
+
48
+ @classmethod
49
+ def register_for_auto_class(cls, auto_class="AutoProcessor"):
50
+ cls._auto_class = auto_class
51
+
52
+ @classmethod
53
+ def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
54
+ import json
55
+ import os
56
+ from transformers import AutoTokenizer
57
+
58
+ kwargs.pop("_from_auto", None)
59
+ kwargs.pop("trust_remote_code", None)
60
+ kwargs.pop("code_revision", None)
61
+ with open(os.path.join(pretrained_model_name_or_path, "preprocessor_config.json"), encoding="utf-8") as f:
62
+ processor_config = json.load(f)
63
+ image_processor = GroundAnythingVLMImageProcessor(
64
+ media_proc_cfg=processor_config["media_proc_cfg"]
65
+ )
66
+ tokenizer = AutoTokenizer.from_pretrained(
67
+ pretrained_model_name_or_path, trust_remote_code=True, **kwargs
68
+ )
69
+ return cls(image_processor=image_processor, tokenizer=tokenizer)
70
+
71
+ def apply_chat_template(self, messages, **kwargs):
72
+ if self.chat_template and "chat_template" not in kwargs:
73
+ kwargs["chat_template"] = self.chat_template
74
+ return self.tokenizer.apply_chat_template(messages, **kwargs)
75
+
76
+ def __call__(
77
+ self,
78
+ text=None,
79
+ images=None,
80
+ return_tensors="pt",
81
+ padding=False,
82
+ **kwargs,
83
+ ):
84
+ return_mm_token_type_ids = kwargs.pop("return_mm_token_type_ids", False)
85
+ if isinstance(text, str):
86
+ text = [text]
87
+
88
+ image_inputs = {}
89
+ if images is not None:
90
+ image_inputs = self.image_processor(
91
+ images=images,
92
+ return_tensors=return_tensors,
93
+ )
94
+ text = list(text)
95
+ image_index = 0
96
+ merge_length = self.image_processor.merge_size**2
97
+ for batch_index, prompt in enumerate(text):
98
+ while self.image_token in prompt:
99
+ grid = image_inputs["image_grid_thw"][image_index]
100
+ num_tokens = int(grid.prod().item()) // merge_length
101
+ prompt = prompt.replace(
102
+ self.image_token, "<|image_placeholder|>" * num_tokens, 1
103
+ )
104
+ image_index += 1
105
+ text[batch_index] = prompt.replace(
106
+ "<|image_placeholder|>", self.image_token
107
+ )
108
+ if image_index != len(image_inputs["image_grid_thw"]):
109
+ raise ValueError(
110
+ "number of image placeholders does not match image inputs"
111
+ )
112
+
113
+ text_inputs = self.tokenizer(
114
+ text,
115
+ return_tensors=return_tensors,
116
+ padding=padding,
117
+ **kwargs,
118
+ )
119
+ if return_mm_token_type_ids:
120
+ input_ids = text_inputs["input_ids"]
121
+ if hasattr(input_ids, "new_zeros"):
122
+ mm_token_type_ids = input_ids.new_zeros(input_ids.shape)
123
+ mm_token_type_ids[input_ids == self.image_token_id] = 1
124
+ else:
125
+ mm_token_type_ids = [
126
+ [int(token == self.image_token_id) for token in row]
127
+ for row in input_ids
128
+ ]
129
+ text_inputs["mm_token_type_ids"] = mm_token_type_ids
130
+
131
+ return BatchFeature(data={**text_inputs, **image_inputs})
132
+
133
+ def batch_decode(self, *args, **kwargs):
134
+ return self.tokenizer.batch_decode(*args, **kwargs)
135
+
136
+ def decode(self, *args, **kwargs):
137
+ return self.tokenizer.decode(*args, **kwargs)
138
+
139
+
140
+ __all__ = ["GroundAnythingVLMProcessor"]
141
+
142
+
143
+ class GroundAnythingProcessor(GroundAnythingVLMProcessor):
144
+ """DLM release processor identity."""
requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ transformers>=5.3,<5.8
2
+ huggingface-hub>=1.3,<2
3
+ safetensors>=0.6
4
+ torch>=2.8
special_tokens_map.json ADDED
The diff for this file is too large to render. See raw diff
 
streammind_gate.py ADDED
@@ -0,0 +1,132 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from dataclasses import dataclass
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+ from mamba_ssm.models.mixer_seq_simple import create_block
7
+ from transformers import Qwen3Config
8
+ from transformers.models.qwen3 import Qwen3ForCausalLM
9
+
10
+
11
+ class PreNet(nn.Module):
12
+ def __init__(self, d_code, d_model):
13
+ super().__init__()
14
+ self.fc3 = nn.Linear(d_code, d_model)
15
+
16
+ def forward(self, x):
17
+ return F.leaky_relu(self.fc3(x))
18
+
19
+
20
+ class PostNet(nn.Module):
21
+ def __init__(self, d_model, n_class):
22
+ super().__init__()
23
+ self.fc3 = nn.Linear(d_model, n_class)
24
+
25
+ def forward(self, x):
26
+ return self.fc3(F.leaky_relu(x))
27
+
28
+
29
+ @dataclass
30
+ class SSMConfig:
31
+ d_model: int = 2560
32
+ n_ssm: int = 1
33
+
34
+
35
+ class VideoMamba(nn.Module):
36
+ def __init__(self, config):
37
+ super().__init__()
38
+ self.ssms = nn.ModuleList(
39
+ [create_block(config.d_model, d_intermediate=0, layer_idx=i) for i in range(config.n_ssm)]
40
+ )
41
+ self.norm_fn = nn.LayerNorm(config.d_model)
42
+
43
+ def forward(self, embeds, inference_params=None):
44
+ hidden_states = embeds
45
+ residual = None
46
+ for ssm in self.ssms:
47
+ hidden_states, residual = ssm(
48
+ hidden_states, residual, inference_params=inference_params
49
+ )
50
+ residual = hidden_states + residual if residual is not None else hidden_states
51
+ return self.norm_fn(residual.to(dtype=self.norm_fn.weight.dtype))
52
+
53
+
54
+ class Qwen3ForCausalLMCls(Qwen3ForCausalLM):
55
+ def forward(self, inputs_embeds=None, labels=None, attention_mask=None, **kwargs):
56
+ outputs = self.model(inputs_embeds=inputs_embeds, attention_mask=attention_mask)
57
+ logits = self.lm_head(outputs.last_hidden_state).float()
58
+ loss = None
59
+ if labels is not None:
60
+ shift_logits = logits[..., :-1, :].contiguous().view(-1, self.config.vocab_size)
61
+ shift_labels = labels[..., 1:].contiguous().view(-1).to(shift_logits.device)
62
+ loss = nn.CrossEntropyLoss(
63
+ weight=torch.tensor([0.15, 0.85], device=shift_logits.device)
64
+ )(shift_logits, shift_labels)
65
+ return {"loss": loss, "logits": logits}
66
+
67
+
68
+ class ClsNet(nn.Module):
69
+ def __init__(self, hidden_size=2560, num_layers=4):
70
+ super().__init__()
71
+ config = Qwen3Config(
72
+ vocab_size=2,
73
+ hidden_size=hidden_size,
74
+ num_hidden_layers=num_layers,
75
+ num_attention_heads=32,
76
+ num_key_value_heads=8,
77
+ intermediate_size=12288,
78
+ head_dim=128,
79
+ max_position_embeddings=8192,
80
+ rms_norm_eps=1e-6,
81
+ tie_word_embeddings=False,
82
+ attention_bias=False,
83
+ )
84
+ self.cls_model = Qwen3ForCausalLMCls(config)
85
+
86
+ def forward(self, x, labels=None, attention_mask=None):
87
+ return self.cls_model(inputs_embeds=x, labels=labels, attention_mask=attention_mask)
88
+
89
+
90
+ class StreamMindGate(nn.Module):
91
+ def __init__(self, hidden_size=2560):
92
+ super().__init__()
93
+ self.pre_net = PreNet(hidden_size, hidden_size)
94
+ self.mamba_model = VideoMamba(SSMConfig(d_model=hidden_size))
95
+ self.post_net = PostNet(hidden_size, hidden_size)
96
+ self.cls_net = ClsNet(hidden_size=hidden_size, num_layers=4)
97
+
98
+ def perception_tokens(self, vision_tokens):
99
+ """Convert [B,T,P,D] visual patches to one EPFE token per time step."""
100
+ x = vision_tokens.mean(dim=2)
101
+ batch, time, dim = x.shape
102
+ x = self.pre_net(x.reshape(batch * time, dim)).reshape(batch, time, dim)
103
+ x = self.mamba_model(x)
104
+ x = self.post_net(x.reshape(batch * time, dim)).reshape(batch, time, dim)
105
+ return x
106
+
107
+ def forward(self, vision_tokens, response_positions=None):
108
+ """Return [B,T,2] silent/speak logits for every EPFE time step."""
109
+ tokens = self.perception_tokens(vision_tokens)
110
+ batch, time, dim = tokens.shape
111
+ target_ids = torch.zeros(batch, time, dtype=torch.long, device=tokens.device)
112
+ if response_positions is not None:
113
+ target_ids[:, torch.as_tensor(response_positions, device=tokens.device) - 1] = 1
114
+ targets = self.cls_net.cls_model.model.embed_tokens(
115
+ target_ids.reshape(batch * time)
116
+ )
117
+ pair = torch.stack((tokens.reshape(batch * time, dim), targets), dim=1)
118
+ rotary = self.cls_net.cls_model.model.rotary_emb
119
+ saved_inv_freq = rotary.inv_freq
120
+ try:
121
+ # Match the training checkpoint, where the full model (including
122
+ # non-persistent Qwen3 RoPE buffers) was cast to BF16.
123
+ rotary.inv_freq = rotary.inv_freq.to(pair.dtype)
124
+ output = self.cls_net(
125
+ pair,
126
+ attention_mask=torch.ones(pair.shape[:2], device=pair.device),
127
+ )
128
+ finally:
129
+ rotary.inv_freq = saved_inv_freq
130
+ # Autoregressive shift: position 0 predicts the target token at
131
+ # position 1, matching StreamMind's logits[..., :-1, :] evaluation.
132
+ return output["logits"][:, 0].reshape(batch, time, 2)
tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7e0398aca93659140fd6345deb335db2330dee89a2aad1228d8604c1952f6dc8
3
+ size 11604906
tokenizer_config.json ADDED
The diff for this file is too large to render. See raw diff
 
vocab.json ADDED
The diff for this file is too large to render. See raw diff