Skywalker0410 commited on
Commit
ba4cb9a
·
verified ·
1 Parent(s): b44d814

Upload model files with empty README

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ tokenizer.json filter=lfs diff=lfs merge=lfs -text
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,11 +0,0 @@
1
- ---
2
- license: apache-2.0
3
- ---
4
-
5
- # GroundAnything
6
-
7
- ## Related repositories
8
-
9
- - [GroundingPI](https://huggingface.co/GroundingPI/GroundingPI)
10
- - [GroundAnything-VLM](https://huggingface.co/GroundingPI/GroundAnything-VLM)
11
- - [GroundAnything](https://huggingface.co/GroundingPI/GroundAnything)
 
 
 
 
 
 
 
 
 
 
 
 
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
+ }
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
+ e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855 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