kingjones777 commited on
Commit
18c1466
·
verified ·
1 Parent(s): 7704dae

Add files using upload-large-folder tool

Browse files
Files changed (50) hide show
  1. LICENSE +21 -0
  2. README.md +239 -0
  3. SHA256SUMS +94 -0
  4. code/.gitignore +10 -0
  5. code/LICENSE +21 -0
  6. code/README.md +279 -0
  7. code/REVISION +1 -0
  8. code/bailingmm_utils.py +504 -0
  9. code/configuration_bailing_moe_v2.py +85 -0
  10. code/configuration_bailingmm2.py +43 -0
  11. code/examples/profiles/generation_edit.json +8 -0
  12. code/examples/profiles/layer_decompose.json +8 -0
  13. code/generate_paired.sh +162 -0
  14. code/image_processing_bailingmm2.py +451 -0
  15. code/infer.py +468 -0
  16. code/inference_profile.py +420 -0
  17. code/mllm_device_map.py +146 -0
  18. code/modeling_bailing_moe_v2.py +2031 -0
  19. code/modeling_bailingmm2.py +814 -0
  20. code/modeling_utils.py +191 -0
  21. code/pe_ling.py +446 -0
  22. code/processing_bailingmm2.py +512 -0
  23. code/quant/__init__.py +1 -0
  24. code/quant/int8_linear.py +221 -0
  25. code/quant/quantize_stream.py +556 -0
  26. code/qwen2_5_vit.py +508 -0
  27. code/requirements-rocm.txt +36 -0
  28. code/requirements.txt +23 -0
  29. code/rocm.patch +0 -0
  30. code/tests/assets/smoke_input.png +0 -0
  31. code/tokenization_bailing.py +1024 -0
  32. connector/config.json +58 -0
  33. connector/generation_config.json +14 -0
  34. mllm/config.json +210 -0
  35. mllm/int8_manifest.json +0 -0
  36. mllm/model.safetensors.index.json +0 -0
  37. mllm/preprocessor_config.json +23 -0
  38. mllm/special_tokens_map.json +44 -0
  39. mllm/tokenizer_config.json +2334 -0
  40. mlp/config.json +17 -0
  41. samples/cabin_upstream.json +33 -0
  42. samples/e2e_daily_grind.caption.txt +1 -0
  43. samples/e2e_daily_grind.json +102 -0
  44. samples/info_water.json +141 -0
  45. samples/poster_jazz.json +86 -0
  46. samples/ui_banking.json +94 -0
  47. scheduler/scheduler_config.json +7 -0
  48. transformer/config.json +35 -0
  49. transformer/diffusion_pytorch_model.safetensors.index.json +526 -0
  50. vae/config.json +25 -0
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2026 inclusionAI
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
README.md ADDED
@@ -0,0 +1,239 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ base_model: inclusionAI/Ming-Image-0.1-Design
4
+ base_model_relation: quantized
5
+ pipeline_tag: text-to-image
6
+ tags:
7
+ - text-to-image
8
+ - rocm
9
+ - amd
10
+ - strix-halo
11
+ - gfx1151
12
+ - int8
13
+ - design
14
+ - rgba
15
+ ---
16
+
17
+ # Ming-Image-0.1-Design — ROCm build (AMD Strix Halo, gfx1151) · INT8 MLLM · paired with Ling-3.0-flash-VL
18
+
19
+ [inclusionAI/Ming-Image-0.1-Design](https://huggingface.co/inclusionAI/Ming-Image-0.1-Design) (text-to-image for
20
+ UI, posters and infographics, RGBA output) made to run on AMD ROCm, with the routed experts of its 17B-parameter
21
+ MoE language model stored as weight-only INT8, and wired to the prompt enhancer its model card names: **Ling-3.0-flash-VL**
22
+ (served from our [ROCmFP4 build](https://huggingface.co/kingjones777/Ling-3.0-flash-VL-MTP-ROCmFP4-GGUF)).
23
+
24
+ Everything below was measured on one AMD Ryzen AI Max+ 395 (Radeon 8060S, gfx1151) — see
25
+ [Reproduction](#reproduction). Nothing here was run on CUDA.
26
+
27
+ ## Results (1024 × 1024, 12 steps, cfg 1.0, seed 42, four prompts)
28
+
29
+ | prompt (1024², 12 steps) | BF16 · upstream code | BF16 · contiguous attention | INT8 · contiguous attention | INT8 · fast attention |
30
+ |---|---:|---:|---:|---:|
31
+ | four-seasons cabin (upstream's example prompt) | 339.9 s | 211.1 s | 181.1 s | 77.8 s |
32
+ | banking-app screen (Ling rewrite) | 501.8 s | 270.0 s | 259.7 s | 88.5 s |
33
+ | jazz-night poster (Ling rewrite) | 355.5 s | 221.4 s | 219.9 s | 79.4 s |
34
+ | water-cycle infographic (Ling rewrite) | 606.7 s | 312.0 s | 250.0 s | 100.3 s |
35
+ | **mean** | 451.0 s | 253.6 s | 227.7 s | 86.5 s |
36
+ | speed-up vs upstream code | 1.00× | 1.78× | 1.98× † | 5.21× |
37
+
38
+ | configuration | PyTorch peak allocated | Ming's own GTT peak | box |
39
+ |---|---:|---:|---|
40
+ | BF16 · upstream code, all components resident | 58.7 GiB | 74.5 GiB | Ling stopped |
41
+ | BF16 · contiguous attention | 58.9 GiB | 74.5 GiB | Ling stopped |
42
+ | BF16 · contiguous · `--release-mllm-after-conditioning` (cabin) | 47.2 GiB | — (baseline not settled) | Ling stopped |
43
+ | **INT8 · contiguous · `--release-mllm-after-conditioning`** | **33.6 GiB** | **35.8 GiB** | Ling-3.0-flash-VL resident (67.5 GiB) → box peak 103.3 GiB of 128 GiB GTT |
44
+ | INT8 · fast attention · release | 33.6 GiB | 35.1 GiB | Ling resident |
45
+
46
+ - **Contiguous attention is the big one:** 1.78× faster end to end, and the images are **byte-identical** to the upstream code path (4/4 prompts plus the `--release` run, compared with `cmp`).
47
+ - † **INT8** runs the same BF16 diffusion transformer, so it does not change the step time: its 227.7 s mean vs 253.6 s for BF16 with the same attention reflects when the runs happened (later, with Ling resident), not the quantization. What INT8 does cost is the conditioning pass, where the weights are dequantized on the fly: 2.4–4.7 s vs 1.1–2.0 s in BF16. INT8 exists for memory: with `--release-mllm-after-conditioning` Ming needs 35.8 GiB at its peak, which fits next to the resident Ling-3.0-flash-VL (67.5 GiB) on a 128 GiB box — BF16 with the upstream code needed 74.5 GiB and did not. Loading takes 30–100 s; the slow end is a cold page cache (first load after the files were written or after Ling was restarted).
48
+ - **Fast attention** (`--attention-bf16-reduction`) is opt-in: 5.21× vs upstream, at the fidelity cost shown below.
49
+ - The very first image on a fresh box is slower once: the reference cabin took 608.4 s the first time and 339.9 s warm (MIOpen tunes the VAE's 3-D convolutions and caches the result in `~/.cache/miopen`).
50
+
51
+ ### Fidelity against the BF16 original
52
+
53
+ Inference on this box is **deterministic**: the same prompt and seed produced a byte-identical PNG twice, in BF16
54
+ and in INT8, so every difference below is caused by the quantization (or by the fast-attention option), not by
55
+ run-to-run noise.
56
+
57
+ | prompt | INT8 SSIM | INT8 PSNR | cond cos (query / VLM tokens) | cond rel L2 (query / VLM) | INT8 fast SSIM | INT8 fast PSNR |
58
+ |---|---:|---:|---:|---:|---:|---:|
59
+ | four-seasons cabin (upstream's example prompt) | 0.9345 | 22.90 dB | 0.98169 / 0.99672 | 0.1931 / 0.0810 | 0.9150 | 22.13 dB |
60
+ | banking-app screen (Ling rewrite) | 0.9666 | 23.76 dB | 0.98638 / 0.99623 | 0.1655 / 0.0868 | 0.9624 | 22.75 dB |
61
+ | jazz-night poster (Ling rewrite) | 0.9450 | 23.31 dB | 0.98874 / 0.99587 | 0.1504 / 0.0909 | 0.9456 | 23.55 dB |
62
+ | water-cycle infographic (Ling rewrite) | 0.8486 | 17.49 dB | 0.99061 / 0.99500 | 0.1378 / 0.1000 | 0.8439 | 17.42 dB |
63
+
64
+ `cos` / `rel L2` compare the conditioning tensors the diffusion transformer receives (the MLLM output — the only
65
+ part that is quantized). SSIM is windowed 7×7 on luminance; PSNR over RGB.
66
+
67
+ ## Run it
68
+
69
+ ```bash
70
+ hf download kingjones777/Ming-Image-0.1-Design-ROCm-INT8 --local-dir Ming-Image-0.1-Design-ROCm-INT8
71
+ cd Ming-Image-0.1-Design-ROCm-INT8
72
+ # caption -> Ling-3.0-flash-VL rewrites it into the layered JSON prompt -> Ming renders it (one image per process)
73
+ bash code/generate_paired.sh --model . --base-url http://127.0.0.1:8090/v1 \
74
+ "A minimalist concert poster for a jazz night called \"Blue Hour\"" --resolution 1024 --output-dir out \
75
+ -- --device-map none --release-mllm-after-conditioning # add --attention-bf16-reduction for fast mode
76
+ ```
77
+
78
+ `PYTHON=/path/to/python` selects the interpreter (it needs a ROCm build of PyTorch and
79
+ [`code/requirements-rocm.txt`](code/requirements-rocm.txt)); `--pe-model` names the chat model your endpoint serves.
80
+ This is the command shape our end-to-end run used (see [Samples](#samples)); 1024² is the resolution we measured —
81
+ `infer.py`'s own default is 2048², which we did not run.
82
+
83
+ `generate_paired.sh` calls `pe_ling.py` (standard library only) against an OpenAI-compatible endpoint
84
+ (default `http://127.0.0.1:8090/v1`, our Ling-3.0-flash-VL llama-server seat), validates the rewrite against the
85
+ schema the upstream system prompt demands (one retry with the errors), then runs `infer.py`. You can also pass
86
+ your own JSON prompt straight to `infer.py --prompt prompt.json`.
87
+
88
+ The code is in [`code/`](code/): upstream `inclusionAI/Ming-Image` at `62c6072` plus the changes below
89
+ ([`code/rocm.patch`](code/rocm.patch)).
90
+
91
+ ## What changed for ROCm, and why
92
+
93
+ 1. **`transformer_engine` removed.** The vision tower imported NVIDIA Transformer Engine (CUDA-only) for a
94
+ single `te.RMSNorm`; it is now a plain RMSNorm with the same parameter name (`weight`), so all 65 vision norm
95
+ tensors in the checkpoint load unchanged. A dead import in `modeling_bailing_moe_v2.py` went too.
96
+ 2. **`--attn-implementation eager` now reaches the towers.** `BailingMM2Config` declared no `sub_configs`, so
97
+ transformers never copied the chosen attention implementation into the vision and language configs; their
98
+ `flash_attention_2` defaults raised `ImportError` at model construction on any box without flash-attn. Declaring
99
+ `sub_configs` fixes it. (Upstream's `--validate-only` cannot catch this — it never builds the model.)
100
+ 3. **Contiguous attention inputs.** The diffusion transformer handed PyTorch's attention permuted views of
101
+ `[B, L, H, D]` tensors. On gfx1151 the only working SDPA kernel is the math one, and with those strides its
102
+ fp32 GEMMs fall onto an 8×8×8 macro-tile: 644 ms per attention call at the cabin prompt's length
103
+ (L = 5,759) versus 363 ms for the same tensors made contiguous, bit-identical output. The profiler put
104
+ SDPA-math at 87.6% of all GPU time before the fix. `diffusion/transformer.py` now passes
105
+ contiguous tensors whenever diffusers' default native backend is active.
106
+ 4. **INT8 MLLM — routed experts only.** Weight-only, per-output-channel symmetric INT8 (fp32 scales) for the
107
+ 256 routed experts of the language model's 19 MoE layers — 14,592 Linears, which hold
108
+ almost all of its weights. Everything that runs on every token stays byte-identical BF16: attention
109
+ (`query_key_value`, `dense`), the shared experts, layer 0's dense MLP, the routers (`gate`, `image_gate`,
110
+ `audio_gate`), embeddings, `lm_head`, norms and the vision tower (691 tensors).
111
+ `mllm/` goes from 34.00 GB to 18.76 GB. Relative weight error
112
+ ‖W − Ŵ‖/‖W‖: mean 0.00833, p99 0.01035, max
113
+ 0.01412. `quant/quantize_stream.py` writes it shard by shard without building the model
114
+ (6.7 min; peak RSS 17.84 GiB, sampled on an earlier run of the same
115
+ tool); `quant/load_int8.py` builds the model on the meta device and loads the INT8 shards straight onto the GPU, so
116
+ BF16 weights for the quantized layers never exist in memory. The scales stay fp32 through `.to(bfloat16)`.
117
+ *Measured alternative:* quantizing the always-on Linears too saves another 0.30 GB but raised the mean
118
+ conditioning error (rel L2, VLM tokens) from 0.090 to 0.108 and changed the cabin prompt's
119
+ surround from scenery to white (SSIM 0.699); mean SSIM over the four prompts 0.877 vs 0.924.
120
+ 5. **Connector stored as bf16.** It shipped as float32 (6.17 GB) but `infer.py` always loads it as bf16; storing
121
+ it pre-rounded halves the download and every tensor equals `fp32.to(bfloat16)` exactly (all 338 checked).
122
+ 6. **`--release-mllm-after-conditioning`** (opt-in, one image per process): the language model, vision tower and
123
+ connector are only needed for the ~1–2 s conditioning pass, so they are freed before the 12 diffusion steps.
124
+ PyTorch peak allocated on the cabin prompt (BF16): 55.6 → 47.2 GiB, byte-identical image.
125
+ 7. **`--attention-bf16-reduction`** (opt-in): lets the math SDPA kernel stay in bf16 instead of upcasting to fp32 —
126
+ 363 → 115 ms per call, attention error vs an fp32 reference 1.658e-03 → 5.032e-03
127
+ (rel L2, random inputs). Image-level cost is in the fidelity table.
128
+
129
+ ## Things that do not work on gfx1151 (measured, so you don't have to)
130
+
131
+ - **Only the math SDPA kernel runs.** AOTriton's efficient and flash kernels are unavailable by default
132
+ (`UNAVAILABLE: No available kernel. Aborting execution.`).
133
+ - **Do not set `TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1`**, even though PyTorch's warning suggests it: the shipped
134
+ `amd-gfx11xx` AOTriton kernels then get selected and fail (`UNAVAILABLE: HIP error: invalid argument`), and a full
135
+ generation crashes with the same `HIP error: invalid argument`.
136
+ - **The first image on a fresh box is slower once.** The same cabin prompt took 608.4 s the
137
+ first time and 339.9 s warm; MIOpen tunes the VAE's 3-D convolutions on first use and its cache
138
+ (`~/.cache/miopen`) grew during that first run.
139
+ - **Inference is deterministic here, BF16 and INT8:** same prompt + seed → byte-identical PNG (both checked), which
140
+ is what makes the fidelity numbers above exact rather than statistical.
141
+
142
+ ## Pairing with Ling-3.0-flash-VL
143
+
144
+ Ming-Image's text-to-image quality depends on a structured, Figma-like JSON prompt (canvas settings +
145
+ positioned layers with colours). Upstream ships the rewriter system prompt (`assets/t2i_rewriter_system_prompt.txt`) and names Ling-3.0-flash-VL as the model to run it.
146
+ `pe_ling.py` sends that prompt verbatim plus your caption to any OpenAI-compatible endpoint, extracts the JSON, validates it, and retries once with the validation errors if needed.
147
+
148
+ Measured against our Ling-3.0-flash-VL seat (llama-server, `ling-3.0-flash-vl-mtp-halo-STRIX_LEAN`, same box) — all three valid on the first attempt:
149
+
150
+ | caption | rewrite time | layers |
151
+ |---|---:|---:|
152
+ | banking-app screen (Ling rewrite) | 45.9 s | 7 |
153
+ | jazz-night poster (Ling rewrite) | 34.8 s | 8 |
154
+ | water-cycle infographic (Ling rewrite) | 49.8 s | 14 |
155
+
156
+ The exact captions and the JSON Ling returned are in [`samples/`](samples/).
157
+
158
+ ## Samples
159
+
160
+ INT8 build, `--release-mllm-after-conditioning`, 1024², 12 steps, cfg 1.0, seed 42 — exactly the runs in
161
+ the tables above. Full-resolution RGBA PNGs and the JSON prompts are in `samples/`.
162
+
163
+ **cabin upstream** — upstream's own structured prompt (`assets/t2i_four_seasons_cabin_prompt.json`), no rewrite · [prompt JSON](samples/cabin_upstream.json)
164
+
165
+ ![cabin_upstream](samples/cabin_upstream.png)
166
+
167
+ **ui banking** — caption → Ling-3.0-flash-VL: *A mobile banking app home screen: a balance card at the top, a recent transactions list, quick-action buttons for send, pay and top up, and a bottom navigation bar. Clean modern fintech style.* · [prompt JSON](samples/ui_banking.json)
168
+
169
+ ![ui_banking](samples/ui_banking.png)
170
+
171
+ **poster jazz** — caption → Ling-3.0-flash-VL: *A minimalist concert poster for a jazz night called "Blue Hour" on Friday, October 3, 8 PM at The Lantern Room, with a saxophone silhouette over a deep blue gradient.* · [prompt JSON](samples/poster_jazz.json)
172
+
173
+ ![poster_jazz](samples/poster_jazz.png)
174
+
175
+ **info water** — caption → Ling-3.0-flash-VL: *An infographic that explains the water cycle in four labeled stages - evaporation, condensation, precipitation, collection - with arrows and simple flat icons.* · [prompt JSON](samples/info_water.json)
176
+
177
+ ![info_water](samples/info_water.png)
178
+
179
+ ### End to end on the box
180
+
181
+ `ming-paired` (our box's wrapper around `generate_paired.sh`, not part of this repo) with Ling-3.0-flash-VL and two other model servers running: it paused 2 of them for the render and restarted them afterwards (halo-bonsai8b, halo-bonsai). Ling rewrote the caption in 48.3 s (9 layers); the whole command took 323 s.
182
+
183
+ *A landing page hero section for a coffee subscription service called Daily Grind: the headline Fresh beans every Monday, a short tagline, a Start your subscription button, and a photo of latte art on the right. Warm earthy palette.* · [prompt JSON](samples/e2e_daily_grind.json)
184
+
185
+ ![end to end](samples/e2e_daily_grind.png)
186
+
187
+ ## Files
188
+
189
+ | path | size | what |
190
+ |---|---:|---|
191
+ | `mllm/` | 18.78 GB | language model + vision tower; the 14,592 routed-expert Linears are **INT8** (`int8_manifest.json` lists them) |
192
+ | `transformer/` | 12.31 GB | diffusion transformer, BF16, unchanged |
193
+ | `connector/` | 3.09 GB | Qwen2 connector, **stored as bf16** (upstream ships fp32; runtime identical) |
194
+ | `vae/` | 253.82 MB | VAE (4-channel RGBA), unchanged |
195
+ | `mlp/` | 124.84 MB | conditioning MLP, unchanged |
196
+ | `scheduler/` | 173 B | flow-matching scheduler config, unchanged |
197
+ | `code/` | 2.31 MB | patched inference code + `rocm.patch` + tools |
198
+ | `samples/` | 4.56 MB | the sample images and prompts shown above |
199
+ | **total** | **34.56 GB** | upstream: 52.88 GB |
200
+
201
+ ## Reproduction
202
+
203
+ ```
204
+ upstream : inclusionAI/Ming-Image-0.1-Design revision 1cd7fac3b0dcb54196fe2cd12b80da09edf8fcf4
205
+ code : inclusionAI/Ming-Image @ 62c6072e1ff15af83f7c4963a0a1954c1424e80e + code/rocm.patch (branch rocm-halo @ f986f7a (upstream inclusionAI/Ming-Image 62c6072 + 3 commits); rocm.patch sha256 03fe16f8ed566bf99ec654caa687c7e659f489df6d077cfb7fc7aee134ebc777)
206
+ python : 3.13.5 · torch 2.10.0 (HIP 7.13.99004) · transformers 4.57.1 · diffusers 0.36.0 · accelerate 1.13.0 · safetensors 0.8.0
207
+ box : amd-halo · AMD RYZEN AI MAX+ 395 w/ Radeon 8060S · 125 GiB RAM visible, 128 GiB GTT · ROCm 7.13.0 · kernel 6.18.35+rex+2-amd64
208
+ power : platform_profile=balanced · governor=powersave · GPU 83–107 W at 100% busy (step probe, 0.5 s samples)
209
+ env : nothing set; in particular TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL is NOT set (it crashes gfx1151)
210
+ date : 2026-09-22 to 2026-09-23 (runs crossed midnight, America/Chicago)
211
+ ```
212
+
213
+ Build the INT8 package from the upstream download:
214
+
215
+ ```bash
216
+ python code/quant/quantize_stream.py <upstream>/mllm <package>/mllm \
217
+ --exclude '\.attention\.|\.shared_experts\.|^model\.model\.layers\.0\.mlp\.' # routed experts only
218
+ python code/tools/convert_connector.py <upstream>/connector <package>/connector # fp32 -> bf16, proven exact
219
+ # transformer/ vae/ mlp/ scheduler/ LICENSE are the upstream files, unchanged
220
+ python code/tools/verify_package.py <upstream> <package> # the checks behind this card
221
+ ```
222
+
223
+ Measure (one process per image, the pairing configuration):
224
+
225
+ ```bash
226
+ cd code && PYTHONPATH=. python tools/ming_bench.py --prompts <prompt.json> --out <dir> -- \
227
+ --model <package> --task text-to-image --resolution 1024 --device-map none \
228
+ --attn-implementation eager --release-mllm-after-conditioning # add --attention-bf16-reduction for fast
229
+ python tools/fidelity_compare.py <bf16-dir> <int8-dir> --json fidelity.json
230
+ ```
231
+
232
+ BF16 references ran with Ling stopped (the box cannot hold both); every INT8 run had Ling-3.0-flash-VL resident.
233
+ GPU memory is the amdgpu `mem_info_gtt_used` peak sampled every second, minus a baseline taken after GTT settled.
234
+
235
+ ## License
236
+
237
+ MIT, same as the original. Model weights, architecture and inference code © 2026 inclusionAI
238
+ ([LICENSE](LICENSE)). The INT8 quantization, the ROCm changes and the pairing scripts are ours; they are listed in
239
+ this card and in `code/rocm.patch`.
SHA256SUMS ADDED
@@ -0,0 +1,94 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ b772c54f1f150a680e7c083a92431c518cf62821a407f6cccdf53f6a07bcf702 LICENSE
2
+ b1a5d04726679cfa37262199478de1c0c1735b2d5ac34ede7ac49c72fd35e10a README.md
3
+ 561cdddbafb1cbabe417d2ad2e90ec1dd28a7c131aaead2cc9ded59e9e27b685 code/.gitignore
4
+ b772c54f1f150a680e7c083a92431c518cf62821a407f6cccdf53f6a07bcf702 code/LICENSE
5
+ bef9e77db8ade9d625b80869a43319beb9e8353927ea279d095f69a2fd7c9dcc code/README.md
6
+ 7416e29e85a7997d7550ed455226473a0a95217faa3523aa9ab2125055fca4ad code/REVISION
7
+ a763ccc5a39740b19fc610581f14041914739feeb83c28b712a72a8a896cc978 code/assets/layer_decompose_5layers.txt
8
+ e445a68e4969c1f2734f00a4223559cf58904c2299019346816571db9dee35c3 code/assets/layer_samples/card_making_decomposition.png
9
+ a52347559312b7ddbf435d042bdc32c83efe8750d3071580e6a8eceee2412112 code/assets/layer_samples/card_making_input.png
10
+ 3e06063e76ffd32b80ad6989352f47006fd7f936eb808ff563ec2392898c1884 code/assets/layer_samples/card_making_prompt.txt
11
+ 387468da6db92aeb773e1776434a5c9b723de7d2f059e3304f4de3a9c862eaf6 code/assets/t2i_four_seasons_cabin_prompt.json
12
+ 62b0c7c0380fdccf34abf490d1da6ad971a2e63d82f99144c5cd4c9a8c80480c code/assets/t2i_rewriter_system_prompt.txt
13
+ 897fd9b0541efcc15e5734c9afd943c37a7b29b37d465556a43b10f5815d05cd code/bailingmm_utils.py
14
+ 7d49dc16d2c31c05ea25e76089c635a9b56c629a6a416c88d8fba4e05a13c0ab code/configuration_bailing_moe_v2.py
15
+ ac329fb7c8959b16b68b774693672d1f991d7fe4489d484dc72b4e389c10706f code/configuration_bailingmm2.py
16
+ e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855 code/diffusion/__init__.py
17
+ 4013c3e09e3684d2ca8bbf1f98144a8ebab2a9408641831b1661d89b4a0aa106 code/diffusion/autoencoder_kl_qwenimage.py
18
+ a31f006c123febdcfdfeb92b1b55f2716b154a98e3c3e13bd7bf825148e9901e code/diffusion/generator.py
19
+ 511f76e9292f82076200b4c493a481782df9a8d3eea38ecc5c453b75d3faebe0 code/diffusion/padding.py
20
+ ab42fc223ddff73f854d8b6595e1207ab184248d32151adbd847c0e9974f2343 code/diffusion/pipeline.py
21
+ 287522eea36b6538e5e5dd38a9cee6a21d47053f9868d79457083cae281f0293 code/diffusion/transformer.py
22
+ b373bf064d60b7dd493b44d8e54fa1d939fc134fc63766c3f9e9d7d405cf31d6 code/examples/profiles/generation_edit.json
23
+ 51f773286518ce8ca9034fb12bf7b527ce52a4fcc09f975e9df04d968276a244 code/examples/profiles/layer_decompose.json
24
+ 99e45db68ad2a119d07feb5efe2da3f788ea33ad5997b7a0c8406c9a6e00d5e8 code/generate_paired.sh
25
+ b54ff75e3452de090d29c0f1bc0ce082c510a77405808ea6a59aebef6a691142 code/image_processing_bailingmm2.py
26
+ e39ee57066c183f8411ae3038f8ac491a82a59bde8a9e75ac41f4ed2bb368f92 code/infer.py
27
+ 84f407d5e323c7261792492ed400bb23e3a189cb3edfa017e6c854263a8ee490 code/inference_profile.py
28
+ 02bae9078a230d2414861368894eee027403f324a94b95b07272d79a75a0bf88 code/mllm_device_map.py
29
+ 9da007026e814d17b1bbb8577a03fe70bc4ab6f43854bf6b7f5fba5f846fd429 code/modeling_bailing_moe_v2.py
30
+ 9816b917c9661020160f1cd6f4b895f8a9c6dfa42a163db7751f9aad29bab00b code/modeling_bailingmm2.py
31
+ 81b0be70fbe32ca6ac34110d58653a18462c0c73bf797579e8e6be56472be156 code/modeling_utils.py
32
+ aa4d34184c0d513de0ce6948bc0c0d17195b521ecfa1de606eb7f437dab4739b code/pe_ling.py
33
+ 96988388c47935b2f54eeb14ca10dab16044dd7dcf0ca67d2c284c46e49f4861 code/processing_bailingmm2.py
34
+ ba6414475da49f4696e414eb48b34da71108519bf5b538bd51acd780da8f3f5e code/quant/__init__.py
35
+ 05be56bc8b26fe63825b0303e9f830a898e95c21b23847cbe7eb8e40616dcc41 code/quant/int8_linear.py
36
+ d80aa13e407519078059f847ce5f0c5117ce67606a2766be65809e3a44999b53 code/quant/load_int8.py
37
+ a0466c47611eed8ed3fdd627f1c0cc00ed9ee6760407b3cf4011e24c6c420870 code/quant/quantize_stream.py
38
+ 3eb9c6eab13676931f1a60856249e0cdebb80e4d55b8d4772c4fe862c69bc50b code/quant/test_int8.py
39
+ 08bc15e548fc07d55626b2707b69253c7cb9e69ffa490ba00d215b6740126ba5 code/qwen2_5_vit.py
40
+ 022e5312f4c0eb2be01c03abbdceb905e0fb9d7005473b51d337f4b755c6d807 code/requirements-rocm.txt
41
+ f82a235b7e7da97b59d4121473e2b002bfef9960344050913e8c75f21e6d1c6e code/requirements.txt
42
+ 03fe16f8ed566bf99ec654caa687c7e659f489df6d077cfb7fc7aee134ebc777 code/rocm.patch
43
+ e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855 code/tests/__init__.py
44
+ d5d02c0618be0e4e13dc74f14bc8376396d0a9ed3c06152d3a78ed35485de019 code/tests/assets/smoke_input.png
45
+ 0cc02c16a9f34b3700bef5c1e04997b71e3cf89696c1ffe6daf13151dd5c4451 code/tests/test_infer_cli.py
46
+ 1d9f5d94101fadb72416320bf073284e5cfc4b30e3ad539b606ff4485c83dfa9 code/tests/test_inference_profile.py
47
+ 35edb0b9a378cd7c7be0e0488b9f6ae66ac01b15cab65a71df6b993c21f17af0 code/tests/test_inference_smoke.py
48
+ 21648462c682f9cd03eddb178907104ffa1e0b6a4e79fb760f755385aea4abf8 code/tests/test_mllm_device_map.py
49
+ 4f2d6347f360399392257caa22da310d2a170f570eb5da1dff4cde63feee3a62 code/tests/test_padding.py
50
+ 9c9edaf036544e209c9b6195233fcb3b0baa3752dc1f59923b058bb654ffc5b8 code/tests/test_runtime_precision.py
51
+ 719b7dbdb29f18d74f0a7fae070383d2d612173f8c5c6d42d5a0496f3071ae98 code/tokenization_bailing.py
52
+ 4f9040a6b8b7917860c4df4ecc9dd56fafead0643458221f6b0179255f03597b code/tools/convert_connector.py
53
+ 47057cbfc5224951617b55c2e6600867145d57c0571c62857f01f1990b96f9f2 code/tools/fidelity_compare.py
54
+ a247b28869d96e3323f47b573fcb0d432af86f67e41db29a37dfbff92ffe1bcf code/tools/ming_bench.py
55
+ f5d885aefc9ec44476e194d32bf614a58d0a20eff15dea3e005cb79994e043ba code/tools/sdpa_layout.py
56
+ 1c54fd63d74608da8f61ef3012cf9420627b7d9d01b337f45653d08a2abe11cf code/tools/step_probe.py
57
+ 7f4ddb78131e9de22c30e2522cf1b667f7be5b2aeceafbc63cf67d3dcc94e00c code/tools/verify_package.py
58
+ cc743612957ccbb78d9812e060e87fe096ea1c0f2ad63519b5eb49f249724f95 connector/config.json
59
+ cd5866254bd02b52464b75af06818afa882298a64d30a3819343a4fa4ed16ca9 connector/generation_config.json
60
+ 1522edb909db45f90d3628aa5cb03805fba81aca1950dfe9385b057c6b636a47 connector/model.safetensors
61
+ f797c39adcd1624e5bfc94f60d2a86095dc40978f8cf230fcd11ada877c597ea mllm/config.json
62
+ 1472f23b9a1bc48af3298875a075fa667ffc2c147c33acad6c3ae6ba77bc13f7 mllm/int8_manifest.json
63
+ fd19afe5acf56675a154a2ea5f79a7844514d7500dd4efe6b7e771524fdf9c52 mllm/model-00001-of-00004.safetensors
64
+ a1f3472a28da6607565cd2e31324c25b17499b61620666d090d3f6a434ebdfa7 mllm/model-00002-of-00004.safetensors
65
+ d9be6de14c1f0fabaab1941b895bee49e97a015a469f7da31e622f6f10ba368e mllm/model-00003-of-00004.safetensors
66
+ c770d343ff03238ec3f3a6c2cb594c7efe286eca5a790612bbe72c4e5155783c mllm/model-00004-of-00004.safetensors
67
+ 7dcaa09cd8129511cf675dc5ad832ebdff251f4a3e155f4ebeac272977d4fa30 mllm/model.safetensors.index.json
68
+ e63066dde0329df1642c8dd9a269d613307804084d12e0787ccc887839339091 mllm/preprocessor_config.json
69
+ 8c383d444e75b12f2643b4695bece7f3dd1d43725d8a26deeb019aa40a8f891d mllm/special_tokens_map.json
70
+ e7ff01708d504f7bf4dbf7f5815adde57bab9a40e7f563ab6ad1acace4464917 mllm/tokenizer.json
71
+ 861c20bea04c2057362f226658ad9209978ff294b8a95591a0d32f69dbd3fc28 mllm/tokenizer_config.json
72
+ 682c34c6c39173b396aa24f444098e495c7b731247d935552dc7d6714154adce mlp/config.json
73
+ 53e47c1ec942749f07025a41be148f5e8364d8587aa2d954969716abe521b77b mlp/model.safetensors
74
+ 387468da6db92aeb773e1776434a5c9b723de7d2f059e3304f4de3a9c862eaf6 samples/cabin_upstream.json
75
+ 1c9aaa45455f78c289bc41d8aa175cc4e01c46b1c57e109be043b9f78de57bc1 samples/cabin_upstream.png
76
+ cf9acb660827aad5de550c74cc825508f1b46db59dab42af07f4c0f82830f303 samples/e2e_daily_grind.caption.txt
77
+ f5c561cce8c455622a8342331b0eaa57d4e276041266ba77dee63332b0b904da samples/e2e_daily_grind.json
78
+ 654d8f20a74959a987db7ea306c03618762f07a103fefdef39305818db2e91e8 samples/e2e_daily_grind.png
79
+ e86c214427617b69d1ac4bf126120edbb75c1a6e70fe37669d908997abc8c4f4 samples/info_water.json
80
+ e9906c15e21da79428bd3dd0e8dbb37eca28d61ffbef49aa3b6801ade1f4d8e6 samples/info_water.png
81
+ 864b07f559fad0bbc38b5dd3ff831debf35a4fe82a33ab39f8d82fdbff8625f8 samples/poster_jazz.json
82
+ 1e7083ba5d40fb97b4269148b127f780a21aae6c7ebc534adb3c1ffc6fd7b3b5 samples/poster_jazz.png
83
+ 6c9ddda781fa033068da8dcc414996f6d592efae322d8b2d45b24d157c34215e samples/ui_banking.json
84
+ 264207a5c73e24dd8292a85a0ad1fb4216dec03998ac98e81da7b85b27754e5b samples/ui_banking.png
85
+ 79bd98cddf1a60d365f3072d713ffad832239f57c837bba0e9da3cff40bfafc5 scheduler/scheduler_config.json
86
+ 9a64223dd9858f5a1120f4593534173b57741bb8ec4bac53f1636058c820abcd transformer/config.json
87
+ 1488ae867d45a4c184fa0df9fde68d8f7bd9c2dea52b611b77fc6bbd685d9d70 transformer/diffusion_pytorch_model-00001-of-00005.safetensors
88
+ 6f9441923d41408c74bd3761a34eb4db2c583fc08d5785a2c7f2260c8e17089e transformer/diffusion_pytorch_model-00002-of-00005.safetensors
89
+ 45cbd9a557c20c7ae5d676049b301f4856046be2f802f3e9e7de477c77a35b8c transformer/diffusion_pytorch_model-00003-of-00005.safetensors
90
+ 73c70763356a57d5d3cd04e91e714cb0b86081a3e922b594659b58f2ab49bc90 transformer/diffusion_pytorch_model-00004-of-00005.safetensors
91
+ 0c178789512c7f5d45a7b0c1c561f868704cd0fa77d4cb32f8d7ea06b2707781 transformer/diffusion_pytorch_model-00005-of-00005.safetensors
92
+ 21f362456d5f8394b516215521d7de8d44f8353e170043f82d902913eef5a2cc transformer/diffusion_pytorch_model.safetensors.index.json
93
+ 8db67ee46bb5810ed819d698a8acfc95936d157fe68b66c9ed16b2d5ea04413e vae/config.json
94
+ 06520463778e64dca1039c7447890065ee220bc408d15412182e5c3e06f304f1 vae/diffusion_pytorch_model.safetensors
code/.gitignore ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ __pycache__/
2
+ *.pyc
3
+ *.pyo
4
+ *.egg-info/
5
+ .venv/
6
+ venv/
7
+ env/
8
+ outputs/
9
+ cmd.log
10
+ .DS_Store
code/LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2026 inclusionAI
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
code/README.md ADDED
@@ -0,0 +1,279 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Ming Image 0.1 Design
2
+
3
+ Ming-Image-0.1-Design is an open-source series for visual-design generation
4
+ and editable layer decomposition.
5
+
6
+ The series includes two 6B-parameter models:
7
+
8
+ - [Ming-Image-0.1-Design](https://huggingface.co/inclusionAI/Ming-Image-0.1-Design)
9
+ generates complete visual designs for UI, infographics, posters, and
10
+ text-rich compositions.
11
+ - [Ming-Image-0.1-Design-Layer](https://huggingface.co/inclusionAI/Ming-Image-0.1-Design-Layer)
12
+ decomposes flattened design images into independently editable transparent
13
+ layers.
14
+
15
+ ![Ming-Image-0.1-Design on the Artificial Analysis UI/UX Design leaderboard](assets/ming-image-design-ui-ux-leaderboard.webp)
16
+
17
+ ## Gallery
18
+
19
+ ### Text-to-image
20
+
21
+ ![Text-to-image showcase](assets/model_cards/design_showcase.webp)
22
+
23
+ ### Transparent-background text-to-image
24
+
25
+ ![Transparent-background text-to-image showcase](assets/model_cards/design_transparency_showcase.webp)
26
+
27
+ ### Layer decomposition
28
+
29
+ ![Six-layer card decomposition showcase](assets/model_cards/layer_showcase.webp)
30
+
31
+ ![Layer-decomposition gallery](assets/model_cards/layer_gallery.webp)
32
+
33
+ ![Layer-decomposition results on the Crello test set](assets/model_cards/layer_performance.webp)
34
+
35
+ ## Requirements
36
+
37
+ - Python 3.10 or newer;
38
+ - CUDA-capable PyTorch for full model inference;
39
+ - a local checkpoint directory or a Hugging Face Hub repository ID.
40
+
41
+ Install the runtime dependencies in a clean environment:
42
+
43
+ ```bash
44
+ pip install -r requirements.txt
45
+ ```
46
+
47
+ `flash_attention_2` is optional. The CLI defaults to `eager` attention because
48
+ the LLM only implements eager and FlashAttention 2 attention classes; the
49
+ diffusion transformer always uses PyTorch SDPA internally. Selecting
50
+ `--attn-implementation sdpa` fails closed at load time.
51
+
52
+ ## Inference
53
+
54
+ The `--model` argument accepts either a local directory or an HF Hub repo ID.
55
+ Hub models are resolved to one immutable local snapshot before any component is
56
+ loaded. Use `--revision` to pin a branch, tag or commit. The examples use the
57
+ published Hub IDs; replace them with local checkpoint directories for offline
58
+ inference.
59
+
60
+ Sampling defaults and public resolution buckets are task-specific:
61
+
62
+ | Task | Steps | CFG | Resolution buckets | Default / recommended |
63
+ | --- | ---: | ---: | --- | ---: |
64
+ | Text-to-image | 12 | 1.0 | 1024, 2048 | 2048 |
65
+ | Layer decomposition | 12 | 2.0 | 512, 1024 | 1024 |
66
+
67
+ Pass `--steps` or `--cfg` to override either value explicitly.
68
+ `--resolution` is optional. A supplied positive integer snaps to the nearest
69
+ bucket supported by the selected task, with ties going to the smaller bucket.
70
+ For faster layer decomposition, explicitly pass `--resolution 512`; use 1024
71
+ for the recommended output quality. Text-to-image output is square at the
72
+ selected bucket. Layer decomposition preserves the reference image's aspect
73
+ ratio while selecting a predefined working size from the effective 1024 or
74
+ 512 bucket.
75
+
76
+ The default and minimum validated deployment is **one GPU with at least 80 GiB
77
+ of memory**, running in BF16. This configuration passes both model families
78
+ end-to-end, and all demos below use it by default.
79
+
80
+ ```bash
81
+ export CUDA_VISIBLE_DEVICES=0
82
+ ```
83
+
84
+ Select `--attn-implementation flash_attention_2` when FlashAttention 2 is
85
+ installed; the portable CLI default remains `eager`.
86
+
87
+ For both tasks, `--prompt` accepts either raw text or a path to a prompt file.
88
+ Use `--validate-only` to check the checkpoint contract and task combination
89
+ without loading model weights:
90
+
91
+ ```bash
92
+ python infer.py --model inclusionAI/Ming-Image-0.1-Design --task text-to-image \
93
+ --prompt "A red circle on a white background" --validate-only
94
+ ```
95
+
96
+ ### Text-to-image demo
97
+
98
+ This demo renders a fixed cabin across spring, summer, autumn, and winter. It
99
+ uses the [exact structured JSON prompt](assets/t2i_four_seasons_cabin_prompt.json)
100
+ without duplicating that long input in the README.
101
+
102
+ ```bash
103
+ python infer.py \
104
+ --model inclusionAI/Ming-Image-0.1-Design \
105
+ --task text-to-image \
106
+ --prompt assets/t2i_four_seasons_cabin_prompt.json \
107
+ --attn-implementation flash_attention_2 \
108
+ --resolution 2048 \
109
+ --output-dir outputs/t2i
110
+ ```
111
+
112
+ ### Text-to-image prompt rewriting
113
+
114
+ An instruction-following VLM can turn a short caption into a precise,
115
+ layout-structured JSON description: Figma-style layers ordered back to front,
116
+ with exact coordinates, hierarchy, color specs, and every rendered string
117
+ quoted verbatim and owned exactly once. Unlike layer decomposition, the
118
+ text-to-image rewriter describes the complete 1:1 canvas.
119
+
120
+ This rewrite is a pre-processing step *outside* `infer.py`. Prompt enhancement
121
+ (PE) can use `Ling-3.0-flash-VL` or `qwen3.8-27B`; run it first, then pass its
122
+ output to `--prompt` as raw text or via an asset file.
123
+
124
+ The released system prompt is stored in
125
+ [`assets/t2i_rewriter_system_prompt.txt`](assets/t2i_rewriter_system_prompt.txt)
126
+ and reproduced below:
127
+
128
+ ```text
129
+ You are a senior visual designer and image-prompt engineer. Expand the user's request into one precise, high-resolution Figma-style caption. Return only one JSON object.
130
+
131
+ Use exactly two top-level keys. `canvas_settings` contains exactly `aspect_ratio`, `ambient_lighting`, and `image_style`. `layers` lists visible groups from background to topmost overlay. Every layer contains exactly `description`, `coordinates`, `hierarchy_and_relation`, and `color_specs`; `color_specs` is an array of hex colors.
132
+
133
+ `coordinates` MUST be one string, never an object or array, in exactly this form: `"cx: 0.500, cy: 0.500, w: 1.000, h: 1.000"`. Values are normalized; each bbox encloses its complete owned object and stays inside the canvas.
134
+
135
+ A layer is one selectable visible semantic group: background, full person, coherent object, panel, card, row, or text block. Prefer the fewest groups that preserve the layout. Keep people and objects intact. Never create invisible parents, guides, placeholders, empty layers, duplicate summaries, or multiple owners for one element.
136
+
137
+ Preserve every user-supplied rendered string character-for-character and as one contiguous string. Unless multiple visible copies are requested, it must occur exactly once across all `description` fields and zero times in `hierarchy_and_relation`. Quote it only where describing its visible rendering; refer to the related subject elsewhere with unquoted semantic wording. Enumerate intended copy, invent extra copy sparingly, and never hide content behind "other text", "remaining labels", or "etc."
138
+
139
+ Describe concrete composition, typography, materials, texture, lighting, pose, and camera treatment without literary filler. Use `hierarchy_and_relation` only for ownership, alignment, containment, stacking, and occlusion.
140
+
141
+ Infer structured layouts first. Use one complete layer per card and state its row and column. A compact secondary table may be one layer only if every header and cell is listed; otherwise use a visible shared frame when present, one complete header, and one complete layer per body row, binding values to columns and stating blanks. Enumerate sequences, schedules, spans, gaps, and vacant tracks in visual order. Do not mistake ordinary alignment for a table.
142
+
143
+ Silently verify schema, string coordinates, Z-order, exact-text counts, geometry, bbox validity, and completeness.
144
+ ```
145
+
146
+ ### Transparent-background generation tip
147
+
148
+ To generate an image with an alpha channel, choose exactly one of the following
149
+ fixed phrases and place it at the beginning of the prompt. Do not combine
150
+ multiple prefixes.
151
+
152
+ - `带透明通道,4通道RGBA图像`
153
+ - `透明背景,alpha通道,无底图`
154
+ - `抠图素材,背景alpha=0`
155
+ - `孤立主体,透明PNG图层`
156
+ - `不要白底,不要棋盘格,只要透明通道`
157
+ - `RGBA, 4-channel, transparent background`
158
+ - `isolated subject, alpha matte, no background`
159
+ - `cutout PNG, alpha=0 outside the object`
160
+ - `transparent canvas, not white, not checkerboard`
161
+ - `production RGBA layer for compositing`
162
+
163
+ ### Layer-decomposition demo
164
+
165
+ This example separates the flattened card design into six transparent RGBA
166
+ layers. See the [source input](assets/layer_samples/card_making_input.png) and
167
+ the [exact six-layer specification](assets/layer_samples/card_making_prompt.txt).
168
+
169
+ ```bash
170
+ python infer.py \
171
+ --model inclusionAI/Ming-Image-0.1-Design-Layer \
172
+ --task layer-decompose \
173
+ --input-image assets/layer_samples/card_making_input.png \
174
+ --prompt assets/layer_samples/card_making_prompt.txt \
175
+ --attn-implementation flash_attention_2 \
176
+ --resolution 1024 \
177
+ --output-dir outputs/layers
178
+ ```
179
+
180
+ The layer count is parsed from a "Decompose this image into N layers" or
181
+ "Number of layers: N" specification in the prompt. When `--prompt` is omitted,
182
+ `--num-layers N` generates the default request `Decompose this image into N
183
+ layers.`
184
+
185
+ The layer model returns the requested layers plus one leading
186
+ composite/full-canvas image. The CLI skips that first image and saves the
187
+ standalone layers as `layer_01.png`, `layer_02.png`, and so on.
188
+
189
+ ### Layer-decomposition prompt rewriting
190
+
191
+ Layer decomposition is driven by an explicit per-layer specification, not a
192
+ free-form caption. The reference pipeline first runs a prompt enhancer: an
193
+ instruction-following VLM rewrites the user's rough layer plan into a precise
194
+ decomposition — concrete colors/positions/shapes per layer, real text quoted
195
+ verbatim on its own front layer, a text-supporting card/panel/badge/banner as
196
+ its own layer directly behind the text, the main subject as its own layer, and
197
+ the background last absorbing the remaining supports and surfaces.
198
+
199
+ This rewrite is a pre-processing step *outside* `infer.py`: run the enhancer
200
+ first, then pass its output to `--prompt` (as raw text or via an asset file).
201
+ The enhancer's output is exactly the format this CLI parses — it regenerates the
202
+ "Decompose this image into N layers" / "Number of layers: N" spec, so the layer
203
+ count flows through the same `--prompt` parsing path described above.
204
+
205
+ Prompt enhancement (PE) can use `Ling-3.0-flash-VL` or `qwen3.8-27B` as the
206
+ instruction-following VLM. The guided prompt below is part of the released
207
+ pipeline and is kept here for reproducibility:
208
+
209
+ ```text
210
+ GUIDED_PROMPT = """You are a graphic-design layer-decomposition expert. You are given ONE flattened design image and a ROUGH layer plan from the user. Rewrite the rough plan into a precise layer decomposition that matches the image.
211
+
212
+ User's rough layer plan:
213
+ {spec}
214
+
215
+ Guidelines:
216
+ - Follow the user's plan EXACTLY: use the same number of layers and the same per-layer role/meaning, in the same order. Layer 1 is the FRONT-most (topmost); the last layer is the background/environment. Stacking the layers back-to-front must reproduce the image.
217
+ - For each layer, write a concrete one-or-two-sentence description grounded in the image: real colors, positions and shapes.
218
+ - TEXT goes in the front layer(s); quote any real text VERBATIM in double quotes and keep its original language (do not translate).
219
+ - A text-supporting CARD / PANEL / BADGE / BANNER is its OWN layer directly behind the text — do not merge it into the text layer or into the background.
220
+ - The MAIN SUBJECT (hero product/photo/illustration) is its own layer.
221
+ - The BACKGROUND/ENVIRONMENT is the LAST layer and ABSORBS supporting props and surfaces under the subject (tables, boards, plates, floors, shadows, gradients, patterns) — these are NOT separate layers.
222
+
223
+ Write the description DIRECTLY about the content; do NOT mention "image" or narrate your reasoning. Output EXACTLY in this format and NOTHING else (N = the number of layers in the user's plan):
224
+
225
+ Decompose this image into N layers with the following specifications:
226
+
227
+ Number of layers: N
228
+ Layer 1: <front-most layer>
229
+ Layer 2: <...>
230
+ Layer N: <background/environment layer>"""
231
+ ```
232
+
233
+ ## Deployment
234
+
235
+ We recommend the following inference frameworks to serve the model:
236
+
237
+ - vLLM-Omni: see the [recipes](https://github.com/vllm-project/vllm-omni/blob/main/recipes/inclusionAI/Ming-Image.md)
238
+ and [installation guide](https://docs.vllm.ai/projects/vllm-omni/en/latest/getting_started/quickstart/).
239
+
240
+ ## Verification
241
+
242
+ Fast contract tests do not require model weights:
243
+
244
+ ```bash
245
+ python -m unittest -v tests.test_inference_profile
246
+ python -m unittest -v tests.test_padding
247
+ ```
248
+
249
+ Full inference is a GPU smoke test and requires both checkpoint families. At a
250
+ minimum, validate the four-seasons-cabin text-to-image case and the six-layer
251
+ card-making decomposition case with fixed seeds. The smoke test defaults to
252
+ the single-GPU placement and FlashAttention 2 configuration shown above:
253
+
254
+ ```bash
255
+ export CUDA_VISIBLE_DEVICES=0
256
+ export MING_GENERATION_MODEL=inclusionAI/Ming-Image-0.1-Design
257
+ export MING_LAYER_MODEL=inclusionAI/Ming-Image-0.1-Design-Layer
258
+ export MING_SMOKE_OUTPUT_DIR=/path/to/persistent/results
259
+ python -m unittest -v tests.test_inference_smoke.InferenceSmokeTest.test_two_step_showcase_smoke
260
+ ```
261
+
262
+ ## Checkpoint contract
263
+
264
+ Each checkpoint declares its runtime capability in `transformer/config.json`:
265
+
266
+ | Family | `alignment_padding_mode` | `multi_frame_output` |
267
+ | --- | --- | ---: |
268
+ | Text-to-image | `"zero_masked"` | `false` |
269
+ | Layer decomposition | `"learned"` | `true` |
270
+
271
+ Both fields must be present together, and the VAE component must be the
272
+ 4-channel `AutoencoderKLQwenImage` (argmax reference encoding). Missing,
273
+ partial, or unknown values are errors; the loader never infers padding
274
+ behavior from a directory name or silently falls back to another mode.
275
+
276
+ Legacy packages that predate the component metadata are loaded strictly from
277
+ a root `inference_profile.json` during the compatibility window; when both
278
+ exist, the component configs are authoritative and any disagreement with the
279
+ legacy file is an error.
code/REVISION ADDED
@@ -0,0 +1 @@
 
 
1
+ branch rocm-halo @ f986f7a (upstream inclusionAI/Ming-Image 62c6072 + 3 commits); rocm.patch sha256 03fe16f8ed566bf99ec654caa687c7e659f489df6d077cfb7fc7aee134ebc777
code/bailingmm_utils.py ADDED
@@ -0,0 +1,504 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import logging
3
+ import math
4
+ import os
5
+ from io import BytesIO
6
+ from tqdm.contrib.concurrent import thread_map
7
+
8
+ import numpy as np
9
+
10
+ import requests
11
+ import torch
12
+
13
+ from PIL import Image
14
+ from typing import Union, Tuple, List
15
+
16
+ def fetch_video(ele):
17
+ raise ValueError(
18
+ "video input is not supported by Ming Image inference; "
19
+ "the public API accepts image and text input only"
20
+ )
21
+
22
+ logger = logging.getLogger(__name__)
23
+
24
+ IMAGE_FACTOR = 28
25
+ MIN_PIXELS = 4 * 28 * 28
26
+ MAX_PIXELS = 1024 * 28 * 28
27
+ MAX_RATIO = 200
28
+
29
+ VideoInput = Union[
30
+ List["Image.Image"],
31
+ "np.ndarray",
32
+ "torch.Tensor",
33
+ List["np.ndarray"],
34
+ List["torch.Tensor"],
35
+ List[List["Image.Image"]],
36
+ List[List["np.ndarray"]],
37
+ List[List["torch.Tensor"]],
38
+ ]
39
+
40
+
41
+ def round_by_factor(number: int, factor: int) -> int:
42
+ """Returns the closest integer to 'number' that is divisible by 'factor'."""
43
+ return round(number / factor) * factor
44
+
45
+ def ceil_by_factor(number: int, factor: int) -> int:
46
+ """Returns the smallest integer greater than or equal to 'number' that is divisible by 'factor'."""
47
+ return math.ceil(number / factor) * factor
48
+
49
+ def floor_by_factor(number: int, factor: int) -> int:
50
+ """Returns the largest integer less than or equal to 'number' that is divisible by 'factor'."""
51
+ return math.floor(number / factor) * factor
52
+
53
+ def is_image(image_file):
54
+ if isinstance(image_file, str) and (image_file.startswith("base64,") or image_file.lower().endswith(
55
+ ('.bmp', '.dib', '.png', '.jpg', '.jpeg', '.pbm', '.pgm', '.ppm', '.tif', '.tiff'))):
56
+ return True
57
+ elif isinstance(image_file, Image.Image):
58
+ return True
59
+ else:
60
+ return False
61
+
62
+ def is_video(video_file):
63
+ if isinstance(video_file, str) and video_file.lower().endswith(
64
+ ('.mp4', '.mkv', '.avi', '.wmv', '.iso', ".webm")):
65
+ return True
66
+ else:
67
+ return False
68
+
69
+ def is_audio(audio_file):
70
+ if isinstance(audio_file, str) and audio_file.lower().endswith(
71
+ (".wav", ".mp3", ".aac", ".flac", ".alac", ".m4a", ".ogg", ".wma", ".aiff", ".amr", ".au")):
72
+ return True
73
+ else:
74
+ return False
75
+
76
+ def smart_resize(
77
+ height: int, width: int, factor: int = IMAGE_FACTOR, min_pixels: int = MIN_PIXELS, max_pixels: int = MAX_PIXELS
78
+ ) -> tuple[int, int]:
79
+ """
80
+ Rescales the image so that the following conditions are met:
81
+
82
+ 1. Both dimensions (height and width) are divisible by 'factor'.
83
+
84
+ 2. The total number of pixels is within the range ['min_pixels', 'max_pixels'].
85
+
86
+ 3. The aspect ratio of the image is maintained as closely as possible.
87
+ """
88
+ if max(height, width) / min(height, width) > MAX_RATIO:
89
+ raise ValueError(
90
+ f"absolute aspect ratio must be smaller than {MAX_RATIO}, got {max(height, width) / min(height, width)}"
91
+ )
92
+ h_bar = max(factor, round_by_factor(height, factor))
93
+ w_bar = max(factor, round_by_factor(width, factor))
94
+ if h_bar * w_bar > max_pixels:
95
+ beta = math.sqrt((height * width) / max_pixels)
96
+ h_bar = floor_by_factor(height / beta, factor)
97
+ w_bar = floor_by_factor(width / beta, factor)
98
+ elif h_bar * w_bar < min_pixels:
99
+ beta = math.sqrt(min_pixels / (height * width))
100
+ h_bar = ceil_by_factor(height * beta, factor)
101
+ w_bar = ceil_by_factor(width * beta, factor)
102
+ return h_bar, w_bar
103
+
104
+ def fetch_image(ele: dict[str, str | Image.Image], size_factor: int = IMAGE_FACTOR) -> Image.Image:
105
+ if "image" in ele:
106
+ image = ele["image"]
107
+ else:
108
+ image = ele["image_url"]
109
+ image_obj = None
110
+ if isinstance(image, Image.Image):
111
+ image_obj = image
112
+ elif image.startswith("http://") or image.startswith("https://"):
113
+ image_obj = Image.open(requests.get(image, stream=True).raw)
114
+ elif image.startswith("file://"):
115
+ image_obj = Image.open(image[7:])
116
+ elif image.startswith("data:image"):
117
+ if "base64," in image:
118
+ _, base64_data = image.split("base64,", 1)
119
+ data = base64.b64decode(base64_data)
120
+ image_obj = Image.open(BytesIO(data))
121
+ else:
122
+ image_obj = Image.open(image)
123
+ if image_obj is None:
124
+ raise ValueError(f"Unrecognized image input, support local path, http url, base64 and PIL.Image, got {image}")
125
+ image = image_obj.convert("RGB")
126
+ ## resize
127
+ if "resized_height" in ele and "resized_width" in ele:
128
+ resized_height, resized_width = smart_resize(
129
+ ele["resized_height"],
130
+ ele["resized_width"],
131
+ factor=size_factor,
132
+ )
133
+ else:
134
+ width, height = image.size
135
+ min_pixels = ele.get("min_pixels", MIN_PIXELS)
136
+ max_pixels = ele.get("max_pixels", MAX_PIXELS)
137
+ resized_height, resized_width = smart_resize(
138
+ height,
139
+ width,
140
+ factor=size_factor,
141
+ min_pixels=min_pixels,
142
+ max_pixels=max_pixels,
143
+ )
144
+ image = image.resize((resized_width, resized_height))
145
+
146
+ return image
147
+
148
+ def fetch_image_wo_resize(ele: dict[str, str | Image.Image], size_factor: int = IMAGE_FACTOR) -> Image.Image:
149
+ if "image" in ele:
150
+ image = ele["image"]
151
+ else:
152
+ image = ele["image_url"]
153
+ image_obj = None
154
+ if isinstance(image, Image.Image):
155
+ image_obj = image
156
+ elif image.startswith("http://") or image.startswith("https://"):
157
+ image_obj = Image.open(requests.get(image, stream=True).raw)
158
+ elif image.startswith("file://"):
159
+ image_obj = Image.open(image[7:])
160
+ elif image.startswith("data:image"):
161
+ if "base64," in image:
162
+ _, base64_data = image.split("base64,", 1)
163
+ data = base64.b64decode(base64_data)
164
+ image_obj = Image.open(BytesIO(data))
165
+ else:
166
+ image_obj = Image.open(image)
167
+ if image_obj is None:
168
+ raise ValueError(f"Unrecognized image input, support local path, http url, base64 and PIL.Image, got {image}")
169
+
170
+ #image = image_obj.convert("RGB")
171
+
172
+ return image_obj
173
+
174
+ def fetch_audio(ele: dict[str, str | torch.Tensor], return_tensor="pt") -> Tuple[Union[torch.Tensor, np.ndarray], int]:
175
+ import torchaudio
176
+
177
+ if "audio" in ele:
178
+ audio = ele["audio"]
179
+ else:
180
+ audio = ele["audio_url"]
181
+
182
+ if isinstance(audio, torch.Tensor):
183
+ waveform = audio
184
+ sample_rate: int = ele.get("sample_rate", 16000)
185
+ elif audio.startswith("http://") or audio.startswith("https://"):
186
+ audio_file = BytesIO(requests.get(audio, stream=True).content)
187
+ waveform, sample_rate = torchaudio.load(audio_file)
188
+ elif audio.startswith("file://"):
189
+ waveform, sample_rate = torchaudio.load(audio[7:])
190
+ else:
191
+ waveform, sample_rate = torchaudio.load(audio)
192
+ if return_tensor == "pt":
193
+ return waveform, sample_rate
194
+ else:
195
+ return waveform.numpy(), sample_rate
196
+
197
+ def extract_vision_info(conversations: list[dict] | list[list[dict]]) -> list[dict]:
198
+ vision_infos = []
199
+ if isinstance(conversations[0], dict):
200
+ conversations = [conversations]
201
+ for conversation in conversations:
202
+ for message in conversation:
203
+ if isinstance(message["content"], list):
204
+ for ele in message["content"]:
205
+ if (
206
+ "image" in ele
207
+ or "image_url" in ele
208
+ or "video" in ele
209
+ or "video_url" in ele
210
+ or "audio" in ele
211
+ or "audio_url" in ele
212
+ or ele["type"] in ["image", "image_url", "video", "video_url", "audio", "audio_url"]
213
+ ):
214
+ vision_infos.append(ele)
215
+ return vision_infos
216
+
217
+
218
+ def process_reference_vision_info(
219
+ conversations: list[dict] | list[list[dict]],
220
+ ) -> list[Image.Image] | None:
221
+ vision_infos = extract_vision_info(conversations)
222
+ ## Read images
223
+ image_inputs = []
224
+
225
+ def inner_process_func(vision_info):
226
+ if "image" in vision_info or "image_url" in vision_info:
227
+ res_list = []
228
+ if "image" in vision_info and isinstance(vision_info["image"], (tuple, list)):
229
+ for i in range(len(vision_info["image"])):
230
+ res_list.append(fetch_image_wo_resize({"type": "image", "image": vision_info["image"][i]}))
231
+ elif "image_url" in vision_info and vision_info["image_url"].get("url", None) is not None:
232
+ vision_info["image_url"] = vision_info["image_url"].get("url")
233
+ res_list.extend([fetch_image_wo_resize(vision_info)])
234
+ else:
235
+ res_list.extend([fetch_image_wo_resize(vision_info)])
236
+ return {'image_inputs':res_list}
237
+ else:
238
+ return None
239
+
240
+ vision_infos_reslist = thread_map(inner_process_func, vision_infos, disable=True)
241
+ for res in vision_infos_reslist:
242
+ if res is None:
243
+ raise ValueError("image, image_url, video, video_url, audio or audio_url should in content.")
244
+ elif 'image_inputs' in res:
245
+ image_inputs.extend(res['image_inputs'])
246
+
247
+ if len(image_inputs) > 1: # multi-image input keeps only the first image as the VAE reference
248
+ image_inputs = [image_inputs[0]]
249
+
250
+ if len(image_inputs) == 0:
251
+ image_inputs = None
252
+
253
+ return image_inputs
254
+
255
+
256
+ def process_vision_info(
257
+ conversations: list[dict] | list[list[dict]],
258
+ ) -> tuple[list[Image.Image] | None, list[torch.Tensor | list[Image.Image]] | None, list[
259
+ torch.Tensor | list[np.ndarray]] | None]:
260
+ vision_infos = extract_vision_info(conversations)
261
+ ## Read images, videos or audios
262
+ image_inputs = []
263
+ video_inputs = []
264
+ audio_inputs = []
265
+
266
+ def inner_process_func(vision_info):
267
+ if "image" in vision_info or "image_url" in vision_info:
268
+ res_list = []
269
+ if "image" in vision_info and isinstance(vision_info["image"], (tuple, list)):
270
+ for i in range(len(vision_info["image"])):
271
+ res_list.append(fetch_image({"type": "image", "image": vision_info["image"][i]}))
272
+ elif "image_url" in vision_info and vision_info["image_url"].get("url", None) is not None:
273
+ vision_info["image_url"] = vision_info["image_url"].get("url")
274
+ res_list.extend([fetch_image(vision_info)])
275
+ else:
276
+ res_list.extend([fetch_image(vision_info)])
277
+ return {'image_inputs':res_list}
278
+
279
+ elif "video" in vision_info or "video_url" in vision_info:
280
+ if "video_url" in vision_info and vision_info["video_url"].get("url", None) is not None:
281
+ data_value = vision_info["video_url"].get("url")
282
+ elif "video" in vision_info and not os.path.isdir(vision_info['video']):
283
+ data_value = vision_info['video']
284
+ else:
285
+ data_value = [os.path.join(vision_info['video'], frame) for frame in sorted(os.listdir(vision_info['video']))]
286
+ vision_info['video']=data_value
287
+ return {"video_inputs": [fetch_video(vision_info)]}
288
+
289
+ elif "audio" in vision_info or "audio_url" in vision_info:
290
+ if "audio" in vision_info and isinstance(vision_info["audio"], (tuple, list)):
291
+ return {"audio_inputs":[fetch_audio(info) for info in vision_info["audio"]]}
292
+ elif "audio_url" in vision_info and vision_info["audio_url"].get("url", None) is not None:
293
+ vision_info["audio_url"] = vision_info["audio_url"].get("url")
294
+ return {"audio_inputs":[fetch_audio(vision_info)]}
295
+ else:
296
+ return {"audio_inputs":[fetch_audio(vision_info)]}
297
+ else:
298
+ return None
299
+
300
+ vision_infos_reslist = thread_map(inner_process_func, vision_infos, disable=True)
301
+ for res in vision_infos_reslist:
302
+ if res is None:
303
+ raise ValueError("image, image_url, video, video_url, audio or audio_url should in content.")
304
+ elif 'image_inputs' in res:
305
+ image_inputs.extend(res['image_inputs'])
306
+ elif 'video_inputs' in res:
307
+ video_inputs.extend(res['video_inputs'])
308
+ elif 'audio_inputs' in res:
309
+ audio_inputs.extend(res['audio_inputs'])
310
+
311
+ if len(image_inputs) == 0:
312
+ image_inputs = None
313
+ if len(video_inputs) == 0:
314
+ video_inputs = None
315
+ if len(audio_inputs) == 0:
316
+ audio_inputs = None
317
+ return image_inputs, video_inputs, audio_inputs
318
+
319
+
320
+ def get_closest_ratio(height: float, width: float, aspect_ratios: dict):
321
+ aspect_ratio = height / width
322
+ closest_ratio = min(aspect_ratios.keys(), key=lambda ratio: abs(float(ratio) - aspect_ratio))
323
+ return aspect_ratios[closest_ratio], float(closest_ratio)
324
+
325
+ def process_ratio(ori_h, ori_w, highres=512):
326
+ ASPECT_RATIO_512 = {
327
+ "0.25": [256, 1024], "0.26": [256, 992], "0.27": [256, 960], "0.28": [256, 928],
328
+ "0.32": [288, 896], "0.33": [288, 864], "0.35": [288, 832], "0.4": [320, 800],
329
+ "0.42": [320, 768], "0.48": [352, 736], "0.5": [352, 704], "0.52": [352, 672],
330
+ "0.5455": [384, 704], "0.57": [384, 672], "0.6": [384, 640], "0.65": [416, 640],
331
+ "0.68": [416, 608], "0.72": [416, 576], "0.78": [448, 576],
332
+ "0.82": [448, 544], "0.88": [480, 544], "0.94": [480, 512],
333
+ "1.0": [512, 512], "1.07": [512, 480], "1.13": [544, 480], "1.21": [544, 448],
334
+ "1.29": [576, 448], "1.38": [576, 416],
335
+ "1.46": [608, 416], "1.5385": [640, 416], "1.67": [640, 384], "1.75": [672, 384],
336
+ "1.8333": [704, 384], "2.0": [704, 352], "2.09": [736, 352], "2.4": [768, 320],
337
+ "2.5": [800, 320], "2.89": [832, 288], "3.0": [864, 288], "3.11": [896, 288],
338
+ "3.62": [928, 256], "3.75": [960, 256], "3.88": [992, 256], "4.0": [1024, 256],
339
+ }
340
+ ASPECT_RATIO_1024 = {
341
+ "0.25": [512, 2048], "0.26": [512, 1984], "0.27": [512, 1920], "0.28": [512, 1856],
342
+ "0.32": [576, 1792], "0.33": [576, 1728], "0.35": [576, 1664], "0.4": [640, 1600],
343
+ "0.42": [640, 1536], "0.48": [704, 1472], "0.5": [704, 1408], "0.52": [704, 1344],
344
+ "0.5581": [768, 1376], "0.5625": [720, 1280], "0.5647": [768, 1360], "0.57": [768, 1344],
345
+ "0.6": [768, 1280], "0.622": [816, 1312], "0.625": [800, 1280], "0.65": [832, 1280],
346
+ "0.6582": [832, 1264], "0.6667": [832, 1248], "0.6709": [848, 1264], "0.68": [832, 1216],
347
+ "0.7013": [864, 1232], "0.72": [832, 1152], "0.7467": [896, 1200], "0.75": [864, 1152],
348
+ "0.7568": [896, 1184], "0.78": [896, 1152], "0.8": [896, 1120], "0.8056": [928, 1152],
349
+ "0.82": [896, 1088], "0.88": [960, 1088], "0.94": [960, 1024], "0.9846": [1024, 1040],
350
+ "1.0": [1024, 1024], "1.07": [1024, 960], "1.13": [1088, 960], "1.21": [1088, 896],
351
+ "1.2414": [1152, 928], "1.25": [1120, 896], "1.2807": [1168, 912], "1.29": [1152, 896],
352
+ "1.3333": [1152, 864], "1.3393": [1200, 896], "1.38": [1152, 832], "1.46": [1216, 832],
353
+ "1.4906": [1264, 848], "1.5": [1248, 832], "1.6": [1280, 800], "1.67": [1280, 768],
354
+ "1.75": [1344, 768], "1.7708": [1360, 768], "1.7778": [1280, 720], "2.0": [1408, 704],
355
+ "2.09": [1472, 704], "2.4": [1536, 640], "2.5": [1600, 640], "2.89": [1664, 576],
356
+ "3.0": [1728, 576], "3.11": [1792, 576], "3.62": [1856, 512], "3.75": [1920, 512],
357
+ "3.88": [1984, 512], "4.0": [2048, 512],
358
+ }
359
+
360
+ ASPECT_RATIO_672 = {
361
+ "0.28": [352, 1280], "0.32": [384, 1184], "0.38": [416, 1088], "0.44": [448, 1024],
362
+ "0.52": [480, 928], "0.5636": [496, 880], "0.57": [512, 896], "0.65": [544, 832],
363
+ "0.6667": [544, 816], "0.75": [576, 768], "0.8": [576, 720], "0.83": [608, 736],
364
+ "0.91": [640, 704], "1.00": [672, 672], "1.10": [704, 640], "1.21": [736, 608],
365
+ "1.25": [720, 576], "1.33": [768, 576], "1.39": [800, 576], "1.5": [816, 544],
366
+ "1.53": [832, 544], "1.69": [864, 512], "1.75": [896, 512], "1.7742": [880, 496],
367
+ "1.93": [928, 480], "2.00": [960, 480],
368
+ "2.21": [992, 448], "2.29": [1024, 448], "2.54": [1056, 416], "2.62": [1088, 416],
369
+ "2.69": [1120, 416], "3.00": [1152, 384], "3.08": [1184, 384], "3.17": [1216, 384],
370
+ "3.55": [1248, 352], "3.64": [1280, 352],
371
+ }
372
+
373
+ ASPECT_RATIO_2048 = {
374
+ "0.25": [1024, 4096], "0.26": [1024, 3968], "0.27": [1024, 3840], "0.28": [1024, 3712],
375
+ "0.32": [1152, 3584], "0.33": [1152, 3456], "0.35": [1152, 3328], "0.4": [1280, 3200],
376
+ "0.42": [1280, 3072], "0.48": [1408, 2944], "0.5": [1408, 2816], "0.52": [1408, 2688],
377
+ "0.5625": [1440, 2560], "0.57": [1536, 2688], "0.6": [1536, 2560], "0.6667": [1664, 2496],
378
+ "0.68": [1664, 2432], "0.72": [1664, 2304], "0.75": [1824, 2432], "0.78": [1792, 2304],
379
+ "0.7917": [1824, 2304], "0.8": [1792, 2240], "0.82": [1792, 2176], "0.88": [1920, 2176],
380
+ "0.94": [1920, 2048],
381
+ "1.0": [2048, 2048], "1.07": [2048, 1920], "1.13": [2176, 1920], "1.21": [2176, 1792],
382
+ "1.25": [2240, 1792], "1.2632": [2304, 1824], "1.29": [2304, 1792],
383
+ "1.3333": [2432, 1824], "1.38": [2304, 1664],
384
+ "1.46": [2432, 1664], "1.5": [2496, 1664], "1.67": [2560, 1536], "1.75": [2688, 1536],
385
+ "1.7778": [2560, 1440], "2.0": [2816, 1408], "2.09": [2944, 1408], "2.4": [3072, 1280],
386
+ "2.5": [3200, 1280], "2.89": [3328, 1152], "3.0": [3456, 1152], "3.11": [3584, 1152],
387
+ "3.62": [3712, 1024], "3.75": [3840, 1024], "3.88": [3968, 1024], "4.0": [4096, 1024],
388
+ "0.2941": [1120, 3808], "0.3043": [1120, 3680], "0.3679": [1248, 3392],
389
+ "0.3846": [1280, 3328], "0.433": [1344, 3104], "0.4574": [1376, 3008],
390
+ "0.4681": [1408, 3008], "0.5465": [1504, 2752], "0.6296": [1632, 2592],
391
+ "0.6582": [1664, 2528], "0.7432": [1760, 2368], "0.8551": [1888, 2208],
392
+ "0.9692": [2016, 2080], "1.0317": [2080, 2016], "1.1695": [2208, 1888],
393
+ "1.3455": [2368, 1760], "1.5192": [2528, 1664], "1.5882": [2592, 1632],
394
+ "1.8298": [2752, 1504], "1.913": [2816, 1472], "2.1364": [3008, 1408],
395
+ "2.186": [3008, 1376], "2.3095": [3104, 1344], "2.6": [3328, 1280],
396
+ "2.7179": [3392, 1248], "3.2857": [3680, 1120], "3.4": [3808, 1120],
397
+ }
398
+
399
+ aspect_ratio_dict = {
400
+ 512 : ASPECT_RATIO_512,
401
+ 672 : ASPECT_RATIO_672,
402
+ 1024 : ASPECT_RATIO_1024,
403
+ 2048 : ASPECT_RATIO_2048,
404
+ }
405
+
406
+ if highres is None or highres is False:
407
+ highres = 512
408
+ elif highres is True:
409
+ highres = 1024
410
+
411
+ aspect_ratio = aspect_ratio_dict[min([i for i in aspect_ratio_dict], key=lambda x: abs(x - highres))]
412
+
413
+ closest_size, _ = get_closest_ratio(ori_h, ori_w, aspect_ratios=aspect_ratio)
414
+ closest_size = list(map(lambda x: int(x), closest_size))
415
+ if closest_size[0] / ori_h > closest_size[1] / ori_w:
416
+ resize_size = closest_size[0], int(ori_w * closest_size[0] / ori_h)
417
+ else:
418
+ resize_size = int(ori_h * closest_size[1] / ori_w), closest_size[1]
419
+ return closest_size, resize_size
420
+
421
+
422
+ def find_first_index_of_consecutive_ones(lst):
423
+ """
424
+ Given a list of 0s and 1s, return the index of the first 1 of each
425
+ consecutive run of 1s.
426
+
427
+ Args:
428
+ lst (list): list of 0s and 1s
429
+
430
+ Returns:
431
+ list: indices of the first 1 of each consecutive run of 1s
432
+ """
433
+ result = []
434
+ i = 0
435
+ n = len(lst)
436
+
437
+ while i < n:
438
+ if lst[i] == 1:
439
+ # find the start of a consecutive run of 1s
440
+ result.append(i)
441
+ # skip the remainder of the consecutive run of 1s
442
+ while i < n and lst[i] == 1:
443
+ i += 1
444
+ else:
445
+ i += 1
446
+
447
+ return result
448
+
449
+ def merge_consecutive_ones(lst, n):
450
+ """
451
+ Given a list of 0s and 1s, merge every n consecutive 1s of each run
452
+ (length >= 1) into a single 1. Each run must have a length divisible by n.
453
+ The relative order of 0s and 1s is preserved.
454
+
455
+ Args:
456
+ lst: list of 0s and 1s
457
+ n: positive integer merge unit size
458
+
459
+ Returns:
460
+ list: the merged list
461
+ """
462
+ assert isinstance(lst, list), "input must be a list"
463
+ assert isinstance(n, int) and n > 0, "n must be a positive integer"
464
+
465
+ # iterate over the list, extract runs of 1s, and verify each run length is divisible by n
466
+ i = 0
467
+ while i < len(lst):
468
+ if lst[i] == 1:
469
+ count = 0
470
+ start = i
471
+ # count the run of consecutive 1s
472
+ while i < len(lst) and lst[i] == 1:
473
+ count += 1
474
+ i += 1
475
+ # every run of 1s must be divisible by n
476
+ assert count % n == 0, f"run of 1s at index {start} has length {count}, not divisible by n={n}"
477
+ else:
478
+ i += 1
479
+
480
+ # build the new list by merging groups
481
+ result = []
482
+ i = 0
483
+ while i < len(lst):
484
+ if lst[i] == 0:
485
+ result.append(0)
486
+ i += 1
487
+ else:
488
+ # process a run of 1s
489
+ count = 0
490
+ while i < len(lst) and lst[i] == 1:
491
+ count += 1
492
+ i += 1
493
+ # merge every n 1s into a single 1
494
+ result.extend([1] * (count // n))
495
+
496
+ return result
497
+
498
+ def get_default_image_gen_hw(image_gen_highres, image_gen_aspect_ratio):
499
+ if image_gen_aspect_ratio is None:
500
+ image_gen_aspect_ratio = 1.0
501
+
502
+ closest_size, _ = process_ratio(ori_h=512, ori_w=int(512.0 * image_gen_aspect_ratio), highres=image_gen_highres)
503
+ h, w = closest_size
504
+ return h, w
code/configuration_bailing_moe_v2.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Bailing MoE model configuration"""
2
+ from transformers.configuration_utils import PretrainedConfig
3
+
4
+ class BailingMoeV2Config(PretrainedConfig):
5
+ model_type = "bailing_moe_v2"
6
+ def __init__(
7
+ self,
8
+ vocab_size=30592,
9
+ hidden_size=1024,
10
+ intermediate_size=None,
11
+ num_hidden_layers=24,
12
+ num_attention_heads=16,
13
+ num_key_value_heads=0,
14
+ hidden_act="silu",
15
+ use_qkv_bias=False, # bailing only
16
+ use_qk_norm=False,
17
+ use_bias=True, # bailing only
18
+ rms_norm_eps=1e-05,
19
+ norm_head=False, # bailing only
20
+ tie_word_embeddings=False, # PretrainedConfig key, here change default value.
21
+ embedding_dropout=0.1,
22
+ attention_dropout=0.1,
23
+ output_dropout=0.1,
24
+ initializer_range=0.02,
25
+ max_position_embeddings=16384,
26
+ rope_theta=10000.0,
27
+ use_cache=True,
28
+ use_sliding_window=False,
29
+ sliding_window=4096,
30
+ max_window_layers=28,
31
+ rope_scaling=None,
32
+ pad_token_id=126081,
33
+ num_experts=16,
34
+ num_shared_experts=0,
35
+ num_experts_per_tok=2,
36
+ n_group=8,
37
+ topk_group=4,
38
+ routed_scaling_factor=2.5,
39
+ moe_intermediate_size=None,
40
+ first_k_dense_replace=0,
41
+ head_dim=None,
42
+ output_router_logits=False,
43
+ partial_rotary_factor=0.5,
44
+ router_type="topN",
45
+ _attn_implementation="flash_attention_2",
46
+ **kwargs,
47
+ ):
48
+ self.num_hidden_layers = num_hidden_layers
49
+ self.vocab_size = vocab_size
50
+ self.hidden_size = hidden_size
51
+ self.intermediate_size = intermediate_size
52
+ self.num_attention_heads = num_attention_heads
53
+ self.num_key_value_heads = num_key_value_heads
54
+ self.hidden_act = hidden_act
55
+ self.use_qkv_bias = use_qkv_bias
56
+ self.use_bias = use_bias
57
+ self.norm_head = norm_head
58
+ self.rms_norm_eps = rms_norm_eps
59
+ self.embedding_dropout = embedding_dropout
60
+ self.attention_dropout = attention_dropout
61
+ self.output_dropout = output_dropout
62
+ self.initializer_range = initializer_range
63
+ self.max_position_embeddings = max_position_embeddings
64
+ self.rope_theta = rope_theta
65
+ self.use_cache = use_cache
66
+ self.use_sliding_window = use_sliding_window
67
+ self.sliding_window = sliding_window
68
+ self.max_window_layers = max_window_layers
69
+ self.head_dim = head_dim or self.hidden_size // self.num_attention_heads
70
+ self.rope_scaling = rope_scaling
71
+ # MoE configs
72
+ self.num_experts = num_experts
73
+ self.num_shared_experts = num_shared_experts
74
+ self.num_experts_per_tok = num_experts_per_tok
75
+ self.n_group = n_group
76
+ self.topk_group = topk_group
77
+ self.moe_intermediate_size = moe_intermediate_size
78
+ self.first_k_dense_replace = first_k_dense_replace
79
+ self.output_router_logits = output_router_logits
80
+ self.routed_scaling_factor = routed_scaling_factor
81
+ self.partial_rotary_factor = partial_rotary_factor
82
+ self.router_type = router_type
83
+ super().__init__(pad_token_id=pad_token_id, tie_word_embeddings=tie_word_embeddings, **kwargs)
84
+ self._attn_implementation = _attn_implementation
85
+
code/configuration_bailingmm2.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2024 ANT Group and the HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ from transformers import PretrainedConfig
17
+ from qwen2_5_vit import Qwen2_5_VLVisionConfig
18
+ from configuration_bailing_moe_v2 import BailingMoeV2Config
19
+
20
+
21
+ class BailingMM2Config(PretrainedConfig):
22
+ model_type = "bailingmm_moe_v2_lite"
23
+ # Declared so transformers' `_attn_implementation` setter recurses into both towers.
24
+ # Without it an explicit attn_implementation (e.g. "eager" on ROCm, which has no
25
+ # flash-attn) never reaches them, and their "flash_attention_2" defaults raise at
26
+ # model construction.
27
+ sub_configs = {"vision_config": Qwen2_5_VLVisionConfig, "llm_config": BailingMoeV2Config}
28
+
29
+ def __init__(
30
+ self,
31
+ mlp_depth=1,
32
+ llm_config: BailingMoeV2Config = None,
33
+ vision_config: Qwen2_5_VLVisionConfig = None,
34
+ audio_config=None,
35
+ **kwargs
36
+ ):
37
+ if audio_config is not None:
38
+ raise ValueError("audio_config is not supported by Ming Image inference")
39
+ self.audio_config = None
40
+ self.vision_config = Qwen2_5_VLVisionConfig(**vision_config) if isinstance(vision_config, dict) else vision_config
41
+ self.llm_config = BailingMoeV2Config(**llm_config) if isinstance(llm_config, dict) else llm_config
42
+ self.mlp_depth = mlp_depth
43
+ super().__init__(**kwargs)
code/examples/profiles/generation_edit.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "schema_version": 1,
3
+ "inference_profile": "generation_edit",
4
+ "alignment_padding_mode": "zero_masked",
5
+ "multi_frame_output": false,
6
+ "vae_input_channels": 4,
7
+ "vae_sample_mode": "argmax"
8
+ }
code/examples/profiles/layer_decompose.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "schema_version": 1,
3
+ "inference_profile": "layer_decompose",
4
+ "alignment_padding_mode": "learned",
5
+ "multi_frame_output": true,
6
+ "vae_input_channels": 4,
7
+ "vae_sample_mode": "argmax"
8
+ }
code/generate_paired.sh ADDED
@@ -0,0 +1,162 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # Paired pipeline: Ling-3.0-flash-VL prompt enhancement -> Ming-Image text-to-image.
3
+ #
4
+ # Stage 1 pe_ling.py caption -> validated structured JSON prompt
5
+ # (system prompt: assets/t2i_rewriter_system_prompt.txt)
6
+ # Stage 2 infer.py --task text-to-image --prompt <json file> -> PNG
7
+ # (infer.py reads --prompt as a file when the path exists)
8
+ #
9
+ # Artifacts land in --output-dir: enhanced_prompt.json (overwritten per run)
10
+ # plus the PNG(s) infer.py writes (image_00.png for text-to-image).
11
+ # Fails loudly at every stage (set -Eeuo pipefail + ERR trap + stage checks).
12
+ set -Eeuo pipefail
13
+ trap 'printf "generate_paired: FAILED at line %d (exit %d)\n" "$LINENO" "$?" >&2' ERR
14
+
15
+ SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)"
16
+ PYTHON="${PYTHON:-python3}"
17
+
18
+ # Local llama-server seat serving Ling-3.0-flash-VL on the target box.
19
+ DEFAULT_BASE_URL="http://127.0.0.1:8090/v1"
20
+ DEFAULT_PE_MODEL="ling-3.0-flash-vl-mtp-halo-STRIX_LEAN"
21
+
22
+ usage() {
23
+ cat <<'EOF'
24
+ Usage: generate_paired.sh --model DIR_OR_REPO CAPTION [options] [-- EXTRA_INFER_ARGS...]
25
+
26
+ Enhances CAPTION with Ling-3.0-flash-VL (pe_ling.py), validates the structured
27
+ JSON rewrite, then renders it with infer.py --task text-to-image.
28
+
29
+ Required:
30
+ CAPTION free-form design caption (positional)
31
+ --model DIR_OR_REPO Ming checkpoint directory or HF repo id
32
+ (may also be set via the MING_MODEL environment variable)
33
+
34
+ Passthrough to infer.py (all optional; infer.py defaults in parentheses):
35
+ --resolution N resolution bucket, 1024 or 2048 for text-to-image;
36
+ other positive values snap to the nearest bucket (2048)
37
+ --seed N generation seed (42)
38
+ --steps N diffusion steps (12)
39
+ -- everything after this is passed to infer.py verbatim
40
+ (e.g. -- --validate-only --dtype float16)
41
+
42
+ Prompt-enhancement endpoint:
43
+ --base-url URL OpenAI-compatible base URL (http://127.0.0.1:8090/v1)
44
+ --pe-model ID chat model id served there
45
+ (ling-3.0-flash-vl-mtp-halo-STRIX_LEAN)
46
+ LITELLM_API_KEY env exported key is sent as a Bearer token (for a gated
47
+ OpenAI-compatible gateway such as LiteLLM)
48
+
49
+ Other:
50
+ --output-dir DIR artifact directory (outputs/paired)
51
+ -h, --help this help
52
+
53
+ Examples:
54
+ ./generate_paired.sh --model /models/Ming-Image-0.1-Design \
55
+ "espresso machine product poster, warm morning light" --resolution 2048
56
+
57
+ LITELLM_API_KEY=sk-... ./generate_paired.sh \
58
+ --base-url http://<gateway-host>:4000/v1 --pe-model <gateway-model-name> \
59
+ --model /models/Ming-Image-0.1-Design "a caption" --seed 7
60
+ EOF
61
+ }
62
+
63
+ die() {
64
+ printf 'generate_paired: %s\n' "$*" >&2
65
+ exit 1
66
+ }
67
+
68
+ model="${MING_MODEL:-}"
69
+ base_url="$DEFAULT_BASE_URL"
70
+ pe_model="$DEFAULT_PE_MODEL"
71
+ output_dir="outputs/paired"
72
+ resolution=""
73
+ seed=""
74
+ steps=""
75
+ caption=""
76
+ extra_infer_args=()
77
+
78
+ while [[ $# -gt 0 ]]; do
79
+ case "$1" in
80
+ --model) [[ $# -ge 2 ]] || die "--model requires a value"; model="$2"; shift 2 ;;
81
+ --base-url) [[ $# -ge 2 ]] || die "--base-url requires a value"; base_url="$2"; shift 2 ;;
82
+ --pe-model) [[ $# -ge 2 ]] || die "--pe-model requires a value"; pe_model="$2"; shift 2 ;;
83
+ --output-dir) [[ $# -ge 2 ]] || die "--output-dir requires a value"; output_dir="$2"; shift 2 ;;
84
+ --resolution) [[ $# -ge 2 ]] || die "--resolution requires a value"; resolution="$2"; shift 2 ;;
85
+ --seed) [[ $# -ge 2 ]] || die "--seed requires a value"; seed="$2"; shift 2 ;;
86
+ --steps) [[ $# -ge 2 ]] || die "--steps requires a value"; steps="$2"; shift 2 ;;
87
+ -h|--help) usage; exit 0 ;;
88
+ --) shift; extra_infer_args+=("$@"); break ;;
89
+ -*) usage >&2; die "unknown option: $1" ;;
90
+ *)
91
+ if [[ -n "$caption" ]]; then
92
+ usage >&2
93
+ die "unexpected extra argument: $1 (CAPTION was already given)"
94
+ fi
95
+ caption="$1"
96
+ shift
97
+ ;;
98
+ esac
99
+ done
100
+
101
+ [[ -n "$caption" ]] || { usage >&2; die "CAPTION is required"; }
102
+ [[ -n "$model" ]] || { usage >&2; die "--model DIR_OR_REPO is required (or set MING_MODEL)"; }
103
+ if [[ -n "$resolution" && ! "$resolution" =~ ^[0-9]+$ ]]; then
104
+ die "--resolution must be a positive integer, got: $resolution"
105
+ fi
106
+ if [[ -n "$seed" && ! "$seed" =~ ^-?[0-9]+$ ]]; then
107
+ die "--seed must be an integer, got: $seed"
108
+ fi
109
+ if [[ -n "$steps" && ! "$steps" =~ ^[0-9]+$ ]]; then
110
+ die "--steps must be a positive integer, got: $steps"
111
+ fi
112
+ [[ -f "$SCRIPT_DIR/pe_ling.py" ]] || die "missing stage-1 script: $SCRIPT_DIR/pe_ling.py"
113
+ [[ -f "$SCRIPT_DIR/infer.py" ]] || die "missing stage-2 script: $SCRIPT_DIR/infer.py"
114
+ command -v "$PYTHON" >/dev/null 2>&1 || die "python interpreter not found: $PYTHON (override with PYTHON=...)"
115
+
116
+ mkdir -p -- "$output_dir" || die "cannot create output directory: $output_dir"
117
+ prompt_json="$output_dir/enhanced_prompt.json"
118
+
119
+ printf '== stage 1/2: prompt enhancement (pe_ling.py, model %s @ %s)\n' "$pe_model" "$base_url" >&2
120
+ "$PYTHON" "$SCRIPT_DIR/pe_ling.py" "$caption" \
121
+ --out "$prompt_json" \
122
+ --base-url "$base_url" \
123
+ --model "$pe_model"
124
+ [[ -s "$prompt_json" ]] || die "prompt enhancement produced no prompt file: $prompt_json"
125
+
126
+ printf '== stage 2/2: Ming-Image text-to-image (infer.py, model %s)\n' "$model" >&2
127
+ infer_args=(
128
+ --model "$model"
129
+ --task text-to-image
130
+ --prompt "$prompt_json"
131
+ --output-dir "$output_dir"
132
+ )
133
+ if [[ -n "$resolution" ]]; then infer_args+=(--resolution "$resolution"); fi
134
+ if [[ -n "$seed" ]]; then infer_args+=(--seed "$seed"); fi
135
+ if [[ -n "$steps" ]]; then infer_args+=(--steps "$steps"); fi
136
+ if [[ ${#extra_infer_args[@]} -gt 0 ]]; then infer_args+=("${extra_infer_args[@]}"); fi
137
+ validate_only=0
138
+ for arg in ${extra_infer_args[@]+"${extra_infer_args[@]}"}; do
139
+ if [[ "$arg" == "--validate-only" ]]; then validate_only=1; fi
140
+ done
141
+ "$PYTHON" "$SCRIPT_DIR/infer.py" "${infer_args[@]}"
142
+
143
+ if [[ "$validate_only" -eq 1 ]]; then
144
+ printf 'generate_paired: --validate-only dry run, no PNG expected; enhanced prompt: %s\n' \
145
+ "$prompt_json" >&2
146
+ exit 0
147
+ fi
148
+
149
+ # infer.py exits non-zero on failure (set -e above); additionally verify the
150
+ # promised PNG artifacts actually exist so a silent no-write still fails
151
+ # loudly. -newer pins the check to THIS run: stage 2 always writes its PNG
152
+ # after stage 1 wrote enhanced_prompt.json, so stale PNGs do not satisfy it.
153
+ pngs=()
154
+ while IFS= read -r png; do
155
+ pngs+=("$png")
156
+ done < <(find "$output_dir" -maxdepth 1 -name '*.png' -type f -newer "$prompt_json" | sort)
157
+ if [[ ${#pngs[@]} -eq 0 ]]; then
158
+ die "infer.py exited 0 but wrote no PNG under $output_dir in this run"
159
+ fi
160
+ printf 'generate_paired: enhanced prompt: %s\n' "$prompt_json" >&2
161
+ printf 'generate_paired: %d PNG(s):\n' "${#pngs[@]}" >&2
162
+ printf '%s\n' "${pngs[@]}"
code/image_processing_bailingmm2.py ADDED
@@ -0,0 +1,451 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2024 The Qwen team, Alibaba Group and the HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
5
+ # and OPT implementations in this library. It has been modified from its
6
+ # original forms to accommodate minor architectural differences compared
7
+ # to GPT-NeoX and OPT used by the Meta AI team that trained the model.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ """Image processor class for BailingMM"""
21
+
22
+ import math
23
+ from typing import Dict, List, Optional, Union
24
+
25
+ import numpy as np
26
+
27
+ from transformers.image_processing_utils import BaseImageProcessor, BatchFeature
28
+ from transformers.image_transforms import (
29
+ convert_to_rgb,
30
+ resize,
31
+ to_channel_dimension_format,
32
+ )
33
+ from bailingmm_utils import VideoInput
34
+ from transformers.image_utils import (
35
+ OPENAI_CLIP_MEAN,
36
+ OPENAI_CLIP_STD,
37
+ ChannelDimension,
38
+ ImageInput,
39
+ PILImageResampling,
40
+ get_image_size,
41
+ infer_channel_dimension_format,
42
+ is_scaled_image,
43
+ is_valid_image,
44
+ make_list_of_images,
45
+ to_numpy_array,
46
+ valid_images,
47
+ validate_preprocess_arguments,
48
+ )
49
+ from transformers.utils import TensorType, is_vision_available, logging
50
+ logger = logging.get_logger(__name__)
51
+
52
+ if is_vision_available():
53
+ from PIL import Image
54
+
55
+ def make_batched_images(images) -> List[List[ImageInput]]:
56
+ """
57
+ Accepts images in list or nested list format, and makes a list of images for preprocessing.
58
+
59
+ Args:
60
+ images (`Union[List[List[ImageInput]], List[ImageInput], ImageInput]`):
61
+ The input image.
62
+
63
+ Returns:
64
+ list: A list of images.
65
+ """
66
+ if isinstance(images, (list, tuple)) and isinstance(images[0], (list, tuple)) and is_valid_image(images[0][0]):
67
+ return [img for img_list in images for img in img_list]
68
+
69
+ elif isinstance(images, (list, tuple)) and is_valid_image(images[0]):
70
+ return images
71
+
72
+ elif is_valid_image(images):
73
+ return [images]
74
+
75
+ raise ValueError(f"Could not make batched images from {images}")
76
+
77
+ # Copied from transformers.models.llava_next_video.image_processing_llava_next_video.make_batched_videos
78
+ def make_batched_videos(videos) -> List[VideoInput]:
79
+ if isinstance(videos, (list, tuple)) and isinstance(videos[0], (list, tuple)) and is_valid_image(videos[0][0]):
80
+ return videos
81
+
82
+ elif isinstance(videos, (list, tuple)) and is_valid_image(videos[0]):
83
+ if isinstance(videos[0], Image.Image):
84
+ return [videos]
85
+ elif len(videos[0].shape) == 4:
86
+ return [list(video) for video in videos]
87
+
88
+ elif is_valid_image(videos) and len(videos.shape) == 4:
89
+ return [list(videos)]
90
+
91
+ raise ValueError(f"Could not make batched video from {videos}")
92
+
93
+ def smart_resize(
94
+ height: int, width: int, factor: int = 28, min_pixels: int = 56 * 56, max_pixels: int = 14 * 14 * 4 * 1280
95
+ ):
96
+ """Rescales the image so that the following conditions are met:
97
+
98
+ 1. Both dimensions (height and width) are divisible by 'factor'.
99
+
100
+ 2. The total number of pixels is within the range ['min_pixels', 'max_pixels'].
101
+
102
+ 3. The aspect ratio of the image is maintained as closely as possible.
103
+
104
+ """
105
+ if height < factor or width < factor:
106
+ raise ValueError(f"height:{height} or width:{width} must be larger than factor:{factor}")
107
+ elif max(height, width) / min(height, width) > 200:
108
+ raise ValueError(
109
+ f"absolute aspect ratio must be smaller than 200, got {max(height, width) / min(height, width)}"
110
+ )
111
+ h_bar = round(height / factor) * factor
112
+ w_bar = round(width / factor) * factor
113
+ if h_bar * w_bar > max_pixels:
114
+ beta = math.sqrt((height * width) / max_pixels)
115
+ h_bar = math.floor(height / beta / factor) * factor
116
+ w_bar = math.floor(width / beta / factor) * factor
117
+ elif h_bar * w_bar < min_pixels:
118
+ beta = math.sqrt(min_pixels / (height * width))
119
+ h_bar = math.ceil(height * beta / factor) * factor
120
+ w_bar = math.ceil(width * beta / factor) * factor
121
+ return h_bar, w_bar
122
+
123
+ class BailingMM2ImageProcessor(BaseImageProcessor):
124
+ r"""
125
+ Constructs a BailingMM2 image processor that dynamically resizes images based on the original images.
126
+
127
+ Args:
128
+ do_resize (`bool`, *optional*, defaults to `True`):
129
+ Whether to resize the image's (height, width) dimensions.
130
+ resample (`PILImageResampling`, *optional*, defaults to `Resampling.BICUBIC`):
131
+ Resampling filter to use when resizing the image.
132
+ do_rescale (`bool`, *optional*, defaults to `True`):
133
+ Whether to rescale the image by the specified scale `rescale_factor`.
134
+ rescale_factor (`int` or `float`, *optional*, defaults to `1/255`):
135
+ Scale factor to use if rescaling the image.
136
+ do_normalize (`bool`, *optional*, defaults to `True`):
137
+ Whether to normalize the image.
138
+ image_mean (`float` or `List[float]`, *optional*, defaults to `[0.48145466, 0.4578275, 0.40821073]`):
139
+ Mean to use if normalizing the image. This is a float or list of floats for each channel in the image.
140
+ image_std (`float` or `List[float]`, *optional*, defaults to `[0.26862954, 0.26130258, 0.27577711]`):
141
+ Standard deviation to use if normalizing the image. This is a float or list of floats for each channel in the image.
142
+ do_convert_rgb (`bool`, *optional*, defaults to `True`):
143
+ Whether to convert the image to RGB.
144
+ min_pixels (`int`, *optional*, defaults to `56 * 56`):
145
+ The min pixels of the image to resize the image.
146
+ max_pixels (`int`, *optional*, defaults to `28 * 28 * 1280`):
147
+ The max pixels of the image to resize the image.
148
+ patch_size (`int`, *optional*, defaults to 14):
149
+ The spacial patch size of the vision encoder.
150
+ temporal_patch_size (`int`, *optional*, defaults to 2):
151
+ The temporal patch size of the vision encoder.
152
+ merge_size (`int`, *optional*, defaults to 2):
153
+ The merge size of the vision encoder to llm encoder.
154
+ """
155
+
156
+ model_input_names = ["pixel_values", "image_grid_thw", "pixel_values_videos", "video_grid_thw"]
157
+
158
+ def __init__(
159
+ self,
160
+ do_resize: bool = True,
161
+ resample: PILImageResampling = PILImageResampling.BICUBIC,
162
+ do_rescale: bool = True,
163
+ rescale_factor: Union[int, float] = 1 / 255,
164
+ do_normalize: bool = True,
165
+ image_mean: Optional[Union[float, List[float]]] = None,
166
+ image_std: Optional[Union[float, List[float]]] = None,
167
+ do_convert_rgb: bool = True,
168
+ min_pixels: int = 78400,
169
+ max_pixels: int = 2007040,
170
+ patch_size: int = 14,
171
+ temporal_patch_size: int = 2,
172
+ merge_size: int = 2,
173
+ **kwargs,
174
+ ) -> None:
175
+ super().__init__(**kwargs)
176
+ self.do_resize = do_resize
177
+ self.resample = resample
178
+ self.do_rescale = do_rescale
179
+ self.rescale_factor = rescale_factor
180
+ self.do_normalize = do_normalize
181
+ self.image_mean = image_mean if image_mean is not None else OPENAI_CLIP_MEAN
182
+ self.image_std = image_std if image_std is not None else OPENAI_CLIP_STD
183
+ self.min_pixels = min_pixels
184
+ self.max_pixels = max_pixels
185
+ self.patch_size = patch_size
186
+ self.temporal_patch_size = temporal_patch_size
187
+ self.merge_size = merge_size
188
+ self.size = {"min_pixels": min_pixels, "max_pixels": max_pixels}
189
+ self.do_convert_rgb = do_convert_rgb
190
+
191
+ def _preprocess(
192
+ self,
193
+ images: Union[ImageInput, VideoInput],
194
+ do_resize: bool = None,
195
+ resample: PILImageResampling = None,
196
+ do_rescale: bool = None,
197
+ rescale_factor: float = None,
198
+ do_normalize: bool = None,
199
+ image_mean: Optional[Union[float, List[float]]] = None,
200
+ image_std: Optional[Union[float, List[float]]] = None,
201
+ do_convert_rgb: bool = None,
202
+ data_format: Optional[ChannelDimension] = ChannelDimension.FIRST,
203
+ input_data_format: Optional[Union[str, ChannelDimension]] = None,
204
+ ):
205
+ """
206
+ Preprocess an image or batch of images. Copy of the `preprocess` method from `CLIPImageProcessor`.
207
+
208
+ Args:
209
+ images (`ImageInput`):
210
+ Image or batch of images to preprocess. Expects pixel values ranging from 0 to 255. If pixel values range from 0 to 1, set `do_rescale=False`.
211
+ vision_info (`List[Dict]`, *optional*):
212
+ Optional list of dictionaries containing additional information about vision inputs.
213
+ do_resize (`bool`, *optional*, defaults to `self.do_resize`):
214
+ Whether to resize the image.
215
+ resample (`PILImageResampling`, *optional*, defaults to `self.resample`):
216
+ Resampling filter to use if resizing the image. This can be one of the `PILImageResampling` enums.
217
+ do_rescale (`bool`, *optional*, defaults to `self.do_rescale`):
218
+ Whether to rescale the image.
219
+ rescale_factor (`float`, *optional*, defaults to `self.rescale_factor`):
220
+ Scale factor to use if rescaling the image.
221
+ do_normalize (`bool`, *optional*, defaults to `self.do_normalize`):
222
+ Whether to normalize the image.
223
+ image_mean (`float` or `List[float]`, *optional*, defaults to `self.image_mean`):
224
+ Mean to use if normalizing the image. Can be a float or a list of floats corresponding to the number of channels in the image.
225
+ image_std (`float` or `List[float]`, *optional*, defaults to `self.image_std`):
226
+ Standard deviation to use if normalizing the image. Can be a float or a list of floats corresponding to the number of channels in the image.
227
+ do_convert_rgb (`bool`, *optional*, defaults to `self.do_convert_rgb`):
228
+ Whether to convert the image to RGB.
229
+ data_format (`ChannelDimension`, *optional*, defaults to `ChannelDimension.FIRST`):
230
+ The channel dimension format for the output image. Can be one of:
231
+ - `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
232
+ - `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
233
+ - Unset: Use the channel dimension format of the input image.
234
+ input_data_format (`ChannelDimension` or `str`, *optional*):
235
+ The channel dimension format for the input image. Can be one of:
236
+ - `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
237
+ - `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
238
+ - `"none"` or `ChannelDimension.NONE`: image in (height, width) format. - `"none"` or `ChannelDimension.NONE`: image in (height, width) format.
239
+ """
240
+ images = make_list_of_images(images)
241
+
242
+ if do_convert_rgb:
243
+ images = [convert_to_rgb(image) for image in images]
244
+
245
+ # All transformations expect numpy arrays.
246
+ images = [to_numpy_array(image) for image in images]
247
+
248
+ if is_scaled_image(images[0]) and do_rescale:
249
+ logger.warning_once(
250
+ "It looks like you are trying to rescale already rescaled images. If the input"
251
+ " images have pixel values between 0 and 1, set `do_rescale=False` to avoid rescaling them again."
252
+ )
253
+ if input_data_format is None:
254
+ # We assume that all images have the same channel dimension format.
255
+ input_data_format = infer_channel_dimension_format(images[0])
256
+
257
+ height, width = get_image_size(images[0], channel_dim=input_data_format)
258
+ resized_height, resized_width = height, width
259
+ processed_images = []
260
+ for image in images:
261
+ if do_resize:
262
+ resized_height, resized_width = smart_resize(
263
+ height,
264
+ width,
265
+ factor=self.patch_size * self.merge_size,
266
+ min_pixels=self.min_pixels,
267
+ max_pixels=self.max_pixels,
268
+ )
269
+ image = resize(
270
+ image, size=(resized_height, resized_width), resample=resample, input_data_format=input_data_format
271
+ )
272
+
273
+ if do_rescale:
274
+ image = self.rescale(image, scale=rescale_factor, input_data_format=input_data_format)
275
+
276
+ if do_normalize:
277
+ image = self.normalize(
278
+ image=image, mean=image_mean, std=image_std, input_data_format=input_data_format
279
+ )
280
+
281
+ image = to_channel_dimension_format(image, data_format, input_channel_dim=input_data_format)
282
+ processed_images.append(image)
283
+
284
+ patches = np.array(processed_images)
285
+ if data_format == ChannelDimension.LAST:
286
+ patches = patches.transpose(0, 3, 1, 2)
287
+ if patches.shape[0] == 1:
288
+ patches = np.tile(patches, (self.temporal_patch_size, 1, 1, 1))
289
+ channel = patches.shape[1]
290
+ grid_t = patches.shape[0] // self.temporal_patch_size
291
+ grid_h, grid_w = resized_height // self.patch_size, resized_width // self.patch_size
292
+ patches = patches.reshape(
293
+ grid_t,
294
+ self.temporal_patch_size,
295
+ channel,
296
+ grid_h // self.merge_size,
297
+ self.merge_size,
298
+ self.patch_size,
299
+ grid_w // self.merge_size,
300
+ self.merge_size,
301
+ self.patch_size,
302
+ )
303
+ patches = patches.transpose(0, 3, 6, 4, 7, 2, 1, 5, 8)
304
+ flatten_patches = patches.reshape(
305
+ grid_t * grid_h * grid_w, channel * self.temporal_patch_size * self.patch_size * self.patch_size
306
+ )
307
+
308
+ return flatten_patches, (grid_t, grid_h, grid_w)
309
+
310
+ def preprocess(
311
+ self,
312
+ images: ImageInput,
313
+ videos: VideoInput = None,
314
+ do_resize: bool = None,
315
+ size: Dict[str, int] = None,
316
+ resample: PILImageResampling = None,
317
+ do_rescale: bool = None,
318
+ rescale_factor: float = None,
319
+ do_normalize: bool = None,
320
+ image_mean: Optional[Union[float, List[float]]] = None,
321
+ image_std: Optional[Union[float, List[float]]] = None,
322
+ do_convert_rgb: bool = None,
323
+ return_tensors: Optional[Union[str, TensorType]] = None,
324
+ data_format: Optional[ChannelDimension] = ChannelDimension.FIRST,
325
+ input_data_format: Optional[Union[str, ChannelDimension]] = None,
326
+ ):
327
+ """
328
+ Args:
329
+ images (`ImageInput`):
330
+ Image to preprocess. Expects a single or batch of images with pixel values ranging from 0 to 255. If
331
+ passing in images with pixel values between 0 and 1, set `do_rescale=False`.
332
+ videos (`VideoInput`):
333
+ Video to preprocess. Expects a single or batch of videos with pixel values ranging from 0 to 255. If
334
+ passing in videos with pixel values between 0 and 1, set `do_rescale=False`.
335
+ do_resize (`bool`, *optional*, defaults to `self.do_resize`):
336
+ Whether to resize the image.
337
+ size (`Dict[str, int]`, *optional*, defaults to `self.size`):
338
+ Size of the image after resizing. Shortest edge of the image is resized to size["shortest_edge"], with
339
+ the longest edge resized to keep the input aspect ratio.
340
+ resample (`int`, *optional*, defaults to `self.resample`):
341
+ Resampling filter to use if resizing the image. This can be one of the enum `PILImageResampling`. Only
342
+ has an effect if `do_resize` is set to `True`.
343
+ do_rescale (`bool`, *optional*, defaults to `self.do_rescale`):
344
+ Whether to rescale the image.
345
+ rescale_factor (`float`, *optional*, defaults to `self.rescale_factor`):
346
+ Rescale factor to rescale the image by if `do_rescale` is set to `True`.
347
+ do_normalize (`bool`, *optional*, defaults to `self.do_normalize`):
348
+ Whether to normalize the image.
349
+ image_mean (`float` or `List[float]`, *optional*, defaults to `self.image_mean`):
350
+ Image mean to use for normalization. Only has an effect if `do_normalize` is set to `True`.
351
+ image_std (`float` or `List[float]`, *optional*, defaults to `self.image_std`):
352
+ Image standard deviation to use for normalization. Only has an effect if `do_normalize` is set to
353
+ `True`.
354
+ do_convert_rgb (`bool`, *optional*, defaults to `self.do_convert_rgb`):
355
+ Whether to convert the image to RGB.
356
+ return_tensors (`str` or `TensorType`, *optional*):
357
+ The type of tensors to return. Can be one of:
358
+ - Unset: Return a list of `np.ndarray`.
359
+ - `TensorType.TENSORFLOW` or `'tf'`: Return a batch of type `tf.Tensor`.
360
+ - `TensorType.PYTORCH` or `'pt'`: Return a batch of type `torch.Tensor`.
361
+ - `TensorType.NUMPY` or `'np'`: Return a batch of type `np.ndarray`.
362
+ - `TensorType.JAX` or `'jax'`: Return a batch of type `jax.numpy.ndarray`.
363
+ data_format (`ChannelDimension` or `str`, *optional*, defaults to `ChannelDimension.FIRST`):
364
+ The channel dimension format for the output image. Can be one of:
365
+ - `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
366
+ - `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
367
+ - Unset: Use the channel dimension format of the input image.
368
+ input_data_format (`ChannelDimension` or `str`, *optional*):
369
+ The channel dimension format for the input image. If unset, the channel dimension format is inferred
370
+ from the input image. Can be one of:
371
+ - `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
372
+ - `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
373
+ - `"none"` or `ChannelDimension.NONE`: image in (height, width) format.
374
+
375
+ """
376
+ do_resize = do_resize if do_resize is not None else self.do_resize
377
+ size = size if size is not None else self.size
378
+ resample = resample if resample is not None else self.resample
379
+ do_rescale = do_rescale if do_rescale is not None else self.do_rescale
380
+ rescale_factor = rescale_factor if rescale_factor is not None else self.rescale_factor
381
+ do_normalize = do_normalize if do_normalize is not None else self.do_normalize
382
+ image_mean = image_mean if image_mean is not None else self.image_mean
383
+ image_std = image_std if image_std is not None else self.image_std
384
+ do_convert_rgb = do_convert_rgb if do_convert_rgb is not None else self.do_convert_rgb
385
+
386
+ if images is not None:
387
+ images = make_batched_images(images)
388
+ if videos is not None:
389
+ videos = make_batched_videos(videos)
390
+
391
+ if images is not None and not valid_images(images):
392
+ raise ValueError(
393
+ "Invalid image type. Must be of type PIL.Image.Image, numpy.ndarray, "
394
+ "torch.Tensor, tf.Tensor or jax.ndarray."
395
+ )
396
+
397
+ validate_preprocess_arguments(
398
+ rescale_factor=rescale_factor,
399
+ do_normalize=do_normalize,
400
+ image_mean=image_mean,
401
+ image_std=image_std,
402
+ do_resize=do_resize,
403
+ size=size,
404
+ resample=resample,
405
+ )
406
+
407
+ if images is not None:
408
+ pixel_values, vision_grid_thws = [], []
409
+ for image in images:
410
+ patches, image_grid_thw = self._preprocess(
411
+ image,
412
+ do_resize=do_resize,
413
+ resample=resample,
414
+ do_rescale=do_rescale,
415
+ rescale_factor=rescale_factor,
416
+ do_normalize=do_normalize,
417
+ image_mean=image_mean,
418
+ image_std=image_std,
419
+ data_format=data_format,
420
+ do_convert_rgb=do_convert_rgb,
421
+ input_data_format=input_data_format,
422
+ )
423
+ pixel_values.extend(patches)
424
+ vision_grid_thws.append(image_grid_thw)
425
+ pixel_values = np.array(pixel_values)
426
+ vision_grid_thws = np.array(vision_grid_thws)
427
+ data = {"pixel_values": pixel_values, "image_grid_thw": vision_grid_thws}
428
+
429
+ if videos is not None:
430
+ pixel_values, vision_grid_thws = [], []
431
+ for images in videos:
432
+ patches, video_grid_thw = self._preprocess(
433
+ images,
434
+ do_resize=do_resize,
435
+ resample=resample,
436
+ do_rescale=do_rescale,
437
+ rescale_factor=rescale_factor,
438
+ do_normalize=do_normalize,
439
+ image_mean=image_mean,
440
+ image_std=image_std,
441
+ data_format=data_format,
442
+ do_convert_rgb=do_convert_rgb,
443
+ input_data_format=input_data_format,
444
+ )
445
+ pixel_values.extend(patches)
446
+ vision_grid_thws.append(video_grid_thw)
447
+ pixel_values = np.array(pixel_values)
448
+ vision_grid_thws = np.array(vision_grid_thws)
449
+ data = {"pixel_values_videos": pixel_values, "video_grid_thw": vision_grid_thws}
450
+
451
+ return BatchFeature(data=data, tensor_type=return_tensors)
code/infer.py ADDED
@@ -0,0 +1,468 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Unified Hugging Face inference entry point for Ming image checkpoints."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import argparse
7
+ import json
8
+ import re
9
+ import sys
10
+ from pathlib import Path
11
+ from typing import Iterable, List
12
+
13
+ from inference_profile import (
14
+ VALID_TASKS,
15
+ load_checkpoint_capabilities,
16
+ resolve_model_directory,
17
+ )
18
+ from mllm_device_map import (
19
+ build_mllm_device_plan,
20
+ load_mllm_num_hidden_layers,
21
+ validate_loaded_layer_devices,
22
+ )
23
+ CODE_DIRECTORY = Path(__file__).resolve().parent
24
+
25
+ TASK_RESOLUTION_BUCKETS = {
26
+ "text-to-image": (1024, 2048),
27
+ "image-edit": (1024,),
28
+ "layer-decompose": (512, 1024),
29
+ }
30
+ TASK_DEFAULT_RESOLUTIONS = {
31
+ "text-to-image": 2048,
32
+ "image-edit": 1024,
33
+ "layer-decompose": 1024,
34
+ }
35
+
36
+
37
+ def parse_args() -> argparse.Namespace:
38
+ parser = argparse.ArgumentParser(
39
+ description=(
40
+ "Run text-to-image, image editing, or layer decomposition with an "
41
+ "explicit checkpoint capability profile."
42
+ )
43
+ )
44
+ parser.add_argument("--model", required=True, help="Local model directory or HF Hub repo ID")
45
+ parser.add_argument("--task", required=True, choices=sorted(VALID_TASKS))
46
+ parser.add_argument("--prompt", help="Generation/edit prompt; optional for layer decomposition")
47
+ parser.add_argument("--input-image", type=Path, help="Required for edit and layer decomposition")
48
+ parser.add_argument("--num-layers", type=int, default=1)
49
+ parser.add_argument("--output-dir", type=Path, default=Path("outputs"))
50
+ parser.add_argument(
51
+ "--resolution",
52
+ type=int,
53
+ help=(
54
+ "Requested resolution bucket. Defaults: text-to-image 2048, "
55
+ "image-edit 1024, layer-decompose 1024. Requests snap to the "
56
+ "nearest bucket supported by the selected task."
57
+ ),
58
+ )
59
+ parser.add_argument(
60
+ "--steps",
61
+ type=int,
62
+ help="Override the checkpoint-family default (generation: 12, layers: 12)",
63
+ )
64
+ parser.add_argument("--seed", type=int, default=42)
65
+ parser.add_argument(
66
+ "--cfg",
67
+ type=float,
68
+ help="Override the checkpoint-family default (generation: 1.0, layers: 2.0)",
69
+ )
70
+ parser.add_argument("--dtype", choices=("bfloat16", "float16", "float32"), default="bfloat16")
71
+ parser.add_argument(
72
+ "--attn-implementation",
73
+ choices=("sdpa", "flash_attention_2", "eager"),
74
+ # The BailingMoeV2 LLM only implements eager and flash_attention_2
75
+ # attention classes; selecting "sdpa" fails closed at load time.
76
+ default="eager",
77
+ )
78
+ parser.add_argument("--device", default="cuda:0")
79
+ parser.add_argument(
80
+ "--device-map",
81
+ choices=("balanced", "none", "auto"),
82
+ default="balanced",
83
+ help="balanced reserves GPU 0 for fixed image modules and shards the MLLM",
84
+ )
85
+ parser.add_argument(
86
+ "--num-gpus",
87
+ type=int,
88
+ default=1,
89
+ help="required visible GPU count for balanced placement (any positive integer)",
90
+ )
91
+ parser.add_argument(
92
+ "--processor",
93
+ help="Optional processor data directory; defaults to <checkpoint>/mllm",
94
+ )
95
+ parser.add_argument("--revision", help="HF Hub model revision")
96
+ parser.add_argument("--cache-dir", type=Path)
97
+ parser.add_argument("--local-files-only", action="store_true")
98
+ parser.add_argument(
99
+ "--validate-only",
100
+ action="store_true",
101
+ help="Validate model profile and task arguments without loading weights",
102
+ )
103
+ parser.add_argument(
104
+ "--attention-bf16-reduction",
105
+ action="store_true",
106
+ help=(
107
+ "Let PyTorch's math attention kernel (the only SDPA kernel that runs on ROCm gfx1151) "
108
+ "stay in bf16 instead of upcasting to fp32: faster, less precise."
109
+ ),
110
+ )
111
+ parser.add_argument(
112
+ "--release-mllm-after-conditioning",
113
+ action="store_true",
114
+ help=(
115
+ "Free the MLLM, vision tower and connector as soon as the conditioning is computed, "
116
+ "before the diffusion steps. Lowers peak memory; one image per process."
117
+ ),
118
+ )
119
+ return parser.parse_args()
120
+
121
+
122
+ def parse_num_layers(text: str) -> int:
123
+ decompose_match = re.search(
124
+ r"decompose this image into\s+(\d+)\s+layers?", text.lower()
125
+ )
126
+ if decompose_match:
127
+ return int(decompose_match.group(1))
128
+ for line in text.splitlines():
129
+ s = line.strip().lower()
130
+ if s.startswith("number of layers:"):
131
+ try:
132
+ return int(s.split(":", 1)[1].strip())
133
+ except ValueError:
134
+ pass
135
+ return 5
136
+
137
+
138
+ def resolve_task_resolution(task: str, requested: int | None) -> int:
139
+ """Resolve a user request to the nearest supported task-level bucket."""
140
+ try:
141
+ buckets = TASK_RESOLUTION_BUCKETS[task]
142
+ except KeyError as exc:
143
+ raise ValueError(f"unsupported task for resolution policy: {task!r}") from exc
144
+ if requested is None:
145
+ return TASK_DEFAULT_RESOLUTIONS[task]
146
+ if isinstance(requested, bool) or not isinstance(requested, int) or requested <= 0:
147
+ raise ValueError("--resolution must be a positive integer")
148
+ return min(buckets, key=lambda value: (abs(value - requested), value))
149
+
150
+
151
+ def _load_prompt(prompt: str) -> str:
152
+ prompt_path = Path(prompt)
153
+ try:
154
+ is_file = prompt_path.is_file()
155
+ except OSError:
156
+ # Long literal prompts can be invalid filesystem paths. In that case,
157
+ # keep treating the argument as prompt text.
158
+ return prompt
159
+ if is_file:
160
+ return prompt_path.read_text(encoding="utf-8")
161
+ return prompt
162
+
163
+
164
+ def _dtype(name: str):
165
+ import torch
166
+
167
+ return {
168
+ "bfloat16": torch.bfloat16,
169
+ "float16": torch.float16,
170
+ "float32": torch.float32,
171
+ }[name]
172
+
173
+
174
+ def _build_messages(task: str, prompt: str, input_image: Path | None):
175
+ content = []
176
+ if input_image is not None:
177
+ content.append({"type": "image", "image": str(input_image)})
178
+ content.append({"type": "text", "text": prompt})
179
+ return [{"role": "HUMAN", "content": content}]
180
+
181
+
182
+ def _normalize_outputs(output) -> List[Image.Image]:
183
+ from PIL import Image
184
+
185
+ if isinstance(output, Image.Image):
186
+ return [output]
187
+ if isinstance(output, (list, tuple)) and all(
188
+ isinstance(item, Image.Image) for item in output
189
+ ):
190
+ return list(output)
191
+ raise TypeError(f"model returned unsupported output type: {type(output)!r}")
192
+
193
+
194
+ def _model_input_device(model):
195
+ try:
196
+ return model.device
197
+ except AttributeError:
198
+ return next(parameter.device for parameter in model.parameters() if not parameter.is_meta)
199
+
200
+
201
+ def _validate_balanced_placement(model, plan, torch) -> None:
202
+ layer_devices = []
203
+ for index, layer in enumerate(model.model.model.layers):
204
+ devices = {parameter.device for parameter in layer.parameters()}
205
+ if len(devices) != 1:
206
+ raise RuntimeError(f"MLLM layer {index} spans devices {sorted(map(str, devices))}")
207
+ device = devices.pop()
208
+ if device.type != "cuda" or device.index is None:
209
+ raise RuntimeError(f"MLLM layer {index} loaded on {device}, expected CUDA")
210
+ layer_devices.append(device.index)
211
+ validate_loaded_layer_devices(layer_devices, plan)
212
+
213
+ fixed_modules = {
214
+ "vision": model.vision,
215
+ "linear_proj": model.linear_proj,
216
+ "connector": getattr(model, "connector", None),
217
+ "proj_in": getattr(model, "proj_in", None),
218
+ "proj_out": getattr(model, "proj_out", None),
219
+ "proj_directvlm": getattr(model, "proj_directvlm", None),
220
+ "query_tokens": getattr(model, "query_tokens_dict", None),
221
+ "diffusion": getattr(model, "diffusion_loss", None),
222
+ }
223
+ expected = {torch.device("cuda:0")}
224
+ for name, module in fixed_modules.items():
225
+ if module is None:
226
+ continue
227
+ devices = {parameter.device for parameter in module.parameters()}
228
+ if devices and devices != expected:
229
+ raise RuntimeError(
230
+ f"fixed module {name} loaded on {sorted(map(str, devices))}, expected cuda:0"
231
+ )
232
+
233
+
234
+ def _save_outputs(
235
+ images: Iterable, output_dir: Path, task: str
236
+ ) -> List[Path]:
237
+ output_dir.mkdir(parents=True, exist_ok=True)
238
+ prefix = "layer" if task == "layer-decompose" else "image"
239
+ skip_first = task == "layer-decompose"
240
+ paths = []
241
+ for index, image in enumerate(images):
242
+ if skip_first and index == 0:
243
+ continue
244
+ output_path = output_dir / f"{prefix}_{index:02d}.png"
245
+ image.save(output_path)
246
+ paths.append(output_path)
247
+ return paths
248
+
249
+
250
+ def main() -> None:
251
+ args = parse_args()
252
+ model_directory = resolve_model_directory(
253
+ args.model,
254
+ revision=args.revision,
255
+ cache_dir=args.cache_dir,
256
+ local_files_only=args.local_files_only,
257
+ )
258
+ profile = load_checkpoint_capabilities(model_directory)
259
+
260
+ has_reference_image = args.input_image is not None
261
+ profile.validate_task(
262
+ args.task,
263
+ has_reference_image=has_reference_image,
264
+ num_layers=args.num_layers,
265
+ )
266
+ effective_resolution = resolve_task_resolution(args.task, args.resolution)
267
+ if args.resolution is not None and args.resolution != effective_resolution:
268
+ print(
269
+ f"resolution {args.resolution} snapped to {effective_resolution} "
270
+ f"for task {args.task}",
271
+ file=sys.stderr,
272
+ )
273
+ sampling = profile.resolve_sampling_parameters(steps=args.steps, cfg=args.cfg)
274
+ if args.input_image is not None and not args.input_image.is_file():
275
+ raise FileNotFoundError(f"input image does not exist: {args.input_image}")
276
+ if args.task != "layer-decompose" and not args.prompt:
277
+ raise ValueError(f"--prompt is required for {args.task}")
278
+
279
+ if args.prompt is not None:
280
+ prompt = _load_prompt(args.prompt)
281
+ else:
282
+ prompt = f"Decompose this image into {args.num_layers} layers."
283
+ num_layers = parse_num_layers(prompt) if args.task == "layer-decompose" else args.num_layers
284
+ if args.validate_only:
285
+ print(
286
+ json.dumps(
287
+ {
288
+ "model": str(model_directory),
289
+ "task": args.task,
290
+ "profile": profile.__dict__,
291
+ "sampling": sampling.__dict__,
292
+ "resolution": {
293
+ "requested": args.resolution,
294
+ "effective": effective_resolution,
295
+ },
296
+ },
297
+ indent=2,
298
+ )
299
+ )
300
+ return
301
+
302
+ model, processor = load_model_and_processor(model_directory, args)
303
+ images = run_generation(
304
+ model,
305
+ processor,
306
+ profile,
307
+ task=args.task,
308
+ prompt=prompt,
309
+ input_image=args.input_image,
310
+ resolution=effective_resolution,
311
+ sampling=sampling,
312
+ seed=args.seed,
313
+ num_layers=num_layers,
314
+ dtype=_dtype(args.dtype),
315
+ )
316
+
317
+ output_paths = _save_outputs(images, args.output_dir, args.task)
318
+ print(json.dumps({"outputs": [str(path.resolve()) for path in output_paths]}, indent=2))
319
+
320
+
321
+ def load_model_and_processor(model_directory: Path, args):
322
+ """Load the model and processor once; shared by the CLI and batch tools."""
323
+ import torch
324
+
325
+ from modeling_bailingmm2 import BailingMM2NativeForConditionalGeneration
326
+ from processing_bailingmm2 import load_bailingmm2_processor
327
+
328
+ # Processor/tokenizer data lives in the package's mllm/ component; the
329
+ # Python implementations stay in this repository (AutoProcessor would
330
+ # require them inside the data directory). --processor overrides the data
331
+ # directory only.
332
+ processor_directory = (
333
+ Path(args.processor).expanduser().resolve()
334
+ if args.processor
335
+ else model_directory / "mllm"
336
+ )
337
+ processor = load_bailingmm2_processor(processor_directory)
338
+
339
+ if getattr(args, "attention_bf16_reduction", False):
340
+ # The math SDPA kernel upcasts bf16 inputs to fp32 by default; this keeps it in bf16.
341
+ torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(True)
342
+
343
+ dtype = _dtype(args.dtype)
344
+ load_kwargs = {
345
+ "torch_dtype": dtype,
346
+ "attn_implementation": args.attn_implementation,
347
+ "load_image_gen": True,
348
+ "image_gen_device": args.device,
349
+ }
350
+ device_plan = None
351
+ if args.device_map == "balanced":
352
+ visible_gpus = torch.cuda.device_count()
353
+ if visible_gpus != args.num_gpus:
354
+ raise RuntimeError(
355
+ f"balanced placement requires {args.num_gpus} visible GPUs, got {visible_gpus}; "
356
+ "set CUDA_VISIBLE_DEVICES before starting Python"
357
+ )
358
+ num_hidden_layers = load_mllm_num_hidden_layers(model_directory / "mllm")
359
+ device_plan = build_mllm_device_plan(num_hidden_layers, visible_gpus)
360
+ load_kwargs["device_map"] = device_plan.device_map
361
+ print(
362
+ f"MLLM device plan: {visible_gpus} GPUs, layer_counts={device_plan.layer_counts}"
363
+ )
364
+ elif args.device_map == "auto":
365
+ load_kwargs["device_map"] = "auto"
366
+
367
+ model = BailingMM2NativeForConditionalGeneration.from_pretrained(
368
+ str(model_directory), **load_kwargs
369
+ )
370
+ if args.device_map == "none":
371
+ model = model.to(device=args.device, dtype=dtype)
372
+ elif device_plan is not None:
373
+ _validate_balanced_placement(model, device_plan, torch)
374
+ if getattr(args, "release_mllm_after_conditioning", False):
375
+ _release_mllm_before_sampling(model)
376
+ return model, processor
377
+
378
+
379
+ def _release_mllm_before_sampling(model) -> None:
380
+ """Free the MLLM-side modules once the conditioning exists (--release-mllm-after-conditioning).
381
+
382
+ Wraps the diffusion sampler: by the time it is called the conditioning tensors are computed,
383
+ so the language model, vision tower and connector are moved to the meta device (releasing
384
+ their memory) before the diffusion steps start. The model cannot generate again afterwards.
385
+ """
386
+ import gc
387
+
388
+ import torch
389
+
390
+ original_sample = model.diffusion_loss.sample
391
+
392
+ def sample_after_release(*args, **kwargs):
393
+ for name in ("model", "vision", "linear_proj", "connector"):
394
+ module = getattr(model, name, None)
395
+ if module is not None:
396
+ module.to("meta")
397
+ gc.collect()
398
+ if torch.cuda.is_available():
399
+ torch.cuda.empty_cache()
400
+ return original_sample(*args, **kwargs)
401
+
402
+ model.diffusion_loss.sample = sample_after_release
403
+
404
+
405
+ def run_generation(
406
+ model,
407
+ processor,
408
+ profile,
409
+ *,
410
+ task: str,
411
+ prompt: str,
412
+ input_image,
413
+ resolution: int,
414
+ sampling,
415
+ seed: int,
416
+ num_layers: int,
417
+ dtype,
418
+ ) -> List:
419
+ """Run one inference with an already-loaded model and processor."""
420
+ import torch
421
+ from PIL import Image
422
+
423
+ resolution = resolve_task_resolution(task, resolution)
424
+
425
+ messages = _build_messages(task, prompt, input_image)
426
+ text = processor.apply_chat_template(messages, add_generation_prompt=True)
427
+ image_inputs, video_inputs, _ = processor.process_vision_info(messages)
428
+
429
+ reference_image = None
430
+ if input_image is not None:
431
+ reference_mode = "RGB" if profile.vae_input_channels == 3 else "RGBA"
432
+ reference_image = Image.open(input_image).convert(reference_mode)
433
+
434
+ inputs = processor(
435
+ text=[text],
436
+ images=image_inputs,
437
+ videos=video_inputs,
438
+ return_tensors="pt",
439
+ image_gen_highres=resolution,
440
+ image_gen_ref_images=reference_image,
441
+ image_gen_input_channels=profile.vae_input_channels,
442
+ )
443
+ input_device = _model_input_device(model)
444
+ inputs = inputs.to(input_device)
445
+ for key, value in inputs.items():
446
+ if isinstance(value, torch.Tensor) and torch.is_floating_point(value):
447
+ inputs[key] = value.to(dtype=dtype)
448
+
449
+ output = model.generate(
450
+ **inputs,
451
+ image_gen=True,
452
+ image_gen_task=task,
453
+ image_gen_seed=seed,
454
+ image_gen_steps=sampling.steps,
455
+ image_gen_cfg=sampling.cfg,
456
+ num_frames_per_prompt=num_layers+1 if task == "layer-decompose" else num_layers,
457
+ )
458
+ images = _normalize_outputs(output)
459
+ expected_outputs = num_layers + 1 if task == "layer-decompose" else 1
460
+ if len(images) != expected_outputs:
461
+ raise RuntimeError(
462
+ f"expected {expected_outputs} output image(s), got {len(images)}"
463
+ )
464
+ return images
465
+
466
+
467
+ if __name__ == "__main__":
468
+ main()
code/inference_profile.py ADDED
@@ -0,0 +1,420 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Checkpoint capability contract for public Ming image inference.
2
+
3
+ New packages declare their capability in ``transformer/config.json`` via the
4
+ ``alignment_padding_mode`` / ``multi_frame_output`` pair; the VAE contract is
5
+ derived from ``vae/config.json``. Legacy packages that predate the component
6
+ metadata are loaded strictly from the root ``inference_profile.json`` during
7
+ the compatibility window. The contract is intentionally strict: selecting
8
+ behavior from directory names, missing state-dict keys, or task names can
9
+ silently load the wrong padding semantics and produce degraded images.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ from dataclasses import dataclass
15
+ import json
16
+ import math
17
+ from pathlib import Path
18
+ from typing import Any, Mapping, Optional, Union
19
+
20
+
21
+ PROFILE_FILENAME = "inference_profile.json"
22
+ PROFILE_SCHEMA_VERSION = 1
23
+
24
+ GENERATION_PROFILE = "generation_edit"
25
+ LAYER_PROFILE = "layer_decompose"
26
+ VALID_PROFILES = {GENERATION_PROFILE, LAYER_PROFILE}
27
+
28
+ LEARNED_PADDING = "learned"
29
+ ZERO_MASKED_PADDING = "zero_masked"
30
+ VALID_PADDING_MODES = {LEARNED_PADDING, ZERO_MASKED_PADDING}
31
+
32
+ VALID_VAE_SAMPLE_MODES = {"sample", "argmax"}
33
+ VALID_TASKS = {"text-to-image", "image-edit", "layer-decompose"}
34
+
35
+ REQUIRED_PROFILE_FIELDS = {
36
+ "schema_version",
37
+ "inference_profile",
38
+ "alignment_padding_mode",
39
+ "multi_frame_output",
40
+ "vae_input_channels",
41
+ "vae_sample_mode",
42
+ }
43
+
44
+
45
+ class InferenceProfileError(ValueError):
46
+ """Raised when a checkpoint does not satisfy the inference contract."""
47
+
48
+
49
+ @dataclass(frozen=True)
50
+ class SamplingParameters:
51
+ steps: int
52
+ cfg: float
53
+
54
+
55
+ DEFAULT_SAMPLING_PARAMETERS = {
56
+ GENERATION_PROFILE: SamplingParameters(steps=12, cfg=1.0),
57
+ LAYER_PROFILE: SamplingParameters(steps=12, cfg=2.0),
58
+ }
59
+
60
+
61
+ @dataclass(frozen=True)
62
+ class InferenceProfile:
63
+ schema_version: int
64
+ inference_profile: str
65
+ alignment_padding_mode: str
66
+ multi_frame_output: bool
67
+ vae_input_channels: int
68
+ vae_sample_mode: str
69
+
70
+ @classmethod
71
+ def from_dict(cls, raw: Mapping[str, Any]) -> "InferenceProfile":
72
+ missing = sorted(REQUIRED_PROFILE_FIELDS.difference(raw))
73
+ if missing:
74
+ raise InferenceProfileError(
75
+ "checkpoint inference profile is missing required fields: "
76
+ + ", ".join(missing)
77
+ )
78
+
79
+ unknown = sorted(set(raw).difference(REQUIRED_PROFILE_FIELDS))
80
+ if unknown:
81
+ raise InferenceProfileError(
82
+ "checkpoint inference profile contains unsupported fields: "
83
+ + ", ".join(unknown)
84
+ )
85
+
86
+ if type(raw["schema_version"]) is not int:
87
+ raise InferenceProfileError("schema_version must be an integer")
88
+ if type(raw["multi_frame_output"]) is not bool:
89
+ raise InferenceProfileError("multi_frame_output must be a boolean")
90
+ if type(raw["vae_input_channels"]) is not int:
91
+ raise InferenceProfileError("vae_input_channels must be an integer")
92
+
93
+ profile = cls(**{key: raw[key] for key in REQUIRED_PROFILE_FIELDS})
94
+ profile.validate()
95
+ return profile
96
+
97
+ def validate(self) -> None:
98
+ if self.schema_version != PROFILE_SCHEMA_VERSION:
99
+ raise InferenceProfileError(
100
+ f"unsupported profile schema_version={self.schema_version}; "
101
+ f"expected {PROFILE_SCHEMA_VERSION}"
102
+ )
103
+ if self.inference_profile not in VALID_PROFILES:
104
+ raise InferenceProfileError(
105
+ f"inference_profile must be one of {sorted(VALID_PROFILES)}, "
106
+ f"got {self.inference_profile!r}"
107
+ )
108
+ if self.alignment_padding_mode not in VALID_PADDING_MODES:
109
+ raise InferenceProfileError(
110
+ "alignment_padding_mode must be one of "
111
+ f"{sorted(VALID_PADDING_MODES)}, got {self.alignment_padding_mode!r}"
112
+ )
113
+ if self.vae_input_channels not in (3, 4):
114
+ raise InferenceProfileError(
115
+ f"vae_input_channels must be 3 or 4, got {self.vae_input_channels}"
116
+ )
117
+ if self.vae_sample_mode not in VALID_VAE_SAMPLE_MODES:
118
+ raise InferenceProfileError(
119
+ f"vae_sample_mode must be one of {sorted(VALID_VAE_SAMPLE_MODES)}, "
120
+ f"got {self.vae_sample_mode!r}"
121
+ )
122
+
123
+ if self.inference_profile == GENERATION_PROFILE:
124
+ if self.alignment_padding_mode != ZERO_MASKED_PADDING:
125
+ raise InferenceProfileError(
126
+ "generation_edit checkpoints must use zero_masked alignment padding"
127
+ )
128
+ if self.multi_frame_output:
129
+ raise InferenceProfileError(
130
+ "generation_edit checkpoints must set multi_frame_output=false"
131
+ )
132
+ if self.vae_input_channels != 4:
133
+ raise InferenceProfileError(
134
+ "generation_edit checkpoints must declare vae_input_channels=4"
135
+ )
136
+ if self.vae_sample_mode != "argmax":
137
+ raise InferenceProfileError(
138
+ "generation_edit checkpoints must declare vae_sample_mode='argmax'"
139
+ )
140
+ else:
141
+ if self.alignment_padding_mode != LEARNED_PADDING:
142
+ raise InferenceProfileError(
143
+ "layer_decompose checkpoints must use learned alignment padding"
144
+ )
145
+ if not self.multi_frame_output:
146
+ raise InferenceProfileError(
147
+ "layer_decompose checkpoints must set multi_frame_output=true"
148
+ )
149
+ if self.vae_input_channels != 4:
150
+ raise InferenceProfileError(
151
+ "layer_decompose checkpoints must declare vae_input_channels=4"
152
+ )
153
+ if self.vae_sample_mode != "argmax":
154
+ raise InferenceProfileError(
155
+ "layer_decompose checkpoints must declare vae_sample_mode='argmax'"
156
+ )
157
+
158
+ def validate_task(
159
+ self,
160
+ task: str,
161
+ *,
162
+ has_reference_image: bool,
163
+ num_layers: int = 1,
164
+ ) -> None:
165
+ if task not in VALID_TASKS:
166
+ raise InferenceProfileError(
167
+ f"task must be one of {sorted(VALID_TASKS)}, got {task!r}"
168
+ )
169
+ if num_layers < 1:
170
+ raise InferenceProfileError("num_layers must be at least 1")
171
+
172
+ if task == "layer-decompose":
173
+ if self.inference_profile != LAYER_PROFILE:
174
+ raise InferenceProfileError(
175
+ "layer-decompose requires a layer_decompose checkpoint"
176
+ )
177
+ if not has_reference_image:
178
+ raise InferenceProfileError(
179
+ "layer-decompose requires an input reference image"
180
+ )
181
+ return
182
+
183
+ if self.inference_profile != GENERATION_PROFILE:
184
+ raise InferenceProfileError(
185
+ f"{task} requires a generation_edit checkpoint"
186
+ )
187
+ if num_layers != 1:
188
+ raise InferenceProfileError(
189
+ f"{task} does not support num_layers={num_layers}; expected 1"
190
+ )
191
+ if task == "text-to-image" and has_reference_image:
192
+ raise InferenceProfileError("text-to-image does not accept a reference image")
193
+ if task == "image-edit" and not has_reference_image:
194
+ raise InferenceProfileError("image-edit requires a reference image")
195
+
196
+ def resolve_sampling_parameters(
197
+ self,
198
+ *,
199
+ steps: Optional[int] = None,
200
+ cfg: Optional[float] = None,
201
+ ) -> SamplingParameters:
202
+ defaults = DEFAULT_SAMPLING_PARAMETERS[self.inference_profile]
203
+ resolved_steps = defaults.steps if steps is None else steps
204
+ resolved_cfg = defaults.cfg if cfg is None else cfg
205
+
206
+ if type(resolved_steps) is not int or resolved_steps < 1:
207
+ raise InferenceProfileError("sampling steps must be an integer >= 1")
208
+ if (
209
+ isinstance(resolved_cfg, bool)
210
+ or not isinstance(resolved_cfg, (int, float))
211
+ or not math.isfinite(float(resolved_cfg))
212
+ or resolved_cfg < 0
213
+ ):
214
+ raise InferenceProfileError("CFG must be a finite number >= 0")
215
+
216
+ return SamplingParameters(
217
+ steps=resolved_steps,
218
+ cfg=float(resolved_cfg),
219
+ )
220
+
221
+
222
+ def load_inference_profile(model_directory: Union[str, Path]) -> InferenceProfile:
223
+ """Legacy parser: strict root ``inference_profile.json`` (see module doc)."""
224
+
225
+ model_directory = Path(model_directory)
226
+ profile_path = model_directory / PROFILE_FILENAME
227
+ if not profile_path.is_file():
228
+ raise InferenceProfileError(
229
+ f"checkpoint must contain {PROFILE_FILENAME}: {profile_path}"
230
+ )
231
+ try:
232
+ raw = json.loads(profile_path.read_text(encoding="utf-8"))
233
+ except (OSError, json.JSONDecodeError) as error:
234
+ raise InferenceProfileError(
235
+ f"failed to read checkpoint inference profile {profile_path}: {error}"
236
+ ) from error
237
+ if not isinstance(raw, dict):
238
+ raise InferenceProfileError("checkpoint inference profile must be a JSON object")
239
+ return InferenceProfile.from_dict(raw)
240
+
241
+
242
+ TRANSFORMER_CONFIG_FILENAME = "transformer/config.json"
243
+ VAE_CONFIG_FILENAME = "vae/config.json"
244
+ CAPABILITY_FIELDS = ("alignment_padding_mode", "multi_frame_output")
245
+ QWEN_VAE_CLASS_NAME = "AutoencoderKLQwenImage"
246
+
247
+ # The only valid capability pairs and the runtime family they select.
248
+ CAPABILITY_PROFILES = {
249
+ (ZERO_MASKED_PADDING, False): GENERATION_PROFILE,
250
+ (LEARNED_PADDING, True): LAYER_PROFILE,
251
+ }
252
+
253
+
254
+ def _read_component_config(model_directory: Path, relative: str) -> Mapping[str, Any]:
255
+ config_path = model_directory / relative
256
+ if not config_path.is_file():
257
+ raise InferenceProfileError(
258
+ f"checkpoint is missing component config {relative}: {config_path}"
259
+ )
260
+ try:
261
+ raw = json.loads(config_path.read_text(encoding="utf-8"))
262
+ except (OSError, json.JSONDecodeError) as error:
263
+ raise InferenceProfileError(
264
+ f"failed to read component config {config_path}: {error}"
265
+ ) from error
266
+ if not isinstance(raw, dict):
267
+ raise InferenceProfileError(f"component config {relative} must be a JSON object")
268
+ return raw
269
+
270
+
271
+ def _derive_vae_contract(model_directory: Path) -> tuple[int, str]:
272
+ """(vae_input_channels, vae_sample_mode) from the VAE component config."""
273
+
274
+ config = _read_component_config(model_directory, VAE_CONFIG_FILENAME)
275
+ class_name = config.get("_class_name")
276
+ if class_name != QWEN_VAE_CLASS_NAME:
277
+ raise InferenceProfileError(
278
+ f"unsupported VAE contract: {VAE_CONFIG_FILENAME} _class_name must be "
279
+ f"{QWEN_VAE_CLASS_NAME!r} for argmax reference encoding, got "
280
+ f"{class_name!r}"
281
+ )
282
+ declared = [
283
+ config[key]
284
+ for key in ("input_channels", "in_channels")
285
+ if key in config
286
+ ]
287
+ if not declared:
288
+ raise InferenceProfileError(
289
+ f"{VAE_CONFIG_FILENAME} must declare input_channels or in_channels"
290
+ )
291
+ if len(set(declared)) != 1:
292
+ raise InferenceProfileError(
293
+ f"{VAE_CONFIG_FILENAME} input_channels and in_channels disagree: "
294
+ f"{declared}"
295
+ )
296
+ channels = declared[0]
297
+ if type(channels) is not int or channels != 4:
298
+ raise InferenceProfileError(
299
+ "the supported public families require a 4-channel VAE, got "
300
+ f"{channels!r} in {VAE_CONFIG_FILENAME}"
301
+ )
302
+ return channels, "argmax"
303
+
304
+
305
+ def _profile_from_components(
306
+ model_directory: Path,
307
+ alignment_padding_mode: Any,
308
+ multi_frame_output: Any,
309
+ ) -> InferenceProfile:
310
+ if type(alignment_padding_mode) is not str:
311
+ raise InferenceProfileError(
312
+ f"{TRANSFORMER_CONFIG_FILENAME} alignment_padding_mode must be a "
313
+ f"string, got {alignment_padding_mode!r}"
314
+ )
315
+ if type(multi_frame_output) is not bool:
316
+ raise InferenceProfileError(
317
+ f"{TRANSFORMER_CONFIG_FILENAME} multi_frame_output must be a "
318
+ f"boolean, got {multi_frame_output!r}"
319
+ )
320
+ pair = (alignment_padding_mode, multi_frame_output)
321
+ if pair not in CAPABILITY_PROFILES:
322
+ raise InferenceProfileError(
323
+ f"unsupported capability pair {pair!r} in "
324
+ f"{TRANSFORMER_CONFIG_FILENAME}; expected one of "
325
+ f"{sorted(CAPABILITY_PROFILES)}"
326
+ )
327
+ channels, sample_mode = _derive_vae_contract(model_directory)
328
+ profile = InferenceProfile(
329
+ schema_version=PROFILE_SCHEMA_VERSION,
330
+ inference_profile=CAPABILITY_PROFILES[pair],
331
+ alignment_padding_mode=alignment_padding_mode,
332
+ multi_frame_output=multi_frame_output,
333
+ vae_input_channels=channels,
334
+ vae_sample_mode=sample_mode,
335
+ )
336
+ profile.validate()
337
+ return profile
338
+
339
+
340
+ def load_checkpoint_capabilities(
341
+ model_directory: Union[str, Path],
342
+ ) -> InferenceProfile:
343
+ """Derive the runtime capability from component configs.
344
+
345
+ New packages declare ``alignment_padding_mode`` and
346
+ ``multi_frame_output`` in ``transformer/config.json``; legacy packages
347
+ without them fall back to the strict root ``inference_profile.json``.
348
+ A missing ``transformer/config.json`` means both fields are absent, so
349
+ the legacy path applies. A partial pair is a hard error. When both the
350
+ component fields and the legacy file exist, the component metadata is
351
+ authoritative and the legacy file may only agree with it.
352
+ """
353
+
354
+ model_directory = Path(model_directory)
355
+ transformer_path = model_directory / TRANSFORMER_CONFIG_FILENAME
356
+ transformer_config = (
357
+ _read_component_config(model_directory, TRANSFORMER_CONFIG_FILENAME)
358
+ if transformer_path.is_file()
359
+ else {}
360
+ )
361
+ present = [field for field in CAPABILITY_FIELDS if field in transformer_config]
362
+ if len(present) == 1:
363
+ raise InferenceProfileError(
364
+ f"{TRANSFORMER_CONFIG_FILENAME} carries only {present[0]!r}; "
365
+ "alignment_padding_mode and multi_frame_output must be declared "
366
+ "together"
367
+ )
368
+ if not present:
369
+ return load_inference_profile(model_directory)
370
+
371
+ profile = _profile_from_components(
372
+ model_directory,
373
+ transformer_config["alignment_padding_mode"],
374
+ transformer_config["multi_frame_output"],
375
+ )
376
+ legacy_path = model_directory / PROFILE_FILENAME
377
+ if legacy_path.is_file():
378
+ legacy = load_inference_profile(model_directory)
379
+ if legacy != profile:
380
+ raise InferenceProfileError(
381
+ f"component configs disagree with legacy {PROFILE_FILENAME}: "
382
+ f"transformer/vae derive {profile!r} but the profile declares "
383
+ f"{legacy!r}"
384
+ )
385
+ return profile
386
+
387
+
388
+ def resolve_model_directory(
389
+ model_name_or_path: Union[str, Path],
390
+ *,
391
+ revision: Optional[str] = None,
392
+ cache_dir: Optional[Union[str, Path]] = None,
393
+ local_files_only: bool = False,
394
+ token: Optional[Union[str, bool]] = None,
395
+ ) -> Path:
396
+ """Resolve a local directory or materialize a Hugging Face Hub snapshot."""
397
+
398
+ candidate = Path(model_name_or_path).expanduser()
399
+ if candidate.is_dir():
400
+ return candidate.resolve()
401
+ if candidate.exists():
402
+ raise ValueError(f"model path must be a directory: {candidate}")
403
+ if candidate.is_absolute():
404
+ raise FileNotFoundError(f"local model directory does not exist: {candidate}")
405
+
406
+ try:
407
+ from huggingface_hub import snapshot_download
408
+ except ImportError as error:
409
+ raise RuntimeError(
410
+ "huggingface_hub is required when --model is a Hub repository ID"
411
+ ) from error
412
+
413
+ snapshot_path = snapshot_download(
414
+ repo_id=str(model_name_or_path),
415
+ revision=revision,
416
+ cache_dir=str(cache_dir) if cache_dir is not None else None,
417
+ local_files_only=local_files_only,
418
+ token=token,
419
+ )
420
+ return Path(snapshot_path).resolve()
code/mllm_device_map.py ADDED
@@ -0,0 +1,146 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Memory-aware device placement for the MLLM inference frontend."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import math
7
+ from dataclasses import dataclass
8
+ from pathlib import Path
9
+
10
+
11
+ DEFAULT_GPU0_RESERVED_LAYER_EQUIVALENTS = 4
12
+
13
+
14
+ class MLLMDeviceMapError(ValueError):
15
+ """Raised when a requested MLLM layout is incomplete or unsafe."""
16
+
17
+
18
+ @dataclass(frozen=True)
19
+ class MLLMDevicePlan:
20
+ num_hidden_layers: int
21
+ n_gpu: int
22
+ gpu0_reserved_layer_equivalents: int
23
+ layer_counts: tuple[int, ...]
24
+ layer_devices: tuple[int, ...]
25
+ device_map: dict[str, int]
26
+
27
+
28
+ def load_mllm_num_hidden_layers(model_directory: str | Path) -> int:
29
+ config_path = Path(model_directory) / "config.json"
30
+ try:
31
+ with config_path.open(encoding="utf-8") as handle:
32
+ config = json.load(handle)
33
+ except (OSError, json.JSONDecodeError) as exc:
34
+ raise MLLMDeviceMapError(f"cannot read MLLM config {config_path}: {exc}") from None
35
+
36
+ llm_config = config.get("llm_config")
37
+ nested = llm_config.get("num_hidden_layers") if isinstance(llm_config, dict) else None
38
+ root = config.get("num_hidden_layers")
39
+ if nested is not None and root is not None and nested != root:
40
+ raise MLLMDeviceMapError(
41
+ "ambiguous decoder depth: "
42
+ f"llm_config.num_hidden_layers={nested!r}, num_hidden_layers={root!r}"
43
+ )
44
+ value = nested if nested is not None else root
45
+ if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
46
+ raise MLLMDeviceMapError(
47
+ f"missing or invalid MLLM num_hidden_layers in {config_path}: {value!r}"
48
+ )
49
+ return value
50
+
51
+
52
+ def allocate_mllm_layer_counts(
53
+ num_hidden_layers: int,
54
+ n_gpu: int,
55
+ *,
56
+ gpu0_reserved_layer_equivalents: int = DEFAULT_GPU0_RESERVED_LAYER_EQUIVALENTS,
57
+ ) -> tuple[int, ...]:
58
+ if (
59
+ isinstance(n_gpu, bool)
60
+ or not isinstance(n_gpu, int)
61
+ or n_gpu <= 0
62
+ ):
63
+ raise MLLMDeviceMapError(f"n_gpu must be a positive integer, got {n_gpu!r}")
64
+ if isinstance(num_hidden_layers, bool) or not isinstance(num_hidden_layers, int) or num_hidden_layers <= 0:
65
+ raise MLLMDeviceMapError(f"num_hidden_layers must be positive, got {num_hidden_layers!r}")
66
+ if (
67
+ isinstance(gpu0_reserved_layer_equivalents, bool)
68
+ or not isinstance(gpu0_reserved_layer_equivalents, int)
69
+ or gpu0_reserved_layer_equivalents < 0
70
+ ):
71
+ raise MLLMDeviceMapError(
72
+ "gpu0_reserved_layer_equivalents must be a non-negative integer"
73
+ )
74
+
75
+ if n_gpu == 1:
76
+ # Single-GPU plan: every decoder layer and every fixed image module
77
+ # shares logical device 0, so nothing is kept sharded.
78
+ return (num_hidden_layers,)
79
+
80
+ effective_target = max(
81
+ gpu0_reserved_layer_equivalents,
82
+ math.ceil((num_hidden_layers + gpu0_reserved_layer_equivalents) / n_gpu),
83
+ )
84
+ gpu0_layers = min(
85
+ num_hidden_layers,
86
+ max(0, effective_target - gpu0_reserved_layer_equivalents),
87
+ )
88
+ base, remainder = divmod(num_hidden_layers - gpu0_layers, n_gpu - 1)
89
+ other_counts = [base] * (n_gpu - 1)
90
+ for index in range(len(other_counts) - remainder, len(other_counts)):
91
+ other_counts[index] += 1
92
+ counts = (gpu0_layers, *other_counts)
93
+ if sum(counts) != num_hidden_layers:
94
+ raise AssertionError(f"invalid internal MLLM allocation: {counts}")
95
+ return counts
96
+
97
+
98
+ def build_mllm_device_plan(
99
+ num_hidden_layers: int,
100
+ n_gpu: int,
101
+ *,
102
+ gpu0_reserved_layer_equivalents: int = DEFAULT_GPU0_RESERVED_LAYER_EQUIVALENTS,
103
+ ) -> MLLMDevicePlan:
104
+ counts = allocate_mllm_layer_counts(
105
+ num_hidden_layers,
106
+ n_gpu,
107
+ gpu0_reserved_layer_equivalents=gpu0_reserved_layer_equivalents,
108
+ )
109
+ layer_devices = tuple(
110
+ device for device, count in enumerate(counts) for _ in range(count)
111
+ )
112
+ device_map = {
113
+ f"model.model.layers.{layer_index}": device
114
+ for layer_index, device in enumerate(layer_devices)
115
+ }
116
+ # Image conditioning and diffusion modules are attached after the base
117
+ # checkpoint load and are intentionally placed on logical CUDA device 0.
118
+ device_map.update(
119
+ {
120
+ "vision": 0,
121
+ "linear_proj": 0,
122
+ "model.model.word_embeddings.weight": 0,
123
+ "model.model.norm.weight": 0,
124
+ "model.lm_head.weight": 0,
125
+ "model.model.norm": 0,
126
+ }
127
+ )
128
+ return MLLMDevicePlan(
129
+ num_hidden_layers=num_hidden_layers,
130
+ n_gpu=n_gpu,
131
+ gpu0_reserved_layer_equivalents=gpu0_reserved_layer_equivalents,
132
+ layer_counts=counts,
133
+ layer_devices=layer_devices,
134
+ device_map=device_map,
135
+ )
136
+
137
+
138
+ def validate_loaded_layer_devices(
139
+ actual_devices: list[int], plan: MLLMDevicePlan
140
+ ) -> None:
141
+ actual = tuple(actual_devices)
142
+ if actual != plan.layer_devices:
143
+ raise MLLMDeviceMapError(
144
+ "loaded MLLM layer placement disagrees with the requested plan: "
145
+ f"actual={actual}, expected={plan.layer_devices}"
146
+ )
code/modeling_bailing_moe_v2.py ADDED
@@ -0,0 +1,2031 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2023 Antgroup and The HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
5
+ # and OPT implementations in this library. It has been modified from its
6
+ # original forms to accommodate minor architectural differences compared
7
+ # to GPT-NeoX and OPT used by the Meta AI team that trained the model.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ """PyTorch BailingMoE model."""
21
+ import math
22
+ import warnings
23
+ import os
24
+ from typing import List, Optional, Tuple, Union
25
+
26
+ import torch
27
+ import torch.nn.functional as F
28
+ import torch.utils.checkpoint
29
+ from torch import nn
30
+ from torch.nn import CrossEntropyLoss
31
+ from transformers.activations import ACT2FN
32
+ from transformers.cache_utils import Cache, DynamicCache
33
+ from transformers.modeling_attn_mask_utils import (
34
+ AttentionMaskConverter,
35
+ _prepare_4d_attention_mask,
36
+ _prepare_4d_causal_attention_mask,
37
+ _prepare_4d_causal_attention_mask_for_sdpa,
38
+ )
39
+ from transformers.modeling_outputs import (
40
+ MoeModelOutputWithPast,
41
+ MoeCausalLMOutputWithPast,
42
+ )
43
+ from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update
44
+ from transformers.modeling_utils import PreTrainedModel
45
+ from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS, is_torch_greater_or_equal_than_1_13
46
+ from transformers.utils import (
47
+ add_start_docstrings,
48
+ add_start_docstrings_to_model_forward,
49
+ is_flash_attn_2_available,
50
+ is_flash_attn_greater_or_equal_2_10,
51
+ logging,
52
+ replace_return_docstrings,
53
+ )
54
+ from transformers.utils.import_utils import is_torch_fx_available
55
+ from configuration_bailing_moe_v2 import BailingMoeV2Config
56
+ from transformers.generation.utils import GenerationMixin
57
+ from modeling_utils import patch_continuous_features, build_modality_mask
58
+
59
+
60
+
61
+ if is_flash_attn_2_available():
62
+ from flash_attn import flash_attn_func, flash_attn_varlen_func
63
+ from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa
64
+
65
+
66
+ # This makes `_prepare_4d_causal_attention_mask` a leaf function in the FX graph.
67
+ # It means that the function will not be traced through and simply appear as a node in the graph.
68
+ if is_torch_fx_available():
69
+ if not is_torch_greater_or_equal_than_1_13:
70
+ import torch.fx
71
+
72
+ _prepare_4d_causal_attention_mask = torch.fx.wrap(_prepare_4d_causal_attention_mask)
73
+
74
+
75
+ logger = logging.get_logger(__name__)
76
+
77
+ _CONFIG_FOR_DOC = "BailingMoeV2Config"
78
+
79
+
80
+ def _get_unpad_data(attention_mask):
81
+ seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
82
+ indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
83
+ max_seqlen_in_batch = seqlens_in_batch.max().item()
84
+ cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.torch.int32), (1, 0))
85
+ return (
86
+ indices,
87
+ cu_seqlens,
88
+ max_seqlen_in_batch,
89
+ )
90
+
91
+
92
+ def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):
93
+ warnings.warn(
94
+ "Calling `transformers.models.BailingMoeV2.modeling_BailingMoeV2._prepare_4d_attention_mask` is deprecated and will be removed in v4.37. Use `transformers.modeling_attn_mask_utils._prepare_4d_attention_mask"
95
+ )
96
+ return _prepare_4d_attention_mask(mask=mask, dtype=dtype, tgt_len=tgt_len)
97
+
98
+
99
+ def _make_causal_mask(
100
+ input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 0
101
+ ):
102
+ warnings.warn(
103
+ "Calling `transformers.models.BailingMoeV2.modeling_BailingMoeV2._make_causal_mask` is deprecated and will be removed in v4.37. Use `transformers.models.BailingMoeV2.modeling_BailingMoeV2.AttentionMaskConverter._make_causal_mask"
104
+ )
105
+ return AttentionMaskConverter._make_causal_mask(
106
+ input_ids_shape=input_ids_shape, dtype=dtype, device=device, past_key_values_length=past_key_values_length
107
+ )
108
+
109
+
110
+ class BailingMoeV2RMSNorm(nn.Module):
111
+ def __init__(self, hidden_size, eps=1e-6):
112
+ """
113
+ BailingMoeV2RMSNorm is equivalent to T5LayerNorm
114
+ """
115
+ super().__init__()
116
+ self.weight = nn.Parameter(torch.ones(hidden_size))
117
+ self.variance_epsilon = eps
118
+
119
+ def forward(self, hidden_states):
120
+ input_dtype = hidden_states.dtype
121
+ hidden_states = hidden_states.to(torch.float32)
122
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
123
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
124
+ return self.weight * hidden_states.to(input_dtype)
125
+
126
+ def reset_parameters(self):
127
+ nn.init.ones_(self.weight) # explicit reset for enabling fsdp
128
+
129
+
130
+ ALL_LAYERNORM_LAYERS.append(BailingMoeV2RMSNorm)
131
+
132
+
133
+ class BailingMoeV2RotaryEmbedding(nn.Module):
134
+ def __init__(self, config: BailingMoeV2Config, device=None):
135
+ super().__init__()
136
+ # BC: "rope_type" was originally "type"
137
+ if hasattr(config, "rope_scaling") and config.rope_scaling is not None:
138
+ self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type"))
139
+ else:
140
+ self.rope_type = "default"
141
+ self.max_seq_len_cached = config.max_position_embeddings
142
+ self.original_max_seq_len = config.max_position_embeddings
143
+
144
+ self.config = config
145
+ self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
146
+
147
+ inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
148
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
149
+ self.original_inv_freq = self.inv_freq
150
+
151
+ @torch.no_grad()
152
+ @dynamic_rope_update # power user: used with advanced RoPE types (e.g. dynamic rope)
153
+ def forward(self, x, position_ids):
154
+ position_ids = position_ids[:, 0, :]
155
+ inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(x.device)
156
+ position_ids_expanded = position_ids[:, None, :].float()
157
+
158
+
159
+ device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
160
+ with torch.autocast(device_type=device_type, enabled=False): # Force float32
161
+ freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
162
+ emb = torch.cat((freqs, freqs), dim=-1)
163
+ cos = emb.cos() * self.attention_scaling
164
+ sin = emb.sin() * self.attention_scaling
165
+
166
+ return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
167
+
168
+ def reset_parameters(self) -> None:
169
+ new_inv_freq, self.attention_scaling = self.rope_init_fn(self.config, None)
170
+ self.inv_freq.copy_(new_inv_freq)
171
+
172
+
173
+ class BailingMoeV2RotaryEmbedding3D(nn.Module):
174
+ def __init__(self, config: BailingMoeV2Config, device=None):
175
+ super().__init__()
176
+ # BC: "rope_type" was originally "type"
177
+ self.rope_init_type = "default"
178
+ self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type"))
179
+ self.max_seq_len_cached = config.max_position_embeddings
180
+ self.original_max_seq_len = config.max_position_embeddings
181
+
182
+ self.config = config
183
+ self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_init_type]
184
+
185
+ inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
186
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
187
+ self.original_inv_freq = self.inv_freq
188
+ self.config = config
189
+
190
+ @torch.no_grad()
191
+ @dynamic_rope_update # power user: used with advanced RoPE types (e.g. dynamic rope)
192
+ def forward(self, x, position_ids):
193
+ if self.rope_type == "3D" or self.rope_type == "video_rope":
194
+ inv_freq_expanded = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1).to(x.device)
195
+ position_ids_expanded = position_ids[:, :, None, :].float()
196
+ else:
197
+ inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(x.device)
198
+ position_ids_expanded = position_ids[:, None, :].float()
199
+
200
+ device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
201
+ with torch.autocast(device_type=device_type, enabled=False): # Force float32
202
+ if self.rope_type == "3D" or self.rope_type == "video_rope":
203
+ freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3)
204
+ else:
205
+ freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
206
+ emb = torch.cat((freqs, freqs), dim=-1)
207
+ cos = emb.cos() * self.attention_scaling
208
+ sin = emb.sin() * self.attention_scaling
209
+
210
+ return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
211
+
212
+ def reset_parameters(self) -> None:
213
+ new_inv_freq, self.attention_scaling = self.rope_init_fn(self.config, None)
214
+ self.inv_freq.copy_(new_inv_freq)
215
+
216
+
217
+ # Copied from transformers.models.llama.modeling_llama.rotate_half
218
+ def rotate_half(x):
219
+ """Rotates half the hidden dims of the input."""
220
+ x1 = x[..., : x.shape[-1] // 2]
221
+ x2 = x[..., x.shape[-1] // 2 :]
222
+ return torch.cat((-x2, x1), dim=-1)
223
+
224
+
225
+ def apply_3d_rotary_pos_emb(
226
+ q,
227
+ k,
228
+ cos,
229
+ sin,
230
+ mrope_section=[8, 12, 12],
231
+ unsqueeze_dim=1,
232
+ rope_type="m_rope",
233
+ rotary_half=True
234
+ ):
235
+ """Applies Rotary Position Embedding with Multimodal Sections to the query and key tensors (https://qwenlm.github.io/blog/qwen2-vl/).
236
+ Explanation:
237
+ Multimodal 3D rotary position embedding is an extension to 1D rotary position embedding. The input embedding
238
+ sequence contains vision (images / videos) embedding and text embedding or just contains text embedding. For
239
+ vision embedding part, we apply rotary position embedding on temporal, height and width dimension separately.
240
+ Here we split the channel dimension to 3 chunks for the temporal, height and width rotary position embedding.
241
+ For text embedding part, we just apply 1D rotary position embedding. The three rotary position index (temporal,
242
+ height and width) of text embedding is always the same, so the text embedding rotary position embedding has no
243
+ difference with modern LLMs.
244
+ Args:
245
+ q (`torch.Tensor`): The query tensor.
246
+ k (`torch.Tensor`): The key tensor.
247
+ cos (`torch.Tensor`): The cosine part of the rotary embedding.
248
+ sin (`torch.Tensor`): The sine part of the rotary embedding.
249
+ mrope_section(`List(int)`):
250
+ Multimodal rope section is for channel dimension of temporal, height and width in rope calculation.
251
+ unsqueeze_dim (`int`, *optional*, defaults to 1):
252
+ The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
253
+ sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
254
+ that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
255
+ k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
256
+ cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
257
+ the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
258
+ rope_type (`str`, *optional*, defaults to "m_rope"):
259
+ rotary_half (`bool`, *optional*, defaults to `False`): Keep half or full tensor for later concatenation
260
+ Returns:
261
+ `tuple(torch.Tensor)` comprising the query and key tensors rotated using the Rotary Position Embedding.
262
+ """
263
+ if rope_type == "3D": # rename rope_type
264
+ rope_type = "m_rope"
265
+
266
+ if rope_type == "m_rope":
267
+ mrope_section = mrope_section * 2
268
+ cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1).unsqueeze(unsqueeze_dim).to(q.device)
269
+ sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1).unsqueeze(unsqueeze_dim).to(q.device)
270
+ elif rope_type == "video_rope":
271
+ mrope_section = list(mrope_section)
272
+ mrope_section = [mrope_section[0], mrope_section[1] + mrope_section[2]]
273
+ mrope_section = mrope_section * 2
274
+ # adjust t last -> (48, 16, 48, 16)
275
+ mrope_section = mrope_section[::-1]
276
+ index = 0
277
+ result_cos = []
278
+ result_sin = []
279
+ # get x1, y1, x2, y2, ..., t1, t2, ...
280
+ for i, section in enumerate(mrope_section):
281
+ if i % 2 == 0:
282
+ for j in range(section):
283
+ row = 1 if j % 2 == 0 else 2
284
+ result_cos.append(cos[row, ..., index: index + 1])
285
+ result_sin.append(sin[row, ..., index: index + 1])
286
+ index += 1
287
+ else:
288
+ result_cos.append(cos[0, ..., index:index + section])
289
+ result_sin.append(sin[0, ..., index:index + section])
290
+ index += section
291
+ cos, sin = torch.cat(result_cos, dim=-1).unsqueeze(dim=unsqueeze_dim).to(q.device), torch.cat(result_sin, dim=-1).unsqueeze(dim=unsqueeze_dim).to(q.device)
292
+ else: # vanilla rope for llm
293
+ cos = cos.unsqueeze(unsqueeze_dim).to(q.device)
294
+ sin = sin.unsqueeze(unsqueeze_dim).to(q.device)
295
+
296
+ if rotary_half:
297
+ rotary_dim = cos.shape[-1]
298
+ q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
299
+ k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
300
+
301
+ # Apply rotary embeddings on the first half
302
+ q_embed = (q_rot * cos) + (rotate_half(q_rot) * sin)
303
+ k_embed = (k_rot * cos) + (rotate_half(k_rot) * sin)
304
+
305
+ # Concatenate back to full shape
306
+ q_embed = torch.cat([q_embed, q_pass], dim=-1)
307
+ k_embed = torch.cat([k_embed, k_pass], dim=-1)
308
+ else:
309
+ q_embed = (q * cos) + (rotate_half(q) * sin)
310
+ k_embed = (k * cos) + (rotate_half(k) * sin)
311
+ return q_embed, k_embed
312
+
313
+ def get_t_scale_rope_index(
314
+ config,
315
+ input_ids: torch.LongTensor,
316
+ image_grid_thw: Optional[torch.LongTensor] = None,
317
+ video_grid_thw: Optional[torch.LongTensor] = None,
318
+ attention_mask: Optional[torch.Tensor] = None,
319
+ scale_factor: float = 1.0,
320
+ second_per_grid_ts: Optional[torch.Tensor] = None,
321
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
322
+ spatial_merge_size = config.spatial_merge_size
323
+ image_token_id = config.image_patch_token
324
+ video_token_id = config.video_patch_token
325
+ image_start_token_id = config.image_start_token
326
+ video_start_token_id = config.video_start_token
327
+ use_abs_time_pos = second_per_grid_ts is not None
328
+
329
+ mrope_position_deltas = []
330
+ if image_grid_thw is not None or video_grid_thw is not None:
331
+ total_input_ids = input_ids
332
+ if attention_mask is None:
333
+ attention_mask = torch.ones_like(total_input_ids)
334
+ position_ids = torch.ones(
335
+ 3,
336
+ input_ids.shape[0],
337
+ input_ids.shape[1],
338
+ dtype=input_ids.dtype,
339
+ device=input_ids.device,
340
+ )
341
+
342
+ image_index, video_index = 0, 0
343
+ attention_mask = attention_mask.to(total_input_ids.device)
344
+
345
+ for i, input_ids in enumerate(total_input_ids):
346
+ if attention_mask is not None:
347
+ input_ids = input_ids[attention_mask[i] == 1]
348
+ image_nums, video_nums = 0, 0
349
+ if image_grid_thw is not None:
350
+ vision_start_indices = torch.argwhere(input_ids == image_start_token_id).squeeze(1)
351
+ vision_tokens = input_ids[vision_start_indices + 1]
352
+ image_nums = (vision_tokens == image_token_id).sum()
353
+ if video_grid_thw is not None:
354
+ vision_start_indices = torch.argwhere(input_ids == video_start_token_id).squeeze(1)
355
+ vision_tokens = input_ids[vision_start_indices + 1]
356
+ video_nums = (vision_tokens == video_token_id).sum()
357
+
358
+ input_tokens = input_ids.tolist()
359
+ llm_pos_ids_list: list = []
360
+ st = 0
361
+ remain_images, remain_videos = image_nums, video_nums
362
+ for _ in range(image_nums + video_nums):
363
+ if image_token_id in input_tokens and remain_images > 0:
364
+ ed_image = input_tokens.index(image_token_id, st)
365
+ else:
366
+ ed_image = len(input_tokens) + 1
367
+ if video_token_id in input_tokens and remain_videos > 0:
368
+ ed_video = input_tokens.index(video_token_id, st)
369
+ else:
370
+ ed_video = len(input_tokens) + 1
371
+ if ed_image < ed_video:
372
+ t, h, w = (
373
+ image_grid_thw[image_index][0],
374
+ image_grid_thw[image_index][1],
375
+ image_grid_thw[image_index][2],
376
+ )
377
+ second_per_grid_t = 0
378
+ image_index += 1
379
+ remain_images -= 1
380
+ ed = ed_image
381
+ else:
382
+ t, h, w = (
383
+ video_grid_thw[video_index][0],
384
+ video_grid_thw[video_index][1],
385
+ video_grid_thw[video_index][2],
386
+ )
387
+ if second_per_grid_ts is not None:
388
+ second_per_grid_t = second_per_grid_ts[video_index]
389
+ else:
390
+ second_per_grid_t = 1.0
391
+ video_index += 1
392
+ remain_videos -= 1
393
+ ed = ed_video
394
+ llm_grid_t, llm_grid_h, llm_grid_w = (
395
+ t.item(),
396
+ h.item() // spatial_merge_size,
397
+ w.item() // spatial_merge_size,
398
+ )
399
+ text_len = ed - st
400
+
401
+ st_idx = llm_pos_ids_list[-1][0].max() + 1 if len(llm_pos_ids_list) > 0 else 0
402
+ llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)
403
+
404
+ # body-diagonal symmetry
405
+ t_index = torch.arange(llm_grid_t).view(-1, 1).expand(
406
+ -1, llm_grid_h * llm_grid_w).flatten()
407
+ h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(
408
+ llm_grid_t, -1, llm_grid_w).flatten() - (llm_grid_h - 1) // 2
409
+ w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(
410
+ llm_grid_t, llm_grid_h, -1).flatten() - (llm_grid_w - 1) // 2
411
+
412
+ # time dim adjust step size
413
+ if use_abs_time_pos:
414
+ t_index = t_index * second_per_grid_t * scale_factor
415
+ else:
416
+ t_index = t_index * scale_factor
417
+ t_index = t_index + text_len + st_idx
418
+
419
+ h_index = h_index + t_index
420
+ w_index = w_index + t_index
421
+ llm_pos_ids_list.append(
422
+ torch.stack([t_index, h_index, w_index]))
423
+ st = ed + llm_grid_t * llm_grid_h * llm_grid_w
424
+
425
+ if st < len(input_tokens):
426
+ # next text token near last video token position = last t + 1
427
+ st_idx = llm_pos_ids_list[-1][0].max() + 1 if len(llm_pos_ids_list) > 0 else 0
428
+ text_len = len(input_tokens) - st
429
+ llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)
430
+
431
+ llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)
432
+ llm_positions = llm_positions.to(dtype=position_ids.dtype, device=position_ids.device)
433
+ position_ids[..., i, attention_mask[i] == 1] = llm_positions
434
+ # generate first token = last t + 1
435
+ mrope_position_deltas.append(llm_positions[0].max() + 1 - len(total_input_ids[i]))
436
+ mrope_position_deltas = torch.tensor(mrope_position_deltas, device=input_ids.device).unsqueeze(1)
437
+ else:
438
+ if attention_mask is not None:
439
+ position_ids = attention_mask.long().cumsum(-1) - 1
440
+ position_ids.masked_fill_(attention_mask == 0, 1)
441
+ position_ids = position_ids.unsqueeze(0).expand(3, -1, -1).to(input_ids.device)
442
+ max_position_ids = position_ids.max(0, keepdim=False)[0].max(-1, keepdim=True)[0]
443
+ mrope_position_deltas = max_position_ids + 1 - attention_mask.shape[-1]
444
+ else:
445
+ position_ids = (
446
+ torch.arange(input_ids.shape[1], device=input_ids.device)
447
+ .view(1, 1, -1)
448
+ .expand(3, input_ids.shape[0], -1)
449
+ )
450
+ mrope_position_deltas = torch.zeros(
451
+ [input_ids.shape[0], 1],
452
+ device=input_ids.device,
453
+ dtype=input_ids.dtype,
454
+ )
455
+
456
+ return position_ids, mrope_position_deltas
457
+
458
+ def get_rope_index(
459
+ config,
460
+ input_ids: Optional[torch.LongTensor] = None,
461
+ image_grid_thw: Optional[torch.LongTensor] = None,
462
+ video_grid_thw: Optional[torch.LongTensor] = None,
463
+ attention_mask: Optional[torch.Tensor] = None,
464
+ second_per_grid_ts: Optional[torch.Tensor] = None,
465
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
466
+ """
467
+ Calculate the 3D rope index based on image and video's temporal, height and width in LLM.
468
+
469
+ Explanation:
470
+ Each embedding sequence contains vision embedding and text embedding or just contains text embedding.
471
+
472
+ For pure text embedding sequence, the rotary position embedding has no difference with modern LLMs.
473
+ Examples:
474
+ input_ids: [T T T T T], here T is for text.
475
+ temporal position_ids: [0, 1, 2, 3, 4]
476
+ height position_ids: [0, 1, 2, 3, 4]
477
+ width position_ids: [0, 1, 2, 3, 4]
478
+
479
+ For vision and text embedding sequence, we calculate 3D rotary position embedding for vision part
480
+ and 1D rotary position embeddin for text part.
481
+ Examples:
482
+ Temporal (Time): 3 patches, representing different segments of the video in time.
483
+ Height: 2 patches, dividing each frame vertically.
484
+ Width: 2 patches, dividing each frame horizontally.
485
+ We also have some important parameters:
486
+ fps (Frames Per Second): The video's frame rate, set to 1. This means one frame is processed each second.
487
+ tokens_per_second: This is a crucial parameter. It dictates how many "time-steps" or "temporal tokens" are conceptually packed into a one-second interval of the video. In this case, we have 25 tokens per second. So each second of the video will be represented with 25 separate time points. It essentially defines the temporal granularity.
488
+ temporal_patch_size: The number of frames that compose one temporal patch. Here, it's 2 frames.
489
+ interval: The step size for the temporal position IDs, calculated as tokens_per_second * temporal_patch_size / fps. In this case, 25 * 2 / 1 = 50. This means that each temporal patch will be have a difference of 50 in the temporal position IDs.
490
+ input_ids: [V V V V V V V V V V V V T T T T T], here V is for vision.
491
+ vision temporal position_ids: [0, 0, 0, 0, 50, 50, 50, 50, 100, 100, 100, 100]
492
+ vision height position_ids: [0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1]
493
+ vision width position_ids: [0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1]
494
+ text temporal position_ids: [101, 102, 103, 104, 105]
495
+ text height position_ids: [101, 102, 103, 104, 105]
496
+ text width position_ids: [101, 102, 103, 104, 105]
497
+ Here we calculate the text start position_ids as the max vision position_ids plus 1.
498
+
499
+ Args:
500
+ input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
501
+ Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
502
+ it.
503
+ image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
504
+ The temporal, height and width of feature shape of each image in LLM.
505
+ video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
506
+ The temporal, height and width of feature shape of each video in LLM.
507
+ second_per_grid_ts (`torch.Tensor` of shape `(num_videos)`, *optional*):
508
+ The time interval (in seconds) for each grid along the temporal dimension in the 3D position IDs.
509
+ attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
510
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
511
+
512
+ - 1 for tokens that are **not masked**,
513
+ - 0 for tokens that are **masked**.
514
+
515
+ Returns:
516
+ position_ids (`torch.LongTensor` of shape `(3, batch_size, sequence_length)`)
517
+ mrope_position_deltas (`torch.Tensor` of shape `(batch_size)`)
518
+ """
519
+ spatial_merge_size = config.spatial_merge_size
520
+ image_token_id = config.image_patch_token
521
+ video_token_id = config.video_patch_token
522
+ image_start_token_id = config.image_start_token
523
+ video_start_token_id = config.video_start_token
524
+
525
+ use_abs_time_pos = second_per_grid_ts is not None
526
+
527
+ mrope_position_deltas = []
528
+ if image_grid_thw is not None or video_grid_thw is not None:
529
+ total_input_ids = input_ids
530
+ if attention_mask is None:
531
+ attention_mask = torch.ones_like(total_input_ids)
532
+ position_ids = torch.ones(
533
+ 3,
534
+ input_ids.shape[0],
535
+ input_ids.shape[1],
536
+ dtype=input_ids.dtype,
537
+ device=input_ids.device,
538
+ )
539
+ image_index, video_index = 0, 0
540
+ attention_mask = attention_mask.to(total_input_ids.device)
541
+ for i, input_ids in enumerate(total_input_ids):
542
+ input_ids = input_ids[attention_mask[i] == 1]
543
+ image_nums, video_nums = 0, 0
544
+ if image_grid_thw is not None:
545
+ vision_start_indices = torch.argwhere(input_ids == image_start_token_id).squeeze(1)
546
+ vision_tokens = input_ids[vision_start_indices + 1]
547
+ image_nums = (vision_tokens == image_token_id).sum()
548
+ if video_grid_thw is not None:
549
+ vision_start_indices = torch.argwhere(input_ids == video_start_token_id).squeeze(1)
550
+ vision_tokens = input_ids[vision_start_indices + 1]
551
+ video_nums = (vision_tokens == video_token_id).sum()
552
+
553
+ input_tokens = input_ids.tolist()
554
+ llm_pos_ids_list: list = []
555
+ st = 0
556
+ remain_images, remain_videos = image_nums, video_nums
557
+ for _ in range(image_nums + video_nums):
558
+ if image_token_id in input_tokens and remain_images > 0:
559
+ ed_image = input_tokens.index(image_token_id, st)
560
+ else:
561
+ ed_image = len(input_tokens) + 1
562
+ if video_token_id in input_tokens and remain_videos > 0:
563
+ ed_video = input_tokens.index(video_token_id, st)
564
+ else:
565
+ ed_video = len(input_tokens) + 1
566
+ if ed_image < ed_video:
567
+ t, h, w = (
568
+ image_grid_thw[image_index][0],
569
+ image_grid_thw[image_index][1],
570
+ image_grid_thw[image_index][2],
571
+ )
572
+ second_per_grid_t = 0
573
+ image_index += 1
574
+ remain_images -= 1
575
+ ed = ed_image
576
+
577
+ else:
578
+ t, h, w = (
579
+ video_grid_thw[video_index][0],
580
+ video_grid_thw[video_index][1],
581
+ video_grid_thw[video_index][2],
582
+ )
583
+ if second_per_grid_ts is not None:
584
+ second_per_grid_t = second_per_grid_ts[video_index]
585
+ else:
586
+ second_per_grid_t = 1.0
587
+ video_index += 1
588
+ remain_videos -= 1
589
+ ed = ed_video
590
+ llm_grid_t, llm_grid_h, llm_grid_w = (
591
+ t.item(),
592
+ h.item() // spatial_merge_size,
593
+ w.item() // spatial_merge_size,
594
+ )
595
+ text_len = ed - st
596
+
597
+ st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0
598
+ llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)
599
+
600
+ range_tensor = torch.arange(llm_grid_t).view(-1, 1)
601
+ expanded_range = range_tensor.expand(-1, llm_grid_h * llm_grid_w)
602
+ if use_abs_time_pos:
603
+ time_tensor = expanded_range * second_per_grid_t * config.tokens_per_second
604
+ time_tensor_long = time_tensor.long()
605
+ else:
606
+ time_tensor_long = expanded_range.long()
607
+ t_index = time_tensor_long.flatten()
608
+
609
+ h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()
610
+ w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()
611
+ llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx)
612
+ st = ed + llm_grid_t * llm_grid_h * llm_grid_w
613
+
614
+ if st < len(input_tokens):
615
+ st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0
616
+ text_len = len(input_tokens) - st
617
+ llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)
618
+
619
+ llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)
620
+ position_ids[..., i, attention_mask[i] == 1] = llm_positions.to(position_ids.device)
621
+ mrope_position_deltas.append(llm_positions.max() + 1 - len(total_input_ids[i]))
622
+ mrope_position_deltas = torch.tensor(mrope_position_deltas, device=input_ids.device).unsqueeze(1)
623
+ else:
624
+ if attention_mask is not None:
625
+ position_ids = attention_mask.long().cumsum(-1) - 1
626
+ position_ids.masked_fill_(attention_mask == 0, 1)
627
+ position_ids = position_ids.unsqueeze(0).expand(3, -1, -1).to(input_ids.device)
628
+ max_position_ids = position_ids.max(0, keepdim=False)[0].max(-1, keepdim=True)[0]
629
+ mrope_position_deltas = max_position_ids + 1 - attention_mask.shape[-1]
630
+ else:
631
+ position_ids = (
632
+ torch.arange(input_ids.shape[1], device=input_ids.device)
633
+ .view(1, 1, -1)
634
+ .expand(3, input_ids.shape[0], -1)
635
+ )
636
+ mrope_position_deltas = torch.zeros(
637
+ [input_ids.shape[0], 1],
638
+ device=input_ids.device,
639
+ dtype=input_ids.dtype,
640
+ )
641
+
642
+ return position_ids, mrope_position_deltas
643
+
644
+
645
+ class BailingMoeV2MLP(nn.Module):
646
+ def __init__(self, config: BailingMoeV2Config, intermediate_size: int):
647
+ super().__init__()
648
+ self.config = config
649
+ self.hidden_size = config.hidden_size
650
+ self.intermediate_size = intermediate_size
651
+
652
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
653
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
654
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
655
+ self.act_fn = ACT2FN[config.hidden_act]
656
+
657
+ def forward(self, x, image_mask=None, audio_mask=None):
658
+ return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
659
+
660
+
661
+ class BailingMoeV2Gate(nn.Module):
662
+ def __init__(self, config):
663
+ super().__init__()
664
+ self.config = config
665
+ self.top_k = config.num_experts_per_tok
666
+ self.num_experts = config.num_experts
667
+
668
+ self.n_group = config.n_group
669
+ self.topk_group = config.topk_group
670
+
671
+ # topk selection algorithm
672
+ self.gating_dim = config.hidden_size
673
+ self.weight = nn.Parameter(torch.empty((self.num_experts, self.gating_dim)))
674
+ self.routed_scaling_factor = config.routed_scaling_factor
675
+ self.bias_update_coeff = 0.001
676
+
677
+ self.expert_bias = torch.nn.Parameter(torch.zeros(self.num_experts), requires_grad=False)
678
+ self.reset_parameters()
679
+
680
+ def reset_parameters(self) -> None:
681
+ import torch.nn.init as init
682
+
683
+ init.kaiming_uniform_(self.weight, a=math.sqrt(5))
684
+
685
+ def update_bias(self, distribution: torch.Tensor):
686
+ with torch.no_grad():
687
+ delta_bias = (distribution.mean() - distribution).sign()
688
+ self.expert_bias.data = self.expert_bias.data + self.bias_update_coeff * delta_bias
689
+
690
+ def group_limited_topk(
691
+ self,
692
+ scores: torch.Tensor,
693
+ ):
694
+ num_tokens, _ = scores.size()
695
+ # Organize the experts into groups
696
+ group_scores = scores.view(num_tokens, self.n_group, -1).topk(2, dim=-1)[0].sum(dim=-1)
697
+ group_idx = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
698
+ group_mask = torch.zeros_like(group_scores)
699
+ group_mask.scatter_(1, group_idx, 1)
700
+
701
+ # Mask the experts based on selection groups
702
+ score_mask = (
703
+ group_mask.unsqueeze(-1)
704
+ .expand(num_tokens, self.n_group, self.num_experts // self.n_group)
705
+ .reshape(num_tokens, -1)
706
+ )
707
+
708
+ masked_scores = scores.masked_fill(~score_mask.bool(), float('-inf'))
709
+ probs, top_indices = torch.topk(masked_scores, k=self.top_k, dim=-1, sorted=False)
710
+
711
+ return probs, top_indices
712
+
713
+ def forward(self, hidden_states):
714
+ # compute gating score
715
+ hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
716
+ logits = F.linear(hidden_states.type(torch.float32), self.weight.type(torch.float32))
717
+
718
+ scores = torch.sigmoid(logits)
719
+
720
+ scores_for_routing = scores + self.expert_bias
721
+ _, topk_idx = self.group_limited_topk(scores_for_routing)
722
+
723
+ scores = torch.gather(scores, dim=1, index=topk_idx).type_as(logits)
724
+
725
+ topk_weight = scores / (scores.sum(dim=-1, keepdim=True) + 1e-20) if self.top_k > 1 else scores
726
+ topk_weight = topk_weight * self.routed_scaling_factor
727
+
728
+ return topk_idx, topk_weight, logits
729
+
730
+
731
+ class BailingMoeV2SparseMoeBlock(nn.Module):
732
+ """
733
+ A mixed expert module containing shared experts.
734
+ """
735
+
736
+ def __init__(self, config: BailingMoeV2Config):
737
+ super().__init__()
738
+ self.config = config
739
+ self.num_experts_per_tok = config.num_experts_per_tok
740
+ self._setup_experts()
741
+ self.router_type = config.router_type
742
+ if self.router_type == "topN":
743
+ logger.info(f"use topN Router")
744
+ self.gate = BailingMoeV2Gate(config)
745
+ elif self.router_type == "MultiRouter":
746
+ logger.info(f"use MultiRouter")
747
+ self.gate = BailingMoeV2Gate(config)
748
+ self.image_gate = BailingMoeV2Gate(config)
749
+ self.audio_gate = BailingMoeV2Gate(config)
750
+ if config.num_shared_experts is not None:
751
+ self.shared_experts = BailingMoeV2MLP(
752
+ config=config, intermediate_size=config.moe_intermediate_size * config.num_shared_experts
753
+ )
754
+
755
+ def _setup_experts(self):
756
+ self.experts = nn.ModuleList(
757
+ [
758
+ BailingMoeV2MLP(config=self.config, intermediate_size=self.config.moe_intermediate_size)
759
+ for _ in range(self.config.num_experts)
760
+ ]
761
+ )
762
+
763
+ def forward(self, hidden_states, image_mask, audio_mask):
764
+ identity = hidden_states
765
+ bsz, seq_len, h = hidden_states.shape
766
+ if self.router_type == "MultiRouter":
767
+ if image_mask is not None:
768
+ if len(image_mask.shape) == 3:
769
+ assert image_mask.shape[-1] == 1
770
+ elif len(image_mask.shape) == 2:
771
+ assert image_mask.shape == hidden_states.shape[:2]
772
+ image_mask = image_mask.unsqueeze(-1)
773
+ else:
774
+ raise ValueError(
775
+ "unexpected image mask shape: "
776
+ f"{tuple(image_mask.shape)}"
777
+ )
778
+ if audio_mask is not None:
779
+ if len(audio_mask.shape) == 3:
780
+ assert audio_mask.shape[-1] == 1
781
+ elif len(audio_mask.shape) == 2:
782
+ assert audio_mask.shape == hidden_states.shape[:2]
783
+ audio_mask = audio_mask.unsqueeze(-1)
784
+ else:
785
+ raise ValueError(
786
+ "unexpected audio mask shape: "
787
+ f"{tuple(audio_mask.shape)}"
788
+ )
789
+ if image_mask is not None and audio_mask is not None:
790
+ assert torch.logical_and(image_mask, audio_mask).sum() == 0
791
+
792
+ image_topk_idx, image_topk_weight, image_router_logits = self.image_gate(hidden_states)
793
+ audio_topk_idx, audio_topk_weight, audio_router_logits = self.audio_gate(hidden_states)
794
+ topk_idx, topk_weight, router_logits = self.gate(hidden_states)
795
+
796
+ if image_mask is not None:
797
+ image_mask = image_mask.view(-1, 1)
798
+ topk_idx = image_topk_idx * image_mask + topk_idx * torch.logical_not(image_mask)
799
+ topk_weight = image_topk_weight * image_mask + topk_weight * torch.logical_not(image_mask)
800
+ router_logits = image_router_logits * image_mask + router_logits * torch.logical_not(image_mask)
801
+ if audio_mask is not None:
802
+ audio_mask = audio_mask.view(-1, 1)
803
+ audio_mask = audio_mask.to(router_logits.device)
804
+ topk_idx = audio_topk_idx * audio_mask + topk_idx * torch.logical_not(audio_mask)
805
+ topk_weight = audio_topk_weight * audio_mask + topk_weight * torch.logical_not(audio_mask)
806
+ router_logits = audio_router_logits * audio_mask + router_logits * torch.logical_not(audio_mask)
807
+ else:
808
+ topk_idx, topk_weight, router_logits = self.gate(hidden_states)
809
+
810
+ hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
811
+ flat_topk_idx = topk_idx.view(-1)
812
+ y = self.moe_infer(hidden_states, topk_idx, topk_weight).view(bsz, seq_len, h)
813
+ if self.config.num_shared_experts is not None:
814
+ y = y + self.shared_experts(identity)
815
+ return y, (router_logits.view(bsz, seq_len, -1), topk_idx.view(bsz, seq_len, -1))
816
+
817
+ @torch.no_grad()
818
+ def moe_infer(self, x, topk_ids, topk_weight):
819
+ cnts = topk_ids.new_zeros((topk_ids.shape[0], len(self.experts)))
820
+ cnts.scatter_(1, topk_ids, 1)
821
+ tokens_per_expert = cnts.sum(dim=0)
822
+ idxs = topk_ids.view(-1).argsort()
823
+ sorted_tokens = x[idxs // topk_ids.shape[1]]
824
+ sorted_tokens_shape = sorted_tokens.shape
825
+ tokens_per_expert = tokens_per_expert.cpu().numpy()
826
+ outputs = []
827
+ start_idx = 0
828
+ for i, num_tokens in enumerate(tokens_per_expert):
829
+ end_idx = start_idx + num_tokens
830
+ if num_tokens == 0:
831
+ continue
832
+ expert = self.experts[i]
833
+ tokens_for_this_expert = sorted_tokens[start_idx:end_idx]
834
+ expert_out = expert(tokens_for_this_expert)
835
+ outputs.append(expert_out.to(x.device))
836
+ start_idx = end_idx
837
+
838
+ outs = torch.cat(outputs, dim=0) if len(outputs) else sorted_tokens.new_empty(0)
839
+ new_x = torch.empty_like(outs)
840
+ new_x[idxs] = outs
841
+ final_out = (
842
+ new_x.view(*topk_ids.shape, -1)
843
+ .type(topk_weight.dtype)
844
+ .mul_(topk_weight.unsqueeze(dim=-1))
845
+ .sum(dim=1)
846
+ .type(new_x.dtype)
847
+ )
848
+ return final_out
849
+
850
+
851
+ # Copied from transformers.models.llama.modeling_llama.repeat_kv
852
+ def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
853
+ """
854
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
855
+ num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
856
+ """
857
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
858
+ if n_rep == 1:
859
+ return hidden_states
860
+ hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
861
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
862
+
863
+
864
+ # Copied from transformers.models.llama.modeling_llama.LlamaAttention with Llama->BailingMoeV2
865
+ class BailingMoeV2Attention(nn.Module):
866
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
867
+
868
+ def __init__(self, config: BailingMoeV2Config, layer_idx: Optional[int] = None):
869
+ super().__init__()
870
+ self.config = config
871
+ self.layer_idx = layer_idx
872
+ if layer_idx is None:
873
+ logger.warning_once(
874
+ f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will "
875
+ "to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` "
876
+ "when creating this class."
877
+ )
878
+
879
+ self.attention_dropout = config.attention_dropout
880
+ self.hidden_size = config.hidden_size
881
+ self.num_heads = config.num_attention_heads
882
+ self.head_dim = config.head_dim or self.hidden_size // self.num_heads
883
+ partial_rotary_factor = config.partial_rotary_factor if hasattr(config, "partial_rotary_factor") else 1.0
884
+ self.rope_dim = int(self.head_dim * partial_rotary_factor)
885
+ self.num_key_value_heads = config.num_key_value_heads
886
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
887
+ self.max_position_embeddings = config.max_position_embeddings
888
+ self.rope_theta = config.rope_theta
889
+ self.is_causal = True
890
+
891
+ self.query_key_value = nn.Linear(
892
+ self.hidden_size,
893
+ (self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,
894
+ bias=config.use_qkv_bias,
895
+ )
896
+
897
+ self.q_norm = BailingMoeV2RMSNorm(self.head_dim, eps=config.rms_norm_eps)
898
+ self.k_norm = BailingMoeV2RMSNorm(self.head_dim, eps=config.rms_norm_eps)
899
+ self.dense = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.use_bias)
900
+
901
+ if self.config.rope_scaling is not None:
902
+ self.rope_scaling = {"type": "mrope", "mrope_section": [8, 12, 12]}
903
+ self.mrope_section = self.rope_scaling["mrope_section"]
904
+
905
+ def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
906
+ return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
907
+
908
+ def forward(
909
+ self,
910
+ hidden_states: torch.Tensor,
911
+ attention_mask: Optional[torch.Tensor] = None,
912
+ position_ids: Optional[torch.LongTensor] = None,
913
+ past_key_value: Optional[Cache] = None,
914
+ output_attentions: bool = False,
915
+ use_cache: bool = False,
916
+ position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
917
+ **kwargs,
918
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
919
+ if "padding_mask" in kwargs:
920
+ warnings.warn(
921
+ "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"
922
+ )
923
+
924
+ bsz, q_len, _ = hidden_states.size()
925
+
926
+ qkv = self.query_key_value(hidden_states)
927
+ qkv = qkv.view(bsz, q_len, self.num_heads + 2 * self.num_key_value_heads, self.head_dim)
928
+
929
+ query_states, key_states, value_states = qkv.split(
930
+ [self.num_heads, self.num_key_value_heads, self.num_key_value_heads], dim=-2
931
+ )
932
+ query_states = query_states.transpose(1, 2)
933
+ key_states = key_states.transpose(1, 2)
934
+ value_states = value_states.transpose(1, 2)
935
+
936
+ if not self.training:
937
+ query_states = query_states.contiguous()
938
+ key_states = key_states.contiguous()
939
+ value_states = value_states.contiguous()
940
+ del qkv
941
+
942
+ query_states = self.q_norm(query_states)
943
+ key_states = self.k_norm(key_states)
944
+
945
+ kv_seq_len = key_states.shape[-2]
946
+ if past_key_value is not None:
947
+ if self.layer_idx is None:
948
+ raise ValueError(
949
+ f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "
950
+ "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "
951
+ "with a layer index."
952
+ )
953
+ kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)
954
+ cos, sin = position_embeddings
955
+ query_states, key_states = apply_3d_rotary_pos_emb(
956
+ query_states,
957
+ key_states,
958
+ cos,
959
+ sin,
960
+ mrope_section=self.rope_scaling["mrope_section"],
961
+ rope_type=self.config.rope_scaling["type"]
962
+ )
963
+
964
+ if past_key_value is not None:
965
+ cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models
966
+ key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
967
+
968
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
969
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
970
+
971
+ attn_weights = torch.matmul(query_states / math.sqrt(self.head_dim), key_states.transpose(2, 3))
972
+
973
+ if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):
974
+ raise ValueError(
975
+ f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"
976
+ f" {attn_weights.size()}"
977
+ )
978
+
979
+ if attention_mask is not None:
980
+ if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):
981
+ raise ValueError(
982
+ f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"
983
+ )
984
+ attn_weights = attn_weights + attention_mask
985
+
986
+ # upcast attention to fp32
987
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
988
+ attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training)
989
+ attn_output = torch.matmul(attn_weights, value_states)
990
+
991
+ if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):
992
+ raise ValueError(
993
+ f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"
994
+ f" {attn_output.size()}"
995
+ )
996
+
997
+ attn_output = attn_output.transpose(1, 2).contiguous()
998
+
999
+ attn_output = attn_output.reshape(bsz, q_len, -1)
1000
+
1001
+ attn_output = self.dense(attn_output)
1002
+
1003
+ if not output_attentions:
1004
+ attn_weights = None
1005
+
1006
+ if not self.training:
1007
+ del query_states, key_states, value_states
1008
+
1009
+ return attn_output, attn_weights, past_key_value
1010
+
1011
+
1012
+ # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2 with Llama->BailingMoeV2
1013
+ class BailingMoeV2FlashAttention2(BailingMoeV2Attention):
1014
+ """
1015
+ BailingMoeV2 flash attention module. This module inherits from `BailingMoeV2Attention` as the weights of the module stays
1016
+ untouched. The only required change would be on the forward pass where it needs to correctly call the public API of
1017
+ flash attention and deal with padding tokens in case the input contains any of them.
1018
+ """
1019
+
1020
+ def __init__(self, *args, **kwargs):
1021
+ super().__init__(*args, **kwargs)
1022
+
1023
+ # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.
1024
+ # flash_attn<2.1 generates top-left aligned causal mask, while what is needed here is bottom-right alignement, that was made default for flash_attn>=2.1. This attribute is used to handle this difference. Reference: https://github.com/Dao-AILab/flash-attention/releases/tag/v2.1.0.
1025
+ # Beware that with flash_attn<2.1, using q_seqlen != k_seqlen (except for the case q_seqlen == 1) produces a wrong mask (top-left).
1026
+ self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10()
1027
+
1028
+ def forward(
1029
+ self,
1030
+ hidden_states: torch.Tensor,
1031
+ attention_mask: Optional[torch.LongTensor] = None,
1032
+ position_ids: Optional[torch.LongTensor] = None,
1033
+ past_key_value: Optional[Cache] = None,
1034
+ output_attentions: bool = False,
1035
+ use_cache: bool = False,
1036
+ position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
1037
+ **kwargs,
1038
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
1039
+ # BailingMoeV2FlashAttention2 attention does not support output_attentions
1040
+ if "padding_mask" in kwargs:
1041
+ warnings.warn(
1042
+ "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"
1043
+ )
1044
+
1045
+ # overwrite attention_mask with padding_mask
1046
+ attention_mask = kwargs.pop("padding_mask")
1047
+
1048
+ output_attentions = False
1049
+
1050
+ bsz, q_len, _ = hidden_states.size()
1051
+
1052
+ # Flash attention requires the input to have the shape
1053
+ # batch_size x seq_length x head_dim x hidden_dim
1054
+ # therefore we just need to keep the original shape
1055
+
1056
+ qkv = self.query_key_value(hidden_states)
1057
+ qkv = qkv.view(bsz, q_len, self.num_heads + 2 * self.num_key_value_heads, self.head_dim)
1058
+
1059
+ query_states, key_states, value_states = qkv.split(
1060
+ [self.num_heads, self.num_key_value_heads, self.num_key_value_heads], dim=-2
1061
+ )
1062
+ query_states = query_states.transpose(1, 2)
1063
+ key_states = key_states.transpose(1, 2)
1064
+ value_states = value_states.transpose(1, 2)
1065
+
1066
+ if not self.training:
1067
+ query_states = query_states.contiguous()
1068
+ key_states = key_states.contiguous()
1069
+ value_states = value_states.contiguous()
1070
+ del qkv
1071
+
1072
+ query_states = self.q_norm(query_states)
1073
+ key_states = self.k_norm(key_states)
1074
+
1075
+ kv_seq_len = key_states.shape[-2]
1076
+ if past_key_value is not None:
1077
+ kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)
1078
+ cos, sin = position_embeddings
1079
+ if self.config.rope_scaling is not None:
1080
+ rope_type = self.config.rope_scaling.get("rope_type", self.config.rope_scaling.get("type"))
1081
+ else:
1082
+ rope_type = "default"
1083
+ query_states, key_states = apply_3d_rotary_pos_emb(
1084
+ query_states,
1085
+ key_states,
1086
+ cos,
1087
+ sin,
1088
+ rope_type=rope_type
1089
+ )
1090
+
1091
+ if past_key_value is not None:
1092
+ cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models
1093
+ key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
1094
+
1095
+ # TODO: These transpose are quite inefficient but Flash Attention requires the layout [batch_size, sequence_length, num_heads, head_dim]. We would need to refactor the KV cache
1096
+ # to be able to avoid many of these transpose/reshape/view.
1097
+ query_states = query_states.transpose(1, 2)
1098
+ key_states = key_states.transpose(1, 2)
1099
+ value_states = value_states.transpose(1, 2)
1100
+
1101
+ dropout_rate = self.attention_dropout if self.training else 0.0
1102
+
1103
+ # In PEFT, usually we cast the layer norms in float32 for training stability reasons
1104
+ # therefore the input hidden states gets silently cast in float32. Hence, we need
1105
+ # cast them back in the correct dtype just to be sure everything works as expected.
1106
+ # This might slow down training & inference so it is recommended to not cast the LayerNorms
1107
+ # in fp32. (BailingMoeV2RMSNorm handles it correctly)
1108
+
1109
+ input_dtype = query_states.dtype
1110
+ if input_dtype == torch.float32:
1111
+ # Handle the case where the model is quantized
1112
+ if hasattr(self.config, "_pre_quantization_dtype"):
1113
+ target_dtype = self.config._pre_quantization_dtype
1114
+ elif torch.is_autocast_enabled():
1115
+ target_dtype = torch.get_autocast_gpu_dtype()
1116
+ else:
1117
+ target_dtype = self.query_key_value.weight.dtype
1118
+
1119
+ logger.warning_once(
1120
+ f"The input hidden states seems to be silently casted in float32, this might be related to"
1121
+ f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"
1122
+ f" {target_dtype}."
1123
+ )
1124
+
1125
+ query_states = query_states.to(target_dtype)
1126
+ key_states = key_states.to(target_dtype)
1127
+ value_states = value_states.to(target_dtype)
1128
+
1129
+ attn_output = self._flash_attention_forward(
1130
+ query_states, key_states, value_states, attention_mask, q_len, dropout=dropout_rate
1131
+ )
1132
+
1133
+ attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()
1134
+ attn_output = self.dense(attn_output)
1135
+
1136
+ if not output_attentions:
1137
+ attn_weights = None
1138
+
1139
+ if not self.training:
1140
+ del query_states, key_states, value_states
1141
+
1142
+ return attn_output, attn_weights, past_key_value
1143
+
1144
+ def _flash_attention_forward(
1145
+ self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None
1146
+ ):
1147
+ """
1148
+ Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token
1149
+ first unpad the input, then computes the attention scores and pad the final attention scores.
1150
+
1151
+ Args:
1152
+ query_states (`torch.Tensor`):
1153
+ Input query states to be passed to Flash Attention API
1154
+ key_states (`torch.Tensor`):
1155
+ Input key states to be passed to Flash Attention API
1156
+ value_states (`torch.Tensor`):
1157
+ Input value states to be passed to Flash Attention API
1158
+ attention_mask (`torch.Tensor`):
1159
+ The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the
1160
+ position of padding tokens and 1 for the position of non-padding tokens.
1161
+ dropout (`int`, *optional*):
1162
+ Attention dropout
1163
+ softmax_scale (`float`, *optional*):
1164
+ The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)
1165
+ query_length (`int`):
1166
+ The length of the query sequence in terms of tokens. This represents the number of tokens in the
1167
+ `query_states` tensor along the sequence dimension. It is used to determine the effective sequence
1168
+ length for attention computations.
1169
+ """
1170
+ if not self._flash_attn_uses_top_left_mask:
1171
+ causal = self.is_causal
1172
+ else:
1173
+ # TODO: Remove the `query_length != 1` check once Flash Attention for RoCm is bumped to 2.1. For details, please see the comment in BailingMoeV2FlashAttention2 __init__.
1174
+ causal = self.is_causal and query_length != 1
1175
+
1176
+ # Contains at least one padding token in the sequence
1177
+ if attention_mask is not None:
1178
+ batch_size = query_states.shape[0]
1179
+ query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(
1180
+ query_states, key_states, value_states, attention_mask, query_length
1181
+ )
1182
+
1183
+ cu_seqlens_q, cu_seqlens_k = cu_seq_lens
1184
+ max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
1185
+
1186
+ attn_output_unpad = flash_attn_varlen_func(
1187
+ query_states,
1188
+ key_states,
1189
+ value_states,
1190
+ cu_seqlens_q=cu_seqlens_q,
1191
+ cu_seqlens_k=cu_seqlens_k,
1192
+ max_seqlen_q=max_seqlen_in_batch_q,
1193
+ max_seqlen_k=max_seqlen_in_batch_k,
1194
+ dropout_p=dropout,
1195
+ softmax_scale=softmax_scale,
1196
+ causal=causal,
1197
+ )
1198
+
1199
+ attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)
1200
+ else:
1201
+ attn_output = flash_attn_func(
1202
+ query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal
1203
+ )
1204
+
1205
+ return attn_output
1206
+
1207
+ def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):
1208
+ indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)
1209
+ batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape
1210
+
1211
+ key_layer = index_first_axis(
1212
+ key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
1213
+ )
1214
+ value_layer = index_first_axis(
1215
+ value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
1216
+ )
1217
+ if query_length == kv_seq_len:
1218
+ query_layer = index_first_axis(
1219
+ query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k
1220
+ )
1221
+ cu_seqlens_q = cu_seqlens_k
1222
+ max_seqlen_in_batch_q = max_seqlen_in_batch_k
1223
+ indices_q = indices_k
1224
+ elif query_length == 1:
1225
+ max_seqlen_in_batch_q = 1
1226
+ cu_seqlens_q = torch.arange(
1227
+ batch_size + 1, dtype=torch.int32, device=query_layer.device
1228
+ ) # There is a memcpy here, that is very bad.
1229
+ indices_q = cu_seqlens_q[:-1]
1230
+ query_layer = query_layer.squeeze(1)
1231
+ else:
1232
+ # The -q_len: slice assumes left padding.
1233
+ attention_mask = attention_mask[:, -query_length:]
1234
+ query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)
1235
+
1236
+ return (
1237
+ query_layer,
1238
+ key_layer,
1239
+ value_layer,
1240
+ indices_q,
1241
+ (cu_seqlens_q, cu_seqlens_k),
1242
+ (max_seqlen_in_batch_q, max_seqlen_in_batch_k),
1243
+ )
1244
+
1245
+
1246
+ ATTENTION_CLASSES = {
1247
+ "eager": BailingMoeV2Attention,
1248
+ "flash_attention_2": BailingMoeV2FlashAttention2,
1249
+ }
1250
+
1251
+
1252
+ class BailingMoeV2DecoderLayer(nn.Module):
1253
+ def __init__(self, config: BailingMoeV2Config, layer_idx: int):
1254
+ super().__init__()
1255
+ self.hidden_size = config.hidden_size
1256
+
1257
+ self.attention = ATTENTION_CLASSES[config._attn_implementation](config=config, layer_idx=layer_idx)
1258
+
1259
+ self.mlp = (
1260
+ BailingMoeV2SparseMoeBlock(config)
1261
+ if (config.num_experts is not None and layer_idx >= config.first_k_dense_replace)
1262
+ else BailingMoeV2MLP(config=config, intermediate_size=config.intermediate_size)
1263
+ )
1264
+ self.input_layernorm = BailingMoeV2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
1265
+ self.post_attention_layernorm = BailingMoeV2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
1266
+
1267
+ def forward(
1268
+ self,
1269
+ hidden_states: torch.Tensor,
1270
+ attention_mask: Optional[torch.Tensor] = None,
1271
+ position_ids: Optional[torch.LongTensor] = None,
1272
+ image_mask: Optional[torch.Tensor] = None,
1273
+ audio_mask: Optional[torch.Tensor] = None,
1274
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
1275
+ output_attentions: Optional[bool] = False,
1276
+ output_router_logits: Optional[bool] = False,
1277
+ use_cache: Optional[bool] = False,
1278
+ position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
1279
+ **kwargs,
1280
+ ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
1281
+ """
1282
+ Args:
1283
+ hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
1284
+ attention_mask (`torch.FloatTensor`, *optional*):
1285
+ attention mask of size `(batch_size, sequence_length)` if flash attention is used or `(batch_size, 1,
1286
+ query_sequence_length, key_sequence_length)` if default attention is used.
1287
+ position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1288
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
1289
+ config.n_positions - 1]`.
1290
+ past_key_value (`Tuple(torch.FloatTensor)`, *optional*):
1291
+ cached past key and value projection states
1292
+ output_attentions (`bool`, *optional*):
1293
+ Whether to return the attentions tensors of all attention layers. See `attentions` under
1294
+ returned tensors for more detail.
1295
+ output_router_logits (`bool`, *optional*):
1296
+ Whether or not to return the logits of all the routers. They are useful for computing the router loss,
1297
+ and should not be returned during inference.
1298
+ use_cache (`bool`, *optional*):
1299
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding
1300
+ (see `past_key_values`).
1301
+ """
1302
+ if "padding_mask" in kwargs:
1303
+ warnings.warn(
1304
+ "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"
1305
+ )
1306
+ residual = hidden_states
1307
+
1308
+ hidden_states = self.input_layernorm(hidden_states)
1309
+
1310
+ # Self Attention
1311
+ hidden_states, self_attn_weights, present_key_value = self.attention(
1312
+ hidden_states=hidden_states,
1313
+ attention_mask=attention_mask,
1314
+ position_ids=position_ids,
1315
+ past_key_value=past_key_value,
1316
+ output_attentions=output_attentions,
1317
+ position_embeddings=position_embeddings,
1318
+ use_cache=use_cache,
1319
+ )
1320
+ hidden_states = residual + hidden_states
1321
+
1322
+ # Fully Connected
1323
+ residual = hidden_states
1324
+ hidden_states = self.post_attention_layernorm(hidden_states)
1325
+ hidden_states = self.mlp(hidden_states, image_mask, audio_mask)
1326
+ #hidden_states = self.mlp(hidden_states)
1327
+ if isinstance(hidden_states, tuple):
1328
+ hidden_states, router_logits = hidden_states
1329
+ else:
1330
+ router_logits = None
1331
+ hidden_states = residual + hidden_states.to(residual.device)
1332
+
1333
+ outputs = (hidden_states,)
1334
+
1335
+ if output_attentions:
1336
+ outputs += (self_attn_weights,)
1337
+
1338
+ if use_cache:
1339
+ outputs += (present_key_value,)
1340
+
1341
+ if output_router_logits:
1342
+ outputs += (router_logits,)
1343
+
1344
+ return outputs
1345
+
1346
+
1347
+ BAILINGMOEV2_START_DOCSTRING = r"""
1348
+ This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
1349
+ library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
1350
+ etc.)
1351
+
1352
+ This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
1353
+ Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
1354
+ and behavior.
1355
+
1356
+ Parameters:
1357
+ config ([`BailingMoeV2Config`]):
1358
+ Model configuration class with all the parameters of the model. Initializing with a config file does not
1359
+ load the weights associated with the model, only the configuration. Check out the
1360
+ [`~PreTrainedModel.from_pretrained`] method to load the model weights.
1361
+ """
1362
+
1363
+
1364
+ @add_start_docstrings(
1365
+ "The bare BailingMoeV2 Model outputting raw hidden-states without any specific head on top.",
1366
+ BAILINGMOEV2_START_DOCSTRING,
1367
+ )
1368
+ class BailingMoeV2PreTrainedModel(PreTrainedModel):
1369
+ config_class = BailingMoeV2Config
1370
+ base_model_prefix = "model"
1371
+ supports_gradient_checkpointing = True
1372
+ _no_split_modules = ["BailingMoeV2DecoderLayer"]
1373
+ _skip_keys_device_placement = "past_key_values"
1374
+ _supports_flash_attn_2 = True
1375
+ _supports_sdpa = True
1376
+ _supports_cache_class = True
1377
+
1378
+ def _init_weights(self, module):
1379
+ std = self.config.initializer_range
1380
+ if isinstance(module, nn.Linear):
1381
+ module.weight.data.normal_(mean=0.0, std=std)
1382
+ if module.bias is not None:
1383
+ module.bias.data.zero_()
1384
+ elif isinstance(module, nn.Embedding):
1385
+ module.weight.data.normal_(mean=0.0, std=std)
1386
+ if module.padding_idx is not None:
1387
+ module.weight.data[module.padding_idx].zero_()
1388
+
1389
+
1390
+ BAILINGMOEV2_INPUTS_DOCSTRING = r"""
1391
+ Args:
1392
+ input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
1393
+ Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
1394
+ it.
1395
+
1396
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
1397
+ [`PreTrainedTokenizer.__call__`] for details.
1398
+
1399
+ [What are input IDs?](../glossary#input-ids)
1400
+ attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
1401
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
1402
+
1403
+ - 1 for tokens that are **not masked**,
1404
+ - 0 for tokens that are **masked**.
1405
+
1406
+ [What are attention masks?](../glossary#attention-mask)
1407
+
1408
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
1409
+ [`PreTrainedTokenizer.__call__`] for details.
1410
+
1411
+ If `past_key_values` is used, optionally only the last `input_ids` have to be input (see
1412
+ `past_key_values`).
1413
+
1414
+ If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
1415
+ and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more
1416
+ information on the default strategy.
1417
+
1418
+ - 1 indicates the head is **not masked**,
1419
+ - 0 indicates the head is **masked**.
1420
+ position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1421
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
1422
+ config.n_positions - 1]`.
1423
+
1424
+ [What are position IDs?](../glossary#position-ids)
1425
+ past_key_values (`Cache` or `tuple(tuple(torch.FloatTensor))`, *optional*):
1426
+ Pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
1427
+ blocks) that can be used to speed up sequential decoding. This typically consists in the `past_key_values`
1428
+ returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.
1429
+
1430
+ Two formats are allowed:
1431
+ - a [`~cache_utils.Cache`] instance;
1432
+ - Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of
1433
+ shape `(batch_size, num_heads, sequence_length, embed_size_per_head)`). This is also known as the legacy
1434
+ cache format.
1435
+
1436
+ The model will output the same cache format that is fed as input. If no `past_key_values` are passed, the
1437
+ legacy cache format will be returned.
1438
+
1439
+ If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't
1440
+ have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`
1441
+ of shape `(batch_size, sequence_length)`.
1442
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
1443
+ Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
1444
+ is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
1445
+ model's internal embedding lookup matrix.
1446
+ use_cache (`bool`, *optional*):
1447
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
1448
+ `past_key_values`).
1449
+ output_attentions (`bool`, *optional*):
1450
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
1451
+ tensors for more detail.
1452
+ output_hidden_states (`bool`, *optional*):
1453
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
1454
+ more detail.
1455
+ return_dict (`bool`, *optional*):
1456
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
1457
+ """
1458
+
1459
+
1460
+ @add_start_docstrings(
1461
+ "The bare BailingMoeV2 Model outputting raw hidden-states without any specific head on top.",
1462
+ BAILINGMOEV2_START_DOCSTRING,
1463
+ )
1464
+ class BailingMoeV2Model(BailingMoeV2PreTrainedModel):
1465
+ """
1466
+ Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`BailingMoeV2DecoderLayer`]
1467
+
1468
+ Args:
1469
+ config: BailingMoeV2Config
1470
+ """
1471
+
1472
+ def __init__(self, config: BailingMoeV2Config):
1473
+ super().__init__(config)
1474
+ self.padding_idx = config.pad_token_id
1475
+ self.vocab_size = config.vocab_size
1476
+
1477
+ self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
1478
+ self.layers = nn.ModuleList(
1479
+ [BailingMoeV2DecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
1480
+ )
1481
+ self._use_sdpa = config._attn_implementation == "sdpa"
1482
+ self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
1483
+ self.norm = BailingMoeV2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
1484
+ if self.config.rope_scaling is not None:
1485
+ self.rotary_emb = BailingMoeV2RotaryEmbedding3D(config=config)
1486
+ else:
1487
+ self.rotary_emb = BailingMoeV2RotaryEmbedding(config=config)
1488
+ self.gradient_checkpointing = False
1489
+ config.spatial_merge_size = 2
1490
+ config.tokens_per_second = 2
1491
+ self.rope_deltas = None
1492
+ # Initialize weights and apply final processing
1493
+ self.post_init()
1494
+
1495
+ def get_input_embeddings(self):
1496
+ return self.word_embeddings
1497
+
1498
+ def set_input_embeddings(self, value):
1499
+ self.word_embeddings = value
1500
+
1501
+ def prompt_wrap_vision(self, input_ids, inputs_embeds, vision_embeds, vision_token_id):
1502
+ if vision_embeds is None or input_ids is None:
1503
+ return inputs_embeds
1504
+
1505
+ if len(vision_embeds.shape) == 3:
1506
+ vision_embeds = vision_embeds.reshape(-1, vision_embeds.shape[-1])
1507
+
1508
+ n_image_tokens = (input_ids == vision_token_id).sum().item()
1509
+ n_image_features = vision_embeds.shape[0]
1510
+
1511
+ if n_image_tokens != n_image_features:
1512
+ raise ValueError(
1513
+ f"Image features and image tokens do not match: tokens: {n_image_tokens}, features {n_image_features}"
1514
+ )
1515
+ vision_mask = (input_ids == vision_token_id).unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device)
1516
+
1517
+ image_embeds = vision_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
1518
+ inputs_embeds = inputs_embeds.masked_scatter(vision_mask, image_embeds)
1519
+
1520
+ return inputs_embeds
1521
+
1522
+ def prompt_wrap_audio(self, input_ids, inputs_embeds, audio_embeds, audio_embeds_lengths, placeholder_audio_loc_lens):
1523
+ inputs_embeds = patch_continuous_features(
1524
+ input_embeddings=inputs_embeds, placeholder_loc_lens=placeholder_audio_loc_lens,
1525
+ encoded_feats=audio_embeds, encoded_feat_lens=audio_embeds_lengths,
1526
+ )
1527
+ router_mask_audio = build_modality_mask(placeholder_audio_loc_lens, inputs_embeds.shape[:-1])
1528
+ router_mask_audio = router_mask_audio.to(inputs_embeds.device)
1529
+ return inputs_embeds, router_mask_audio
1530
+
1531
+ def prompt_wrap_navit(self, input_ids, config, query_embeds_image=None, query_embeds_video=None, query_embeds_audio=None,
1532
+ query_embeds_audio_lengths=None, placeholder_audio_loc_lens=None, target_embeds=None):
1533
+ inputs_embeds = self.word_embeddings(input_ids)
1534
+ vision_mask = None
1535
+ audio_mask = None
1536
+ if query_embeds_image is None and query_embeds_video is None and query_embeds_audio is None and target_embeds is None:
1537
+ return inputs_embeds, vision_mask, audio_mask
1538
+
1539
+ if query_embeds_image is not None:
1540
+ inputs_embeds = self.prompt_wrap_vision(input_ids, inputs_embeds, query_embeds_image, config.image_patch_token)
1541
+ if query_embeds_video is not None:
1542
+ inputs_embeds = self.prompt_wrap_vision(input_ids, inputs_embeds, query_embeds_video, config.video_patch_token)
1543
+
1544
+ image_mask = input_ids == config.image_patch_token
1545
+ video_mask = input_ids == config.video_patch_token
1546
+ vision_mask = (image_mask + video_mask) > 0
1547
+ vision_mask = vision_mask.unsqueeze(-1).to(input_ids.device)
1548
+
1549
+ if query_embeds_audio is not None:
1550
+ inputs_embeds, audio_mask = self.prompt_wrap_audio(
1551
+ input_ids, inputs_embeds, query_embeds_audio, query_embeds_audio_lengths, placeholder_audio_loc_lens,
1552
+ )
1553
+
1554
+ return inputs_embeds, vision_mask, audio_mask
1555
+
1556
+
1557
+ @add_start_docstrings_to_model_forward(BAILINGMOEV2_INPUTS_DOCSTRING)
1558
+ def forward(
1559
+ self,
1560
+ input_ids: torch.LongTensor = None,
1561
+ attention_mask: Optional[torch.Tensor] = None,
1562
+ position_ids: Optional[torch.LongTensor] = None,
1563
+ query_embeds_image: Optional[torch.Tensor] = None,
1564
+ query_embeds_video: Optional[torch.Tensor] = None,
1565
+ query_embeds_audio: Optional[torch.Tensor] = None,
1566
+ query_embeds_audio_lengths: Optional[torch.Tensor] = None,
1567
+ placeholder_audio_loc_lens: Optional[torch.Tensor] = None,
1568
+ target_embeds: Optional[torch.Tensor] = None,
1569
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1570
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1571
+ image_grid_thw: Optional[torch.Tensor] = None,
1572
+ image_grid_thw_video: Optional[torch.Tensor] = None,
1573
+ use_cache: Optional[bool] = None,
1574
+ output_attentions: Optional[bool] = None,
1575
+ output_hidden_states: Optional[bool] = None,
1576
+ output_router_logits: Optional[bool] = None,
1577
+ return_dict: Optional[bool] = None,
1578
+ second_per_grid_ts: Optional[torch.Tensor] = None,
1579
+ image_mask=None,
1580
+ audio_mask=None,
1581
+ **kwargs,
1582
+ ) -> Union[Tuple, MoeModelOutputWithPast]:
1583
+
1584
+ if inputs_embeds is not None:
1585
+ words_embeddings = inputs_embeds
1586
+ input_shape = inputs_embeds.size()[:2]
1587
+ else:
1588
+ if (query_embeds_image is None and query_embeds_video is None and query_embeds_audio is None and target_embeds is None) or input_ids.size(1) == 1:
1589
+ words_embeddings = self.word_embeddings(input_ids.clip(0, self.word_embeddings.weight.shape[0] - 1))
1590
+ input_shape = input_ids.size()
1591
+ image_mask = None
1592
+ audio_mask = None
1593
+ else:
1594
+ words_embeddings, image_mask, audio_mask = self.prompt_wrap_navit(
1595
+ input_ids.clip(0, self.word_embeddings.weight.shape[0] - 1), self.config, query_embeds_image, query_embeds_video, query_embeds_audio,
1596
+ query_embeds_audio_lengths, placeholder_audio_loc_lens, target_embeds, # noqa
1597
+ )
1598
+
1599
+ input_shape = words_embeddings.size()[:2]
1600
+ embeddings = words_embeddings
1601
+
1602
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1603
+ output_hidden_states = (
1604
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1605
+ )
1606
+ output_router_logits = (
1607
+ output_router_logits if output_router_logits is not None else self.config.output_router_logits
1608
+ )
1609
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
1610
+
1611
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1612
+
1613
+ # retrieve input_ids and inputs_embeds
1614
+ if input_ids is not None and inputs_embeds is not None:
1615
+ assert input_ids.size(1) == inputs_embeds.size(1), "{} vs {}".format(
1616
+ input_ids.size,
1617
+ inputs_embeds.size,
1618
+ )
1619
+ batch_size, seq_length = inputs_embeds.shape[:2]
1620
+ #raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
1621
+ elif input_ids is not None:
1622
+ batch_size, seq_length = input_ids.shape[:2]
1623
+ elif inputs_embeds is not None:
1624
+ batch_size, seq_length = inputs_embeds.shape[:2]
1625
+ else:
1626
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
1627
+
1628
+ if self.gradient_checkpointing and self.training:
1629
+ if use_cache:
1630
+ logger.warning_once(
1631
+ "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`transformers."
1632
+ )
1633
+ use_cache = False
1634
+
1635
+ past_key_values_length = 0
1636
+ if use_cache:
1637
+ use_legacy_cache = not isinstance(past_key_values, Cache)
1638
+ if use_legacy_cache:
1639
+ past_key_values = DynamicCache.from_legacy_cache(past_key_values)
1640
+ past_key_values_length = past_key_values.get_usable_length(seq_length)
1641
+
1642
+ if position_ids is None:
1643
+ device = input_ids.device if input_ids is not None else inputs_embeds.device
1644
+ position_ids = torch.arange(
1645
+ past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device
1646
+ )
1647
+ position_ids = position_ids.unsqueeze(0)
1648
+
1649
+ if self.rotary_emb.rope_type == "video_rope" and input_ids.size(1) != 1:
1650
+ position_ids, rope_deltas = get_t_scale_rope_index(
1651
+ self.config,
1652
+ input_ids,
1653
+ image_grid_thw,
1654
+ image_grid_thw_video,
1655
+ attention_mask,
1656
+ scale_factor=2.0,
1657
+ second_per_grid_ts=second_per_grid_ts,
1658
+ )
1659
+ self.rope_deltas = rope_deltas
1660
+ elif self.rotary_emb.rope_type == "3D" and input_ids.size(1) != 1:
1661
+ position_ids, rope_deltas = get_rope_index(
1662
+ self.config,
1663
+ input_ids,
1664
+ image_grid_thw,
1665
+ image_grid_thw_video,
1666
+ attention_mask,
1667
+ second_per_grid_ts,
1668
+ )
1669
+ self.rope_deltas = rope_deltas
1670
+ elif self.rotary_emb.rope_type == "3D" or self.rotary_emb.rope_type == "video_rope": # decode stage
1671
+ batch_size, seq_length = input_ids.shape
1672
+ if past_key_values and self.rope_deltas:
1673
+ delta = past_key_values[0][1].shape[2] + self.rope_deltas
1674
+ elif past_key_values:
1675
+ delta = torch.tensor(past_key_values[0][1].shape[2])
1676
+ else:
1677
+ delta = torch.tensor(0)
1678
+ delta = delta.to(input_ids.device)
1679
+ position_ids = torch.arange(seq_length, device=input_ids.device)
1680
+ position_ids = position_ids.view(1, -1).expand(batch_size, -1)
1681
+ position_ids = position_ids.add(delta)
1682
+ position_ids = position_ids.unsqueeze(0).expand(3, -1, -1)
1683
+
1684
+ inputs_embeds = embeddings
1685
+ if self._use_flash_attention_2:
1686
+ # 2d mask is passed through the layers
1687
+ attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None
1688
+ elif self._use_sdpa and not output_attentions:
1689
+ # output_attentions=True can not be supported when using SDPA, and we fall back on
1690
+ # the manual implementation that requires a 4D causal mask in all cases.
1691
+ attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(
1692
+ attention_mask,
1693
+ (batch_size, seq_length),
1694
+ inputs_embeds,
1695
+ past_key_values_length,
1696
+ )
1697
+ else:
1698
+ # 4d mask is passed through the layers
1699
+ attention_mask = _prepare_4d_causal_attention_mask(
1700
+ attention_mask, (batch_size, seq_length), inputs_embeds, past_key_values_length
1701
+ )
1702
+
1703
+ # embed positions
1704
+ hidden_states = inputs_embeds
1705
+
1706
+ # create position embeddings to be shared across the decoder layers
1707
+ position_embeddings = self.rotary_emb(hidden_states, position_ids)
1708
+
1709
+ # decoder layers
1710
+ all_hidden_states = () if output_hidden_states else None
1711
+ all_self_attns = () if output_attentions else None
1712
+ all_router_logits = () if output_router_logits else None
1713
+ next_decoder_cache = None
1714
+
1715
+ for decoder_layer in self.layers:
1716
+ if output_hidden_states:
1717
+ all_hidden_states += (hidden_states,)
1718
+
1719
+ if self.gradient_checkpointing and self.training:
1720
+ layer_outputs = self._gradient_checkpointing_func(
1721
+ decoder_layer.__call__,
1722
+ hidden_states,
1723
+ attention_mask,
1724
+ position_ids,
1725
+ image_mask,
1726
+ audio_mask,
1727
+ past_key_values,
1728
+ output_attentions,
1729
+ output_router_logits,
1730
+ use_cache,
1731
+ position_embeddings,
1732
+ )
1733
+ else:
1734
+ layer_outputs = decoder_layer(
1735
+ hidden_states,
1736
+ attention_mask=attention_mask,
1737
+ position_ids=position_ids,
1738
+ image_mask=image_mask,
1739
+ audio_mask=audio_mask,
1740
+ past_key_value=past_key_values,
1741
+ output_attentions=output_attentions,
1742
+ output_router_logits=output_router_logits,
1743
+ use_cache=use_cache,
1744
+ position_embeddings=position_embeddings,
1745
+ )
1746
+ hidden_states = layer_outputs[0]
1747
+
1748
+ if use_cache:
1749
+ next_decoder_cache = layer_outputs[2 if output_attentions else 1]
1750
+
1751
+ if output_attentions:
1752
+ all_self_attns += (layer_outputs[1],)
1753
+
1754
+ if output_router_logits and layer_outputs[-1] is not None:
1755
+ all_router_logits += (layer_outputs[-1],)
1756
+
1757
+ hidden_states = self.norm(hidden_states)
1758
+
1759
+ # add hidden states from the last decoder layer
1760
+ if output_hidden_states:
1761
+ all_hidden_states += (hidden_states,)
1762
+
1763
+ next_cache = None
1764
+ if use_cache:
1765
+ next_cache = next_decoder_cache.to_legacy_cache() if use_legacy_cache else next_decoder_cache
1766
+ if not return_dict:
1767
+ return tuple(
1768
+ v
1769
+ for v in [hidden_states, next_cache, all_hidden_states, all_self_attns, all_router_logits]
1770
+ if v is not None
1771
+ )
1772
+ return MoeModelOutputWithPast(
1773
+ last_hidden_state=hidden_states,
1774
+ past_key_values=next_cache,
1775
+ hidden_states=all_hidden_states,
1776
+ attentions=all_self_attns,
1777
+ router_logits=all_router_logits,
1778
+ )
1779
+
1780
+
1781
+ class BailingMoeV2ForCausalLM(BailingMoeV2PreTrainedModel, GenerationMixin):
1782
+ _tied_weights_keys = ["lm_head.weight"]
1783
+
1784
+ def __init__(self, config: BailingMoeV2Config):
1785
+ super().__init__(config)
1786
+ self.model = BailingMoeV2Model(config)
1787
+ self.vocab_size = config.vocab_size
1788
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
1789
+
1790
+ # Initialize weights and apply final processing
1791
+ self.post_init()
1792
+
1793
+ def get_input_embeddings(self):
1794
+ return self.model.word_embeddings
1795
+
1796
+ def set_input_embeddings(self, value):
1797
+ self.model.word_embeddings = value
1798
+
1799
+ def get_output_embeddings(self):
1800
+ return self.lm_head
1801
+
1802
+ def set_output_embeddings(self, new_embeddings):
1803
+ self.lm_head = new_embeddings
1804
+
1805
+ def set_decoder(self, decoder):
1806
+ self.model = decoder
1807
+
1808
+ def get_decoder(self):
1809
+ return self.model
1810
+
1811
+ @add_start_docstrings_to_model_forward(BAILINGMOEV2_INPUTS_DOCSTRING)
1812
+ @replace_return_docstrings(output_type=MoeCausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
1813
+ def forward(
1814
+ self,
1815
+ input_ids: torch.LongTensor = None,
1816
+ attention_mask: Optional[torch.Tensor] = None,
1817
+ query_embeds_image: Optional[torch.Tensor] = None,
1818
+ query_embeds_video: Optional[torch.Tensor] = None,
1819
+ query_embeds_audio: Optional[torch.Tensor] = None,
1820
+ query_embeds_audio_lengths: Optional[torch.Tensor] = None,
1821
+ placeholder_audio_loc_lens: Optional[torch.Tensor] = None,
1822
+ position_ids: Optional[torch.LongTensor] = None,
1823
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1824
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1825
+ image_grid_thw: Optional[torch.Tensor] = None,
1826
+ image_grid_thw_video: Optional[torch.Tensor] = None,
1827
+ labels: Optional[torch.LongTensor] = None,
1828
+ use_cache: Optional[bool] = None,
1829
+ output_attentions: Optional[bool] = None,
1830
+ output_hidden_states: Optional[bool] = None,
1831
+ output_router_logits: Optional[bool] = None,
1832
+ return_dict: Optional[bool] = None,
1833
+ second_per_grid_ts: Optional[torch.Tensor] = None,
1834
+ num_logits_to_keep: Optional[int] = 0,
1835
+ image_mask=None,
1836
+ audio_mask=None,
1837
+ **kwargs,
1838
+ ) -> Union[Tuple, MoeCausalLMOutputWithPast]:
1839
+ r"""
1840
+ Args:
1841
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1842
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1843
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1844
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1845
+
1846
+ Returns:
1847
+
1848
+ Example:
1849
+
1850
+ ```python
1851
+ >>> from transformers import AutoTokenizer
1852
+
1853
+ >>> model = BailingMoeV2ForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)
1854
+ >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)
1855
+
1856
+ >>> prompt = "Hey, are you conscious? Can you talk to me?"
1857
+ >>> inputs = tokenizer(prompt, return_tensors="pt")
1858
+
1859
+ >>> # Generate
1860
+ >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
1861
+ >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
1862
+ "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."
1863
+ ```"""
1864
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1865
+ output_hidden_states = (
1866
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1867
+ )
1868
+ output_router_logits = (
1869
+ output_router_logits if output_router_logits is not None else self.config.output_router_logits
1870
+ )
1871
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1872
+ # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
1873
+ outputs = self.model(
1874
+ input_ids=input_ids,
1875
+ attention_mask=attention_mask,
1876
+ position_ids=position_ids,
1877
+ query_embeds_image=query_embeds_image,
1878
+ query_embeds_video=query_embeds_video,
1879
+ query_embeds_audio=query_embeds_audio,
1880
+ query_embeds_audio_lengths=query_embeds_audio_lengths,
1881
+ placeholder_audio_loc_lens=placeholder_audio_loc_lens,
1882
+ past_key_values=past_key_values,
1883
+ inputs_embeds=inputs_embeds,
1884
+ image_grid_thw=image_grid_thw,
1885
+ image_grid_thw_video=image_grid_thw_video,
1886
+ use_cache=use_cache,
1887
+ output_attentions=output_attentions,
1888
+ output_hidden_states=output_hidden_states,
1889
+ output_router_logits=output_router_logits,
1890
+ return_dict=return_dict,
1891
+ second_per_grid_ts=second_per_grid_ts,
1892
+ image_mask=image_mask,
1893
+ audio_mask=audio_mask,
1894
+ **kwargs,
1895
+ )
1896
+
1897
+ hidden_states = outputs[0]
1898
+
1899
+ logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :])
1900
+ # logits = logits.float()
1901
+
1902
+ loss = None
1903
+ aux_loss = None
1904
+
1905
+ if labels is not None:
1906
+ # Shift so that tokens < n predict n
1907
+ shift_logits = logits[..., :-1, :].contiguous()
1908
+ shift_labels = labels[..., 1:].contiguous()
1909
+ # Flatten the tokens
1910
+ loss_fct = CrossEntropyLoss(inplace_backward=True)
1911
+ shift_logits = shift_logits.view(-1, self.config.vocab_size)
1912
+ shift_labels = shift_labels.view(-1)
1913
+ # Enable model parallelism
1914
+ shift_labels = shift_labels.to(shift_logits.device)
1915
+ loss = loss_fct(shift_logits, shift_labels)
1916
+
1917
+ if not return_dict:
1918
+ output = (logits,) + outputs[1:]
1919
+ if output_router_logits:
1920
+ output = (aux_loss,) + output
1921
+ return (loss,) + output if loss is not None else output
1922
+
1923
+ return MoeCausalLMOutputWithPast(
1924
+ loss=loss,
1925
+ aux_loss=aux_loss,
1926
+ logits=logits,
1927
+ past_key_values=outputs.past_key_values,
1928
+ hidden_states=outputs.hidden_states,
1929
+ attentions=outputs.attentions,
1930
+ router_logits=outputs.router_logits,
1931
+ )
1932
+
1933
+ def prepare_inputs_for_generation(
1934
+ self,
1935
+ input_ids,
1936
+ query_embeds_image: torch.Tensor = None,
1937
+ query_embeds_video: torch.Tensor = None,
1938
+ query_embeds_audio: torch.Tensor = None,
1939
+ query_embeds_audio_lengths: torch.Tensor = None,
1940
+ placeholder_audio_loc_lens: torch.Tensor = None,
1941
+ image_mask=None,
1942
+ audio_mask=None,
1943
+ past_key_values=None,
1944
+ attention_mask=None,
1945
+ inputs_embeds=None,
1946
+ cache_position=None,
1947
+ position_ids=None,
1948
+ use_cache=True,
1949
+ image_grid_thw=None,
1950
+ image_grid_thw_video=None,
1951
+ second_per_grid_ts=None,
1952
+ is_audio_generation_mode=False,
1953
+ num_logits_to_keep=0,
1954
+ **kwargs,
1955
+ ):
1956
+ if past_key_values is not None:
1957
+ if isinstance(past_key_values, Cache):
1958
+ cache_length = past_key_values.get_seq_length()
1959
+ past_length = past_key_values.seen_tokens
1960
+ max_cache_length = (
1961
+ past_key_values.get_max_length()
1962
+ if hasattr(past_key_values, "get_max_length")
1963
+ else past_key_values.get_max_cache_shape()
1964
+ )
1965
+ else:
1966
+ cache_length = past_length = past_key_values[0][0].shape[2]
1967
+ max_cache_length = None
1968
+
1969
+ # Keep only the unprocessed tokens:
1970
+ # 1 - If the length of the attention_mask exceeds the length of input_ids, then we are in a setting where
1971
+ # some of the inputs are exclusivelly passed as part of the cache (e.g. when passing input_embeds as input)
1972
+ if attention_mask is not None and attention_mask.shape[1] > input_ids.shape[1]:
1973
+ input_ids = input_ids[:, -(attention_mask.shape[1] - past_length) :]
1974
+ # 2 - If the past_length is smaller than input_ids', then input_ids holds all input tokens. We can discard
1975
+ # input_ids based on the past_length.
1976
+ elif past_length < input_ids.shape[1]:
1977
+ input_ids = input_ids[:, past_length:]
1978
+ # 3 - Otherwise (past_length >= input_ids.shape[1]), let's assume input_ids only has unprocessed tokens.
1979
+
1980
+ # If we are about to go beyond the maximum cache length, we need to crop the input attention mask.
1981
+ if (
1982
+ max_cache_length is not None
1983
+ and attention_mask is not None
1984
+ and cache_length + input_ids.shape[1] > max_cache_length
1985
+ ):
1986
+ attention_mask = attention_mask[:, -max_cache_length:]
1987
+
1988
+ position_ids = kwargs.get("position_ids", None)
1989
+ if attention_mask is not None and position_ids is None:
1990
+ # create position_ids on the fly for batch generation
1991
+ position_ids = attention_mask.long().cumsum(-1) - 1
1992
+ position_ids.masked_fill_(attention_mask == 0, 1)
1993
+ if past_key_values:
1994
+ position_ids = position_ids[:, -input_ids.shape[1] :]
1995
+ position_ids = position_ids.unsqueeze(0)
1996
+
1997
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1998
+ if inputs_embeds is not None and past_key_values is None:
1999
+ model_inputs = {"inputs_embeds": inputs_embeds, "input_ids": None}
2000
+ else:
2001
+ model_inputs = {"input_ids": input_ids, "inputs_embeds": None}
2002
+
2003
+ model_inputs.update(
2004
+ {
2005
+ "query_embeds_image": query_embeds_image,
2006
+ "query_embeds_video": query_embeds_video,
2007
+ "query_embeds_audio": query_embeds_audio,
2008
+ "position_ids": position_ids,
2009
+ "past_key_values": past_key_values,
2010
+ "use_cache": use_cache,
2011
+ "attention_mask": attention_mask,
2012
+ "image_grid_thw": image_grid_thw,
2013
+ "image_grid_thw_video": image_grid_thw_video,
2014
+ "image_mask": image_mask,
2015
+ "audio_mask": audio_mask,
2016
+ "query_embeds_audio_lengths": query_embeds_audio_lengths,
2017
+ "placeholder_audio_loc_lens": placeholder_audio_loc_lens,
2018
+ "num_logits_to_keep": num_logits_to_keep,
2019
+ }
2020
+ )
2021
+ return model_inputs
2022
+
2023
+ @staticmethod
2024
+ def _reorder_cache(past_key_values, beam_idx):
2025
+ reordered_past = ()
2026
+ for layer_past in past_key_values:
2027
+ reordered_past += (
2028
+ tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past),
2029
+ )
2030
+ return reordered_past
2031
+
code/modeling_bailingmm2.py ADDED
@@ -0,0 +1,814 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # coding=utf-8
3
+ # Copyright (c) Ant Group. All rights reserved.
4
+
5
+ from typing import List, Optional, Union
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+ from PIL import Image
11
+
12
+ from diffusers.models.normalization import RMSNorm
13
+ from transformers import PreTrainedModel
14
+ from transformers.utils import logging
15
+ from configuration_bailingmm2 import BailingMM2Config
16
+ from modeling_bailing_moe_v2 import BailingMoeV2ForCausalLM
17
+ from bailingmm_utils import process_ratio, find_first_index_of_consecutive_ones, merge_consecutive_ones
18
+ from inference_profile import load_checkpoint_capabilities, resolve_model_directory
19
+ import os
20
+ from copy import deepcopy
21
+
22
+ # vision encoder
23
+ from qwen2_5_vit import Qwen2_5_VisionTransformer
24
+
25
+ logger = logging.get_logger(__name__)
26
+
27
+ _CONFIG_FOR_DOC = "BailingMM2Config"
28
+
29
+
30
+ class BailingMM2NativeForConditionalGeneration(PreTrainedModel):
31
+ config_class = BailingMM2Config
32
+ base_model_prefix = "model"
33
+ supports_gradient_checkpointing = True
34
+ _skip_keys_device_placement = "past_key_values"
35
+ _supports_flash_attn_2 = True
36
+
37
+ def __init__(
38
+ self,
39
+ config: BailingMM2Config,
40
+ empty_load=False,
41
+ ):
42
+ super().__init__(config)
43
+ self.config: BailingMM2Config = config
44
+ self.vision = None
45
+
46
+ self.llm_dytpe = torch.bfloat16
47
+
48
+ if empty_load:
49
+ self.model = None
50
+ return
51
+
52
+ if self.config.vision_config:
53
+ self.vision = Qwen2_5_VisionTransformer(self.config.vision_config)
54
+
55
+ self.model = BailingMoeV2ForCausalLM(self.config.llm_config)
56
+
57
+ mlp_modules_img = [nn.Linear(self.vision.image_emb_dim, self.model.config.hidden_size)]
58
+ for _ in range(1, self.config.mlp_depth):
59
+ mlp_modules_img.append(nn.GELU())
60
+ mlp_modules_img.append(nn.Linear(self.model.config.hidden_size, self.model.config.hidden_size))
61
+ self.linear_proj = nn.Sequential(*mlp_modules_img)
62
+
63
+ self.post_init()
64
+
65
+
66
+ def extract_image_feature(self, pixel_values, grid_thw):
67
+ with torch.cuda.amp.autocast(dtype=torch.bfloat16):
68
+ image_embeds = self.vision(pixel_values, grid_thw=grid_thw)
69
+ image_embeds = self.linear_proj(image_embeds)
70
+ image_embeds = F.normalize(image_embeds, dim=-1)
71
+ return image_embeds
72
+
73
+
74
+ @torch.no_grad()
75
+ def generate(
76
+ self,
77
+ input_ids: torch.LongTensor = None,
78
+ attention_mask: Optional[torch.Tensor] = None,
79
+ position_ids: Optional[torch.LongTensor] = None,
80
+ use_cache: Optional[bool] = None,
81
+ pixel_values: Optional[torch.FloatTensor] = None,
82
+ pixel_values_videos: Optional[torch.FloatTensor] = None,
83
+ audio_feats: Optional[torch.FloatTensor] = None,
84
+ audio_feats_lengths: Optional[torch.LongTensor] = None,
85
+ audio_placeholder_loc_lens: Optional[torch.LongTensor] = None,
86
+ image_grid_thw: Optional[torch.LongTensor] = None,
87
+ video_grid_thw: Optional[torch.LongTensor] = None,
88
+ past_key_values: Optional[List[torch.Tensor]] = None,
89
+ num_logits_to_keep: Optional[int] = 0,
90
+ image_gen: Optional[bool] = False,
91
+ image_gen_pixel_values_reference: Optional[torch.FloatTensor] = None,
92
+ image_gen_negative_input_ids: Optional[torch.LongTensor] = None,
93
+ image_gen_negative_attention_mask: Optional[torch.Tensor] = None,
94
+ image_gen_steps: Optional[int] = None,
95
+ image_gen_seed: Optional[int] = None,
96
+ image_gen_cfg: Optional[float] = None,
97
+ image_gen_image_cfg: Optional[float] = 1.0,
98
+ image_gen_cfg_mode: Optional[int] = 1,
99
+ image_gen_height: Optional[int] = None,
100
+ image_gen_width: Optional[int] = None,
101
+ image_gen_llm_hidden_states: Optional[torch.LongTensor] = None,
102
+ image_gen_negative_llm_hidden_states: Optional[torch.LongTensor] = None,
103
+ image_gen_text: Optional[list] = None,
104
+ image_gen_highres = 512,
105
+ image_gen_only_extract_hidden_states = False,
106
+ image_gen_condition_embeds=None,
107
+ image_gen_negative_condition_embeds=None,
108
+ image_gen_condition_embeds_2=None,
109
+ image_gen_negative_condition_embeds_2=None,
110
+ image_gen_return_batch=False,
111
+ image_gen_task=None,
112
+ num_frames_per_prompt=1,
113
+ **generate_kwargs,
114
+ ):
115
+ if audio_feats is not None or audio_feats_lengths is not None:
116
+ raise ValueError("audio input is not supported by Ming Image inference")
117
+ if image_gen_image_cfg not in (None, 1.0):
118
+ raise ValueError(
119
+ "image_gen_image_cfg is not supported by this inference path; "
120
+ "guidance is controlled by image_gen_cfg only"
121
+ )
122
+ image_embeds, video_embeds, audio_embeds, audio_embeds_lengths = None, None, None, None
123
+
124
+ if image_gen:
125
+ if not hasattr(self, "inference_profile"):
126
+ raise RuntimeError(
127
+ "image modules were loaded without a checkpoint inference profile"
128
+ )
129
+ self.inference_profile.validate_task(
130
+ image_gen_task,
131
+ has_reference_image=image_gen_pixel_values_reference is not None,
132
+ num_layers=num_frames_per_prompt,
133
+ )
134
+ sampling = self.inference_profile.resolve_sampling_parameters(
135
+ steps=image_gen_steps,
136
+ cfg=image_gen_cfg,
137
+ )
138
+ image_gen_steps = sampling.steps
139
+ image_gen_cfg = sampling.cfg
140
+ if image_gen_pixel_values_reference is not None:
141
+ input_channels = image_gen_pixel_values_reference.shape[1]
142
+ expected_channels = self.inference_profile.vae_input_channels
143
+ if input_channels % expected_channels != 0:
144
+ raise ValueError(
145
+ "reference image channels do not match the checkpoint "
146
+ f"VAE contract: input={input_channels}, expected a "
147
+ f"multiple of {expected_channels}"
148
+ )
149
+ condition_embeds, negative_condition_embeds = None, None
150
+ condition_embeds_2, negative_condition_embeds_2 = None, None
151
+ if (image_gen_condition_embeds is not None) or (image_gen_condition_embeds_2 is not None):
152
+ if image_gen_condition_embeds is not None:
153
+ condition_embeds = image_gen_condition_embeds
154
+ negative_condition_embeds = condition_embeds * 0.0 if image_gen_negative_condition_embeds is None else image_gen_negative_condition_embeds
155
+
156
+ if image_gen_condition_embeds_2 is not None:
157
+ condition_embeds_2 = image_gen_condition_embeds_2
158
+ negative_condition_embeds_2 = condition_embeds_2 * 0.0 if image_gen_negative_condition_embeds_2 is None else image_gen_negative_condition_embeds_2
159
+
160
+ else:
161
+ if image_gen_llm_hidden_states is None:
162
+ assert self.model is not None
163
+ assert self.vision is not None
164
+ if pixel_values is not None:
165
+ image_embeds = self.extract_image_feature(pixel_values, grid_thw=image_grid_thw)
166
+
167
+ assert self.loaded_image_gen_modules is True, "please add `load_image_gen=True` in from_pretrained() method"
168
+ assert position_ids is None
169
+
170
+
171
+ condition_embeds, condition_embeds_2 = self.get_condition_embeds_for_image_gen(
172
+ input_ids=input_ids,
173
+ attention_mask=attention_mask,
174
+ image_embeds=image_embeds,
175
+ position_ids=position_ids,
176
+ use_cache=use_cache,
177
+ image_grid_thw=image_grid_thw,
178
+ llm_hidden_states=image_gen_llm_hidden_states,
179
+ )
180
+ if condition_embeds is not None:
181
+ negative_condition_embeds = condition_embeds * 0.0
182
+
183
+ if condition_embeds_2 is not None:
184
+ negative_condition_embeds_2 = condition_embeds_2 * 0.0
185
+
186
+ # negative prompt feature is deprecated
187
+ # negative_condition_embeds = self.get_learnable_token_embeds_for_image_gen(
188
+ # input_ids=image_gen_negative_input_ids,
189
+ # attention_mask=image_gen_negative_attention_mask,
190
+ # image_embeds=image_embeds,
191
+ # position_ids=position_ids,
192
+ # use_cache=use_cache,
193
+ # image_grid_thw=image_grid_thw,
194
+ # llm_hidden_states=image_gen_negative_llm_hidden_states,
195
+ # ) if ((image_gen_negative_input_ids is not None) or (image_gen_negative_llm_hidden_states is not None)) else condition_embeds * 0.0
196
+
197
+
198
+
199
+ if image_gen_only_extract_hidden_states:
200
+ return condition_embeds, negative_condition_embeds, condition_embeds_2, negative_condition_embeds_2
201
+
202
+ assert (condition_embeds is not None) or (condition_embeds_2 is not None)
203
+ if (condition_embeds is not None) and (condition_embeds_2 is not None):
204
+ assert condition_embeds.shape[0] == condition_embeds_2.shape[0]
205
+
206
+ bsz = condition_embeds.shape[0] if condition_embeds is not None else condition_embeds_2.shape[0]
207
+
208
+ if image_gen_height is None or image_gen_width is None:
209
+ if isinstance(image_gen_highres, int):
210
+ image_gen_height, image_gen_width = [image_gen_highres] * bsz, [image_gen_highres] * bsz
211
+ elif image_gen_highres is True:
212
+ image_gen_height, image_gen_width = [1024] * bsz, [1024] * bsz
213
+ else:
214
+ image_gen_height, image_gen_width = [512] * bsz, [512] * bsz
215
+ elif isinstance(image_gen_height, torch.Tensor) or isinstance(image_gen_width, torch.Tensor):
216
+ assert isinstance(image_gen_height, torch.Tensor), image_gen_height
217
+ assert isinstance(image_gen_width, torch.Tensor), image_gen_width
218
+ image_gen_height = image_gen_height.cpu().tolist()
219
+ image_gen_width = image_gen_width.cpu().tolist()
220
+ assert len(image_gen_height) == bsz
221
+ assert len(image_gen_width) == bsz
222
+ elif isinstance(image_gen_height, int) or isinstance(image_gen_width, int):
223
+ assert isinstance(image_gen_height, int), image_gen_height
224
+ assert isinstance(image_gen_width, int), image_gen_width
225
+ image_gen_height = [image_gen_height] * bsz
226
+ image_gen_width = [image_gen_width] * bsz
227
+ else:
228
+ assert isinstance(image_gen_height, list), image_gen_height
229
+ assert isinstance(image_gen_width, list), image_gen_width
230
+ assert len(image_gen_height) == bsz
231
+ assert len(image_gen_width) == bsz
232
+
233
+
234
+ image_gen_height_diffusion_list = []
235
+ image_gen_width_diffusion_list = []
236
+ image_gen_output_resize_height = []
237
+ image_gen_output_resize_width = []
238
+ for height, width in zip(image_gen_height, image_gen_width):
239
+ closest_size, resize_size = process_ratio(ori_h=height, ori_w=width, highres=image_gen_highres)
240
+ height, width = closest_size
241
+ image_gen_height_diffusion_list.append(height)
242
+ image_gen_width_diffusion_list.append(width)
243
+ height, width = resize_size
244
+ image_gen_output_resize_height.append(height)
245
+ image_gen_output_resize_width.append(width)
246
+
247
+ image_gen_height = image_gen_height_diffusion_list[0]
248
+ assert all([i == image_gen_height for i in image_gen_height_diffusion_list])
249
+ image_gen_width = image_gen_width_diffusion_list[0]
250
+ assert all([i == image_gen_width for i in image_gen_width_diffusion_list])
251
+
252
+ if image_gen_pixel_values_reference is not None:
253
+ assert (image_gen_height, image_gen_width) == (image_gen_pixel_values_reference.shape[-2], image_gen_pixel_values_reference.shape[-1])
254
+
255
+ if image_gen_seed is None or image_gen_seed < 0:
256
+ from datetime import datetime
257
+ image_gen_seed = datetime.now().microsecond % 1000
258
+
259
+ sample_kwargs = {
260
+ "steps": image_gen_steps,
261
+ "seed": image_gen_seed,
262
+ "cfg": image_gen_cfg,
263
+ "height": image_gen_height,
264
+ "width": image_gen_width,
265
+ "cfg_mode": image_gen_cfg_mode,
266
+ "ref_x": image_gen_pixel_values_reference,
267
+ "encoder_hidden_states": condition_embeds,
268
+ "directvlm_hidden_states": condition_embeds_2,
269
+ "num_frames_per_prompt": num_frames_per_prompt,
270
+ }
271
+
272
+ image = self.diffusion_loss.sample(
273
+ **sample_kwargs,
274
+ )
275
+ if image_gen_task == "layer-decompose":
276
+ output_size = (
277
+ image_gen_output_resize_width[0],
278
+ image_gen_output_resize_height[0],
279
+ )
280
+ image = [item.resize(output_size, Image.LANCZOS) for item in image]
281
+ else:
282
+ image = [
283
+ item.resize((width, height), Image.LANCZOS)
284
+ for item, width, height in zip(
285
+ image,
286
+ image_gen_output_resize_width,
287
+ image_gen_output_resize_height,
288
+ )
289
+ ]
290
+
291
+ if (
292
+ image_gen_task != "layer-decompose"
293
+ and not image_gen_return_batch
294
+ and len(image) == 1
295
+ ):
296
+ image = image[0]
297
+
298
+ return image
299
+
300
+ if pixel_values is not None:
301
+ image_embeds = self.extract_image_feature(pixel_values, grid_thw=image_grid_thw)
302
+ if pixel_values_videos is not None:
303
+ video_embeds = self.extract_image_feature(pixel_values_videos, grid_thw=video_grid_thw)
304
+ with torch.cuda.amp.autocast(dtype=torch.bfloat16):
305
+ outputs = self.model.generate(
306
+ input_ids=input_ids,
307
+ query_embeds_image=image_embeds,
308
+ query_embeds_video=video_embeds,
309
+ query_embeds_audio=audio_embeds,
310
+ query_embeds_audio_lengths=audio_embeds_lengths,
311
+ placeholder_audio_loc_lens=audio_placeholder_loc_lens,
312
+ image_grid_thw=image_grid_thw,
313
+ image_grid_thw_video=video_grid_thw,
314
+ attention_mask=attention_mask,
315
+ position_ids=position_ids,
316
+ past_key_values=past_key_values,
317
+ use_cache=use_cache,
318
+ num_logits_to_keep=num_logits_to_keep,
319
+ **generate_kwargs,
320
+ )
321
+ return outputs
322
+
323
+ def load_image_gen_modules(self, inference_model_path, torch_dtype=torch.float32, load_image_gen_diffusion=True, load_image_gen_others=True, device=None):
324
+ inference_model_path = str(resolve_model_directory(inference_model_path))
325
+ if os.path.exists(os.path.join(inference_model_path, "byt5")):
326
+ raise ValueError(
327
+ "Ming Image inference does not support a byt5 component; "
328
+ "the public checkpoint layout has no byt5/ directory."
329
+ )
330
+ self.inference_profile = load_checkpoint_capabilities(inference_model_path)
331
+ if device is not None:
332
+ device = torch.device(device)
333
+ elif self.model is not None:
334
+ device = self.model.device
335
+ else:
336
+ device = torch.device(torch.cuda.current_device())
337
+ logger.info(f"load_image_gen_modules device={device}")
338
+ from transformers import AutoModelForCausalLM
339
+ from safetensors.torch import load_file
340
+ temp_state_dict = load_file(
341
+ os.path.join(inference_model_path, "mlp", "model.safetensors")
342
+ )
343
+ with open(os.path.join(inference_model_path, 'mlp', 'config.json'), 'r') as f:
344
+ import json
345
+ metax_config = json.load(f)
346
+ diffusion_c_input_dim = metax_config.get("diffusion_c_input_dim", 2048)
347
+ self.img_gen_scales = metax_config.get("img_gen_scales", [4, 8, 16])
348
+ self.connector_norm = metax_config.get("connector_norm", True)
349
+ self.use_vlm_directvlm_condition = metax_config.get(
350
+ "use_vlm_directvlm_condition", False
351
+ )
352
+ self.use_learnable_token_condition = metax_config.get(
353
+ "use_learnable_token_condition", True
354
+ )
355
+ self.selected_hidden_states_layers = metax_config.get(
356
+ "selected_hidden_states_layers"
357
+ )
358
+ self.diffusion_inner_dim = metax_config.get("diffusion_inner_dim")
359
+
360
+ if load_image_gen_others:
361
+ self.connector = None
362
+ self.query_tokens_dict = nn.ParameterDict()
363
+ # cumulative token index across the scales
364
+ self.scale_indices = []
365
+ current_idx = 0
366
+ for scale in self.img_gen_scales:
367
+ num_tokens = scale * scale
368
+ scale_name = f"{scale}x{scale}"
369
+ #weights = temp_state_dict[f"query_tokens_dict.{scale_name}"]
370
+ self.query_tokens_dict[scale_name] = nn.Parameter(
371
+ torch.nn.functional.normalize(torch.randn(num_tokens, self.config.llm_config.hidden_size), dim=-1)
372
+ )
373
+ current_idx += scale * scale
374
+ self.scale_indices.append(current_idx)
375
+
376
+ self.query_tokens_dict.to(torch_dtype).to(device)
377
+
378
+ if self.use_learnable_token_condition:
379
+ modified_state_dict_query_tokens = {
380
+ f"{scale}x{scale}": temp_state_dict[f"query_tokens_dict.{scale}x{scale}"]
381
+ for scale in self.img_gen_scales
382
+ }
383
+
384
+ self.query_tokens_dict.load_state_dict(modified_state_dict_query_tokens, strict=True)
385
+
386
+ # self.norm_query_embeds = True
387
+ # load connector
388
+ self.connector = AutoModelForCausalLM.from_pretrained(inference_model_path, subfolder='connector', torch_dtype=torch_dtype)
389
+ for layer in self.connector.model.layers:
390
+ layer.self_attn.is_causal = False
391
+ self.connector.to(device)
392
+
393
+
394
+ self.proj_in = nn.Linear(self.config.llm_config.hidden_size, self.connector.config.hidden_size)
395
+ self.proj_out = nn.Linear(self.connector.config.hidden_size, diffusion_c_input_dim)
396
+
397
+ modified_state_dict_in = {
398
+ 'weight': temp_state_dict['proj_in.weight'],
399
+ 'bias': temp_state_dict['proj_in.bias']
400
+ }
401
+ self.proj_in.load_state_dict(modified_state_dict_in, strict=True)
402
+ modified_state_dict_out = {
403
+ 'weight': temp_state_dict['proj_out.weight'],
404
+ 'bias': temp_state_dict['proj_out.bias']
405
+ }
406
+ self.proj_out.load_state_dict(modified_state_dict_out, strict=True)
407
+ self.proj_in.to(device=device, dtype=torch_dtype)
408
+ self.proj_out.to(device=device, dtype=torch_dtype)
409
+
410
+ self.proj_directvlm = None
411
+ if self.use_vlm_directvlm_condition:
412
+ directvlm_dim = self.model.config.hidden_size
413
+ if self.selected_hidden_states_layers is not None:
414
+ directvlm_dim = directvlm_dim * len(self.selected_hidden_states_layers)
415
+
416
+ self.proj_directvlm = nn.Sequential(RMSNorm(directvlm_dim, eps=1e-5), nn.Linear(directvlm_dim, self.diffusion_inner_dim, bias=True))
417
+
418
+ modified_state_dict_directvlm = {
419
+ '0.weight': temp_state_dict["proj_directvlm.0.weight"],
420
+ '1.weight': temp_state_dict["proj_directvlm.1.weight"],
421
+ '1.bias': temp_state_dict["proj_directvlm.1.bias"],
422
+ }
423
+ self.proj_directvlm.load_state_dict(modified_state_dict_directvlm, strict=True)
424
+ self.proj_directvlm.to(device=device, dtype=torch_dtype)
425
+
426
+ if load_image_gen_diffusion:
427
+ diffusion_mlp_state_dict = {
428
+ key[len("mlp.") :] : temp_state_dict[key]
429
+ for key in temp_state_dict if key.startswith("mlp.")
430
+ }
431
+ from diffusion.generator import ImageGenerator
432
+
433
+ self.diffusion_loss = ImageGenerator(
434
+ model_path=inference_model_path,
435
+ scheduler_path=inference_model_path,
436
+ vision_dim=diffusion_c_input_dim,
437
+ mlp_state_dict=diffusion_mlp_state_dict,
438
+ torch_dtype=torch_dtype,
439
+ device=device,
440
+ use_identity_mlp=metax_config.get("use_identity_mlp", False),
441
+ text_encoder_norm=metax_config.get("text_encoder_norm", False),
442
+ inference_profile=self.inference_profile,
443
+ )
444
+ self.diffusion_loss.to(device)
445
+ self.loaded_image_gen_modules = True
446
+ @classmethod
447
+ def _from_int8_checkpoint(cls, vlm_directory, device, **kwargs):
448
+ """Load an mllm/ component written by quant/quantize_stream.py (weight-only int8).
449
+
450
+ The model is built with its parameters on the meta device, the Linear modules listed
451
+ in int8_manifest.json become Int8Linear shells, and every stored tensor is loaded
452
+ straight onto `device`, so BF16 weights for the quantized modules never exist in memory.
453
+ """
454
+ from accelerate import init_empty_weights
455
+ from quant.load_int8 import load_int8_mllm_
456
+
457
+ device_map = kwargs.pop("device_map", None)
458
+ if device_map is not None:
459
+ # infer.py's default "balanced" plan on a single-GPU box maps every module to GPU 0;
460
+ # that is honoured. Splitting the int8 model across devices is not supported.
461
+ targets = set(device_map.values()) if isinstance(device_map, dict) else {device_map}
462
+ if len(targets) != 1 or not isinstance(next(iter(targets)), int):
463
+ raise ValueError(
464
+ "the int8 mllm checkpoint loads onto a single GPU; device_map targets "
465
+ f"{sorted(map(str, targets))} (use --device-map none)"
466
+ )
467
+ device = torch.device("cuda", next(iter(targets)))
468
+ supported = ("torch_dtype", "dtype", "attn_implementation")
469
+ unsupported = sorted(key for key in kwargs if key not in supported)
470
+ if unsupported:
471
+ raise ValueError(
472
+ f"the int8 mllm checkpoint loads onto a single device; unsupported arguments: {unsupported}"
473
+ )
474
+ device = torch.device(device) if device is not None else torch.device("cpu")
475
+ if device.type == "cuda" and device.index is None:
476
+ device = torch.device("cuda", torch.cuda.current_device())
477
+ config = BailingMM2Config.from_pretrained(vlm_directory)
478
+ with init_empty_weights():
479
+ model = cls._from_config(config, **kwargs)
480
+ report = load_int8_mllm_(model, vlm_directory, device)
481
+ # Buffers built in __init__ (rotary inv_freq) are not stored in the checkpoint; they follow
482
+ # the weights. Int8Linear keeps its fp32 scales through this and any later dtype cast.
483
+ model.to(device)
484
+ logger.info(f"int8 mllm loaded from {vlm_directory}: {report}")
485
+ model.tie_weights()
486
+ model.eval()
487
+ return model
488
+
489
+ @classmethod
490
+ def from_pretrained(
491
+ cls,
492
+ pretrained_model_name_or_path: Optional[Union[str, os.PathLike]],
493
+ *model_args,
494
+ **kwargs,
495
+ ):
496
+ load_image_gen = False
497
+ if "load_image_gen" in kwargs:
498
+ load_image_gen = kwargs["load_image_gen"]
499
+ del kwargs["load_image_gen"]
500
+ load_image_gen_diffusion = True
501
+ if "load_image_gen_diffusion" in kwargs:
502
+ load_image_gen_diffusion = kwargs["load_image_gen_diffusion"]
503
+ del kwargs["load_image_gen_diffusion"]
504
+
505
+ load_image_gen_others = True
506
+ if "load_image_gen_others" in kwargs:
507
+ load_image_gen_others = kwargs["load_image_gen_others"]
508
+ del kwargs["load_image_gen_others"]
509
+ load_vlm = True
510
+ if "load_vlm" in kwargs:
511
+ load_vlm = kwargs["load_vlm"]
512
+ del kwargs["load_vlm"]
513
+ image_gen_device = kwargs.pop("image_gen_device", None)
514
+ vlm_directory = pretrained_model_name_or_path
515
+ if load_image_gen:
516
+ pretrained_model_name_or_path = str(
517
+ resolve_model_directory(
518
+ pretrained_model_name_or_path,
519
+ revision=kwargs.get("revision"),
520
+ cache_dir=kwargs.get("cache_dir"),
521
+ local_files_only=kwargs.get("local_files_only", False),
522
+ token=kwargs.get("token"),
523
+ )
524
+ )
525
+ # The package root keeps connector/mlp/transformer/vae/scheduler;
526
+ # the MLLM itself (config, weights, tokenizer data) lives in mllm/.
527
+ vlm_directory = os.path.join(pretrained_model_name_or_path, "mllm")
528
+ if not os.path.isdir(vlm_directory):
529
+ raise FileNotFoundError(
530
+ "checkpoint package is missing the mllm/ component: "
531
+ f"{vlm_directory}. Migrate the package to the component "
532
+ "layout before loading."
533
+ )
534
+ if load_vlm and os.path.exists(os.path.join(vlm_directory, "int8_manifest.json")):
535
+ model = cls._from_int8_checkpoint(vlm_directory, image_gen_device, **kwargs)
536
+ elif load_vlm:
537
+ model = super().from_pretrained(
538
+ vlm_directory,
539
+ *model_args,
540
+ **kwargs,
541
+ )
542
+ else:
543
+ model = cls(
544
+ BailingMM2Config.from_dict(BailingMM2Config.get_config_dict(vlm_directory)[0]),
545
+ empty_load=True,
546
+ )
547
+ if load_image_gen:
548
+ model.load_image_gen_modules(
549
+ pretrained_model_name_or_path,
550
+ torch_dtype=kwargs["torch_dtype"] if "torch_dtype" in kwargs else torch.float32,
551
+ load_image_gen_diffusion=load_image_gen_diffusion,
552
+ load_image_gen_others=load_image_gen_others,
553
+ device=image_gen_device,
554
+ )
555
+ return model
556
+
557
+ def append_input_ids_with_multiscale_learnable_tokens(
558
+ self,
559
+ text_ids,
560
+ attention_mask,
561
+ scales,
562
+ start_token_id,
563
+ end_token_id,
564
+ patch_token_id,
565
+ ):
566
+ default_scaled_tokens = []
567
+ default_scaled_attn_masks = []
568
+ default_gen_masks = []
569
+ for scale in scales:
570
+ default_scaled_tokens.append(start_token_id)
571
+ default_scaled_tokens.extend([patch_token_id for _ in range(scale * scale)])
572
+ default_scaled_tokens.append(end_token_id)
573
+ default_scaled_attn_masks.extend([1 for _ in range(scale * scale + 2)])
574
+ default_gen_masks.append(0)
575
+ default_gen_masks.extend([1 for _ in range(scale * scale)])
576
+ default_gen_masks.append(0)
577
+
578
+ text_ids_list = text_ids.cpu().tolist()
579
+ attention_mask_list = attention_mask.cpu().tolist()
580
+
581
+ new_text_ids_list = []
582
+ new_attention_mask_list = []
583
+ gen_mask_list = []
584
+ new_labels_list = []
585
+ for text_ids_one_batch, attention_mask_one_batch in zip(
586
+ text_ids_list, attention_mask_list
587
+ ):
588
+ assert len(text_ids_one_batch) == len(attention_mask_one_batch)
589
+
590
+ padding_start = 0
591
+ for idx, value in enumerate(attention_mask_one_batch):
592
+ if value == 0:
593
+ break
594
+
595
+ padding_start += 1
596
+
597
+ new_text_ids_list.append(text_ids_one_batch[:padding_start] + deepcopy(default_scaled_tokens) + text_ids_one_batch[padding_start:])
598
+ new_labels_list.append([ -100 for _ in range(padding_start)] + [1 for _ in range(len(default_scaled_tokens))] + [-100 for _ in range(len(text_ids_one_batch[padding_start:]))] )
599
+
600
+ new_attention_mask_list.append(attention_mask_one_batch[:padding_start] + deepcopy(default_scaled_attn_masks) + attention_mask_one_batch[padding_start:])
601
+ gen_mask_list.append(
602
+ [0 for _ in range(len(attention_mask_one_batch[:padding_start]))] + \
603
+ deepcopy(default_gen_masks) + \
604
+ [0 for _ in range(len(attention_mask_one_batch[padding_start:]))]
605
+ )
606
+
607
+ text_ids_append_lq = torch.tensor(new_text_ids_list, dtype=text_ids.dtype).to(text_ids.device)
608
+ attention_mask_append_lq = torch.tensor(new_attention_mask_list, dtype=attention_mask.dtype).to(attention_mask.device)
609
+ gen_mask = torch.tensor(gen_mask_list, dtype=attention_mask.dtype).to(attention_mask.device)
610
+ labels = torch.tensor(new_labels_list, dtype=text_ids.dtype).to(text_ids.device)
611
+
612
+ assert attention_mask_append_lq.shape == text_ids_append_lq.shape
613
+ assert labels.shape == text_ids_append_lq.shape
614
+ assert gen_mask.shape == text_ids_append_lq.shape
615
+ return text_ids_append_lq, labels, attention_mask_append_lq, gen_mask
616
+
617
+ def appand_learnable_tokens(
618
+ self,
619
+ text_ids,
620
+ gen_mask,
621
+ image_embeds,
622
+ image_grid_thw,
623
+ patch_token_id,
624
+ ):
625
+ query_tokens_embeds = torch.cat(
626
+ [self.query_tokens_dict[f"{scale}x{scale}"] for scale in self.img_gen_scales],
627
+ dim=0,
628
+ )
629
+ if image_embeds is not None:
630
+ query_tokens_embeds = query_tokens_embeds.to(image_embeds.dtype).to(image_embeds.device)
631
+
632
+ assert text_ids.shape == gen_mask.shape
633
+ text_ids_aslist = text_ids.cpu().view(-1).tolist()
634
+ gen_mask_aslist = gen_mask.cpu().view(-1).tolist()
635
+ is_patch_list = [1 if i == patch_token_id else 0 for i in text_ids_aslist]
636
+ idxes_start_of_patch = find_first_index_of_consecutive_ones(is_patch_list)
637
+ isgen_indicators = merge_consecutive_ones([1 if gen_mask_aslist[i] else 0 for i in idxes_start_of_patch], len(self.img_gen_scales))
638
+ if any([i == 0 for i in isgen_indicators]):
639
+ assert image_grid_thw is not None
640
+ assert image_grid_thw.ndim == 2
641
+ assert image_embeds is not None
642
+ assert image_embeds.ndim == 2
643
+
644
+ new_image_grid_thw = []
645
+ new_image_embeds = []
646
+ cum_image_token = 0
647
+ cnt_input_image = 0
648
+
649
+ for is_gen in isgen_indicators:
650
+ if is_gen:
651
+ for scale in self.img_gen_scales:
652
+ new_image_grid_thw.append([1, 2, scale * scale * 2])
653
+
654
+ new_image_embeds.append(query_tokens_embeds)
655
+ else:
656
+ thw = image_grid_thw[cnt_input_image].tolist()
657
+ assert thw[0] == 1
658
+ assert thw[1] % 2 == 0 # h
659
+ assert thw[2] % 2 == 0 # w
660
+ n_image_token = (thw[1] // 2) * (thw[2] // 2)
661
+ image_embed_one = image_embeds[cum_image_token : cum_image_token + n_image_token, :]
662
+ new_image_embeds.append(image_embed_one)
663
+ new_image_grid_thw.append(thw)
664
+ cnt_input_image += 1
665
+ cum_image_token += n_image_token
666
+
667
+ if image_grid_thw is not None:
668
+ assert cnt_input_image == image_grid_thw.shape[0]
669
+ assert cum_image_token == image_embeds.shape[0]
670
+ else:
671
+ assert cnt_input_image == 0
672
+ assert cum_image_token == 0
673
+
674
+ new_image_grid_thw = torch.tensor(new_image_grid_thw, dtype=text_ids.dtype).to(text_ids.device)
675
+ new_image_embeds = torch.cat(new_image_embeds, dim=0).to(text_ids.device)
676
+
677
+ total_patch_token = 0
678
+ for bid in range(new_image_grid_thw.shape[0]):
679
+ thw = new_image_grid_thw[bid].tolist()
680
+ assert thw[0] == 1
681
+ assert thw[1] % 2 == 0
682
+ assert thw[2] % 2 == 0
683
+ patch_h = thw[1] // 2
684
+ patch_w = thw[2] // 2
685
+ n_patch_token = patch_h * patch_w
686
+ total_patch_token += n_patch_token
687
+
688
+ # if torch.distributed.get_rank() == 0:
689
+ # embed()
690
+ # torch.distributed.barrier()
691
+
692
+ assert total_patch_token == new_image_embeds.shape[0], f"{total_patch_token}, vs. {new_image_embeds.shape}"
693
+
694
+ return new_image_grid_thw, new_image_embeds
695
+
696
+ def get_condition_embeds_for_image_gen(
697
+ self,
698
+ input_ids,
699
+ attention_mask,
700
+ image_embeds,
701
+ position_ids,
702
+ use_cache,
703
+ image_grid_thw,
704
+ llm_hidden_states,
705
+ ):
706
+ input_ids, labels, attention_mask, gen_mask = self.append_input_ids_with_multiscale_learnable_tokens(
707
+ input_ids,
708
+ attention_mask,
709
+ self.img_gen_scales,
710
+ self.config.llm_config.image_patch_token + 1,
711
+ self.config.llm_config.image_patch_token + 2,
712
+ self.config.llm_config.image_patch_token,
713
+ )
714
+
715
+ if llm_hidden_states is None:
716
+ image_grid_thw, image_embeds = self.appand_learnable_tokens(
717
+ input_ids,
718
+ gen_mask,
719
+ image_embeds,
720
+ image_grid_thw,
721
+ self.config.llm_config.image_patch_token,
722
+ )
723
+
724
+ with torch.cuda.amp.autocast(dtype=torch.bfloat16):
725
+ if image_embeds is None or input_ids.size(1) == 1:
726
+ words_embeddings = self.model.get_input_embeddings()(input_ids.clip(0, self.model.get_input_embeddings().weight.shape[0] - 1))
727
+ image_mask = None
728
+ audio_mask = None
729
+ else:
730
+ words_embeddings, image_mask, audio_mask = self.model.model.prompt_wrap_navit(
731
+ input_ids=input_ids.clip(0, self.model.get_input_embeddings().weight.shape[0] - 1),
732
+ config=self.model.model.config,
733
+ query_embeds_image=image_embeds,
734
+ )
735
+
736
+ assert input_ids.size(1) == words_embeddings.size(1), "{} vs {}".format(
737
+ input_ids.size,
738
+ words_embeddings.size,
739
+ )
740
+
741
+ # if torch.distributed.get_rank() == 3:
742
+ # embed()
743
+ # torch.distributed.barrier()
744
+
745
+ outputs = self.model.forward(
746
+ input_ids=input_ids,
747
+ attention_mask=attention_mask,
748
+ position_ids=position_ids,
749
+ past_key_values=None,
750
+ inputs_embeds=words_embeddings,
751
+ image_grid_thw=image_grid_thw,
752
+ use_cache=False,
753
+ image_mask=image_mask,
754
+ audio_mask=None,
755
+ output_hidden_states=True,
756
+ )
757
+ hidden_states = outputs.hidden_states[-1]
758
+ else:
759
+ hidden_states = llm_hidden_states
760
+
761
+ directvlm_hidden_states = None
762
+ if self.use_vlm_directvlm_condition:
763
+ # use hidden states
764
+ use_input_mask = torch.lt(labels, 0).int().to(attention_mask.dtype) * attention_mask
765
+ assert use_input_mask.ndim == 2
766
+ directvlm_max_valid_ind = use_input_mask.cumsum(-1).argmax(-1).max().item() + 1
767
+ #directvlm_max_valid_ind = min(directvlm_max_valid_ind, self.max_vlm_directvlm_length)
768
+ use_input_mask = use_input_mask[:, :directvlm_max_valid_ind]
769
+
770
+ if self.selected_hidden_states_layers is not None:
771
+ directvlm_hidden_states = torch.cat([
772
+ outputs.hidden_states[layer_i].to(labels.device)[:, :directvlm_max_valid_ind, :] * use_input_mask.unsqueeze(-1)
773
+ for layer_i in self.selected_hidden_states_layers
774
+ ], dim=-1)
775
+ else:
776
+ directvlm_hidden_states = outputs.hidden_states[-1].to(labels.device)[:, :directvlm_max_valid_ind, :] * use_input_mask.unsqueeze(-1)
777
+
778
+ directvlm_hidden_states = directvlm_hidden_states.detach()
779
+ directvlm_hidden_states = self.proj_directvlm(directvlm_hidden_states)
780
+
781
+ scale_embeds = None
782
+ if self.use_learnable_token_condition:
783
+ with torch.cuda.amp.autocast(dtype=torch.bfloat16):
784
+ gen_mask = gen_mask.unsqueeze(-1).expand(gen_mask.shape[0], gen_mask.shape[1], hidden_states.shape[-1]).to(hidden_states.device).bool()
785
+ hidden_states_gen = torch.masked_select(hidden_states, gen_mask).view(hidden_states.shape[0], -1, hidden_states.shape[-1])
786
+ # split hidden_states into per-scale representations
787
+ scale_start_idxes = [0] + self.scale_indices[:-1]
788
+ scale_end_idxes = self.scale_indices
789
+ assert scale_end_idxes[-1] == hidden_states_gen.shape[1]
790
+
791
+ scale, scale_start_idx, scale_end_idx = [
792
+ i for i in zip(self.img_gen_scales, scale_start_idxes, scale_end_idxes)
793
+ ][-1]
794
+
795
+ scale_hidden = hidden_states_gen[:, scale_start_idx : scale_end_idx, :]
796
+ scale_embeds = self.proj_in(scale_hidden)
797
+ seq_shape = scale_embeds.shape
798
+ with torch.cuda.amp.autocast(dtype=torch.bfloat16):
799
+ scale_embeds = self.connector(
800
+ inputs_embeds=scale_embeds,
801
+ attention_mask=torch.ones(seq_shape[0],1,seq_shape[1],seq_shape[1]).to(scale_embeds.device),
802
+ output_hidden_states=True
803
+ ).hidden_states[-1]
804
+
805
+ scale_embeds = self.proj_out(scale_embeds)
806
+ # normalize
807
+ if self.connector_norm:
808
+ scale_embeds = torch.nn.functional.normalize(scale_embeds, dim=-1)
809
+
810
+ return scale_embeds, directvlm_hidden_states
811
+
812
+ __all__ = [
813
+ "BailingMM2NativeForConditionalGeneration"
814
+ ]
code/modeling_utils.py ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ from typing import Optional
4
+
5
+ import logging
6
+
7
+ logger = logging.getLogger(__name__)
8
+
9
+ class Transpose(nn.Module):
10
+ def __init__(self, dim0: int, dim1: int):
11
+ super().__init__()
12
+ self.dim0 = dim0
13
+ self.dim1 = dim1
14
+
15
+ def forward(self, x):
16
+ return x.transpose(self.dim0, self.dim1)
17
+
18
+ def patch_continuous_features(
19
+ input_embeddings: torch.Tensor,
20
+ placeholder_loc_lens: torch.Tensor,
21
+ encoded_feats: torch.Tensor,
22
+ encoded_feat_lens: torch.Tensor,
23
+ ):
24
+ """
25
+ Patch continuous features into input embeddings, while keeping a valid gradient flow.
26
+
27
+ input_embeddings: torch.Tensor, size = [B, C?, T, D]
28
+ placeholder_loc_lens: torch.LongTensor, size = [B, N, 2]
29
+ Each 2-tuple represents (start, length) of a placeholder.
30
+ encoded_feats: torch.Tensor, size = [B, L1 + L2 + ... + LN, ...]
31
+ encoded_feat_lens: torch.LongTensor, size = [B, N]
32
+
33
+ Example ('X' for patch placeholder tokens):
34
+ Inputs:
35
+ input_embeddings = [[1, 2, 3, X, X, X, 4, 5, 6, X, X, X, 7, 8]]
36
+ placeholder_loc_lens = [[[3, 3]], [[9, 3]]]
37
+ encoded_feats = [[A, A, A, B, B]]
38
+ encoded_feat_lens = [[3], [2]]
39
+ Outputs:
40
+ embeddings = [[1, 2, 3, A, A, A, 4, 5, 6, B, B, X, 7, 8]]
41
+ """
42
+ batch_size = input_embeddings.size(0)
43
+ for i in range(batch_size):
44
+ audio_feat_start = 0
45
+ for j in range(placeholder_loc_lens.shape[1]):
46
+ placeholder_start: int = int(placeholder_loc_lens[i, j, 0].item())
47
+ placeholder_len: int = int(placeholder_loc_lens[i, j, 1].item())
48
+ if placeholder_len <= 0:
49
+ break
50
+ feat_len = int(encoded_feat_lens[i, j].item())
51
+ real_feat_len = feat_len
52
+ if feat_len > placeholder_len:
53
+ logger.warning(
54
+ f"Feature length ({feat_len}) > placeholder length ({placeholder_len}). "
55
+ "This is not expected; please check estimate_audio_feature_length(). "
56
+ "We truncate the feature to avoid errors."
57
+ )
58
+ feat_len = placeholder_len
59
+ target_len = min(feat_len, placeholder_len)
60
+ input_embeddings[i, placeholder_start:placeholder_start + target_len] = encoded_feats[i, audio_feat_start:audio_feat_start + target_len]
61
+ audio_feat_start += real_feat_len
62
+ return input_embeddings
63
+
64
+ def build_modality_mask(placeholder_loc_lens: torch.Tensor, shape: torch.Size):
65
+ mask = torch.zeros(shape, dtype=torch.bool)
66
+ for i in range(placeholder_loc_lens.shape[0]):
67
+ for j in range(placeholder_loc_lens.shape[1]):
68
+ start: int = int(placeholder_loc_lens[i, j, 0].item())
69
+ length: int = int(placeholder_loc_lens[i, j, 1].item())
70
+ if length <= 0:
71
+ break
72
+ mask[i, start:start + length] = True
73
+ return mask
74
+
75
+ def encode_audio_segments(
76
+ encoder,
77
+ proj_layer,
78
+ wav_feats=None,
79
+ wav_feats_lengths=None,
80
+ waveforms=None,
81
+ waveforms_lengths=None,
82
+ use_waveform=False,
83
+ audio_config=None,
84
+ ):
85
+ """
86
+ Apply audio encoder to input audio features in wrapped format.
87
+ See the documentation of unwrap_feats() for details about 'wrapped format'.
88
+ """
89
+
90
+ # Forward audio encoder.
91
+ if use_waveform:
92
+ assert waveforms is not None and waveforms_lengths is not None
93
+ # Unwrap the waveforms so each waveform is placed at an independent row.
94
+ waveform_segs_batch, waveform_seg_lengths = unwrap_feats(waveforms, waveforms_lengths)
95
+ audio_feats_seg, audio_feat_seg_lengths = encoder(waveform_segs_batch, waveform_seg_lengths)[:2]
96
+ else:
97
+ assert wav_feats is not None and wav_feats_lengths is not None
98
+ # Unwrap the features so the feature of each waveform is placed at an independent row.
99
+ feat_segs_batch, feat_seg_lengths = unwrap_feats(wav_feats, wav_feats_lengths)
100
+ # for whisper encoder
101
+ # feat_segs_batch: [B, T, n_mels]
102
+ # feat_seg_lengths: [B]
103
+ audio_feats_seg = encoder(feat_segs_batch)
104
+ audio_feats_seg_proj = proj_layer(audio_feats_seg.transpose(-1, -2)).transpose(-1, -2)
105
+ feat_seg_lengths = feat_seg_lengths.to(feat_segs_batch.device)
106
+ # whisper encoder conv
107
+ audio_feat_seg_lengths = (feat_seg_lengths - 3 + 2 * 1) // 2 + 1
108
+ # project layer conv
109
+ audio_feat_seg_lengths = (audio_feat_seg_lengths - audio_config.ds_kernel_size + 2 *
110
+ (audio_config.ds_kernel_size//2)) // audio_config.ds_stride + 1
111
+
112
+ # Wrap the features so the 1st dim represents batch_size.
113
+ input_lengths = waveforms_lengths if use_waveform else wav_feats_lengths
114
+ assert input_lengths is not None
115
+ audio_feats, _, audio_feats_lengths = wrap_feats(audio_feats_seg, input_lengths, audio_feat_seg_lengths)
116
+ audio_feats_proj, _, audio_feats_lengths2 = wrap_feats(audio_feats_seg_proj, input_lengths, audio_feat_seg_lengths)
117
+ assert torch.all(audio_feats_lengths == audio_feats_lengths2), f"{audio_feats_lengths}, {audio_feats_lengths2}"
118
+
119
+ return audio_feats_proj, audio_feats, audio_feats_lengths
120
+
121
+ def unwrap_feats(feats: torch.Tensor, feats_lengths: torch.Tensor):
122
+ """
123
+ The input feats are in the "wrapped" format, which means that features from (at most) N audios are concatenated
124
+ as a single sample feats[i]. In this case, each row of feats_lengths contains the lengths of the concatenated
125
+ feature. This function unwraps the features.
126
+ For samples with less than N segments, one should pad feats_lengths with 0. The result will contain valid
127
+ segments only.
128
+
129
+ feats: torch.Tensor, size = [B, L1 + L2 + ... + LN, ...]
130
+ feats_lengths: torch.LongTensor, size = [B, N]
131
+
132
+ Example ('X' for padding):
133
+ Inputs:
134
+ feats = [[A, A, A, A, X],
135
+ [B, B, C, C, C]]
136
+ feats_lengths = [[4, 0],
137
+ [2, 3]]
138
+ Outputs:
139
+ feat_segs = [[A, A, A, A],
140
+ [B, B, X, X],
141
+ [C, C, C, X]]
142
+ feat_seg_lengths = [4, 2, 3]
143
+ """
144
+ feat_segs = []
145
+ feat_seg_lengths = []
146
+ for i in range(feats_lengths.shape[0]):
147
+ feat_index = 0
148
+ for j in range(feats_lengths.shape[1]):
149
+ feat_len = feats_lengths[i, j].item()
150
+ if feat_len == 0: break
151
+ feat_segs.append(feats[i, feat_index:feat_index + feat_len])
152
+ feat_seg_lengths.append(feat_len)
153
+ feat_index += feat_len
154
+ feat_segs_batch = torch.nn.utils.rnn.pad_sequence(feat_segs, True).to(feats.device)
155
+ feat_seg_lengths = torch.tensor(feat_seg_lengths, dtype=torch.long, device=feats.device)
156
+ return feat_segs_batch, feat_seg_lengths
157
+
158
+ def wrap_feats(feat_segs: torch.Tensor, feats_lengths: torch.Tensor, feats_seg_lengths: Optional[torch.Tensor] = None):
159
+ """
160
+ Wrap segmented features back to the wrapped format.
161
+ This function is the inverse operation of unwrap_feats(). See its documentation for details.
162
+ Note that the feats_lengths value does not matter a lot. We only check the location of the first 0 to determine the
163
+ number of feature segments.
164
+ """
165
+ feat_idx = 0
166
+ feats_buffer = []
167
+ feats_locs_buffer = []
168
+ feats_lengths_buffer = []
169
+ for i in range(feats_lengths.shape[0]):
170
+ feat_buffer = []
171
+ feat_locs_buffer = []
172
+ feat_lengths_buffer = []
173
+ feat_total_len = 0
174
+ for j in range(feats_lengths.shape[1]):
175
+ feat_len = feats_lengths[i, j].item()
176
+ if feat_len == 0:
177
+ break
178
+ if feats_seg_lengths is not None:
179
+ feat_len = feats_seg_lengths[feat_idx].item()
180
+ feat_buffer.append(feat_segs[feat_idx, :feat_len])
181
+ feat_locs_buffer.append(feat_total_len)
182
+ feat_lengths_buffer.append(feat_len)
183
+ feat_idx += 1
184
+ feat_total_len += feat_len
185
+ feats_buffer.append(torch.cat(feat_buffer))
186
+ feats_locs_buffer.append(torch.tensor(feat_locs_buffer, dtype=torch.long))
187
+ feats_lengths_buffer.append(torch.tensor(feat_lengths_buffer, dtype=torch.long))
188
+ feats = torch.nn.utils.rnn.pad_sequence(feats_buffer, True).to(feat_segs.device)
189
+ feats_locs = torch.nn.utils.rnn.pad_sequence(feats_locs_buffer, True).to(feats_lengths.device)
190
+ feats_new_lengths = torch.nn.utils.rnn.pad_sequence(feats_lengths_buffer, True).to(feats_lengths.device)
191
+ return feats, feats_locs, feats_new_lengths
code/pe_ling.py ADDED
@@ -0,0 +1,446 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Prompt enhancement (PE) for Ming-Image text-to-image via a Ling-3.0-flash-VL seat.
3
+
4
+ Per the README, PE is a pre-processing step *outside* ``infer.py``: an
5
+ instruction-following VLM rewrites a short caption into the structured
6
+ Figma-style JSON prompt that the text-to-image pipeline consumes, and the
7
+ result is passed to ``infer.py --prompt`` as raw text or via a file.
8
+
9
+ This module drives any OpenAI-compatible ``/chat/completions`` endpoint using
10
+ only the standard library (``urllib``): by default the local llama-server seat
11
+ serving Ling-3.0-flash-VL, optionally the LiteLLM lab gateway (Bearer auth via
12
+ ``--api-key`` or the ``LITELLM_API_KEY`` environment variable). The rewriter
13
+ system prompt is read verbatim from ``assets/t2i_rewriter_system_prompt.txt``.
14
+
15
+ The reply is parsed robustly (```json fences and surrounding prose are
16
+ tolerated), then validated against the schema the system prompt demands. On a
17
+ parse or validation failure the request is retried exactly once with the
18
+ errors appended to the user turn; if that still fails, PromptEnhancementError
19
+ is raised with the errors. Invalid JSON is never passed through silently.
20
+
21
+ CLI:
22
+ python pe_ling.py "a caption" --out prompt.json \
23
+ [--base-url http://127.0.0.1:8090/v1] \
24
+ [--model ling-3.0-flash-vl-mtp-halo-STRIX_LEAN]
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import argparse
30
+ import json
31
+ import os
32
+ import re
33
+ import sys
34
+ import time
35
+ import urllib.error
36
+ import urllib.request
37
+ from pathlib import Path
38
+ from typing import Any, Dict, List, Optional, Tuple
39
+
40
+ CODE_DIRECTORY = Path(__file__).resolve().parent
41
+ SYSTEM_PROMPT_PATH = CODE_DIRECTORY / "assets" / "t2i_rewriter_system_prompt.txt"
42
+
43
+ # The Ling-3.0-flash-VL seat already served on the target box (llama-server,
44
+ # OpenAI-compatible, thinking disabled); both endpoints speak the same
45
+ # /chat/completions protocol.
46
+ DEFAULT_BASE_URL = "http://127.0.0.1:8090/v1"
47
+ DEFAULT_MODEL = "ling-3.0-flash-vl-mtp-halo-STRIX_LEAN"
48
+ API_KEY_ENV = "LITELLM_API_KEY"
49
+
50
+ # Low temperature: the rewrite is a deterministic schema transformation, not
51
+ # creative sampling.
52
+ DEFAULT_TEMPERATURE = 0.2
53
+ # The upstream example rewrite (assets/t2i_four_seasons_cabin_prompt.json) is
54
+ # ~5 KB (~2k tokens); dense multi-layer infographic rewrites run several times
55
+ # longer, so leave generous headroom for a complete JSON object.
56
+ DEFAULT_MAX_TOKENS = 16384
57
+ # A multi-thousand-token completion on the local seat can take minutes.
58
+ DEFAULT_TIMEOUT_SECONDS = 600.0
59
+
60
+ REPAIR_INSTRUCTION = "Return only the corrected JSON object: no prose, no code fences."
61
+
62
+ CANVAS_SETTINGS_KEYS = ("aspect_ratio", "ambient_lighting", "image_style")
63
+ LAYER_KEYS = ("description", "coordinates", "hierarchy_and_relation", "color_specs")
64
+ COORDINATE_FIELDS = ("cx", "cy", "w", "h")
65
+
66
+ # `coordinates` must be ONE string of the form
67
+ # "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000". The upstream example also uses
68
+ # bare integers ("h: 1"), so accept any decimal spelling and enforce the
69
+ # [0, 1] range on the parsed value. Whitespace around ':' and ',' is
70
+ # tolerated; the key order is fixed.
71
+ _COORDINATE_NUMBER = r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)"
72
+ COORDINATES_RE = re.compile(
73
+ rf"^\s*cx:\s*(?P<cx>{_COORDINATE_NUMBER})\s*,"
74
+ rf"\s*cy:\s*(?P<cy>{_COORDINATE_NUMBER})\s*,"
75
+ rf"\s*w:\s*(?P<w>{_COORDINATE_NUMBER})\s*,"
76
+ rf"\s*h:\s*(?P<h>{_COORDINATE_NUMBER})\s*$"
77
+ )
78
+
79
+ # Hex colors: #RGB, #RGBA, #RRGGBB, #RRGGBBAA (the upstream example uses
80
+ # #RRGGBB; the alpha forms keep RGBA-design outputs from failing validation).
81
+ HEX_COLOR_RE = re.compile(
82
+ r"^#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{4}|[0-9a-fA-F]{6}|[0-9a-fA-F]{8})$"
83
+ )
84
+
85
+
86
+ class PromptEnhancementError(RuntimeError):
87
+ """PE failed: transport/protocol error, or schema failure after the retry."""
88
+
89
+ def __init__(
90
+ self,
91
+ message: str,
92
+ errors: Optional[List[str]] = None,
93
+ reply: Optional[str] = None,
94
+ ):
95
+ super().__init__(message)
96
+ self.errors = list(errors or [])
97
+ self.reply = reply
98
+
99
+
100
+ def load_system_prompt(path: Path = SYSTEM_PROMPT_PATH) -> str:
101
+ """Return the released rewriter system prompt, verbatim."""
102
+ return path.read_text(encoding="utf-8")
103
+
104
+
105
+ def extract_json_object(text: str) -> Dict[str, Any]:
106
+ """Return the first complete top-level JSON object found in ``text``.
107
+
108
+ Models sometimes wrap JSON in ```json fences or add prose around it.
109
+ Scanning every ``{`` position with ``JSONDecoder.raw_decode`` (which
110
+ decodes a document at an offset and ignores trailing data) recovers the
111
+ object in all of those shapes. Raises ValueError when no complete JSON
112
+ object is present, e.g. a reply truncated mid-object.
113
+ """
114
+ decoder = json.JSONDecoder()
115
+ position = text.find("{")
116
+ while position != -1:
117
+ try:
118
+ document, _ = decoder.raw_decode(text, position)
119
+ except ValueError:
120
+ position = text.find("{", position + 1)
121
+ continue
122
+ return document
123
+ snippet = text.strip()
124
+ if len(snippet) > 300:
125
+ snippet = snippet[:300] + "..."
126
+ raise ValueError(
127
+ f"reply contains no complete top-level JSON object "
128
+ f"({len(text)} characters); starts with: {snippet!r}"
129
+ )
130
+
131
+
132
+ def _check_exact_keys(
133
+ mapping: Dict[str, Any], expected: Tuple[str, ...], path: str, errors: List[str]
134
+ ) -> None:
135
+ missing = [key for key in expected if key not in mapping]
136
+ unexpected = [key for key in mapping if key not in expected]
137
+ if missing:
138
+ errors.append(f"{path}: missing required key(s): {', '.join(missing)}")
139
+ if unexpected:
140
+ errors.append(
141
+ f"{path}: unexpected key(s): {', '.join(unexpected)} "
142
+ f"(exactly {', '.join(expected)} are required)"
143
+ )
144
+
145
+
146
+ def _check_non_empty_string(value: Any, path: str, errors: List[str]) -> None:
147
+ if not isinstance(value, str):
148
+ errors.append(f"{path}: expected a string, got {type(value).__name__}")
149
+ elif not value.strip():
150
+ errors.append(f"{path}: string is empty")
151
+
152
+
153
+ def _check_coordinates(value: Any, path: str, errors: List[str]) -> None:
154
+ if not isinstance(value, str):
155
+ errors.append(
156
+ f"{path}: must be ONE string of the form "
157
+ f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000', got {type(value).__name__}"
158
+ )
159
+ return
160
+ match = COORDINATES_RE.match(value)
161
+ if match is None:
162
+ errors.append(
163
+ f"{path}: {value!r} is not of the form "
164
+ f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000'"
165
+ )
166
+ return
167
+ for field in COORDINATE_FIELDS:
168
+ number = float(match.group(field))
169
+ if not 0.0 <= number <= 1.0:
170
+ errors.append(f"{path}: {field}={match.group(field)} is outside [0, 1]")
171
+
172
+
173
+ def _check_color_specs(value: Any, path: str, errors: List[str]) -> None:
174
+ if not isinstance(value, list):
175
+ errors.append(
176
+ f"{path}: expected a list of hex colors, got {type(value).__name__}"
177
+ )
178
+ return
179
+ for index, color in enumerate(value):
180
+ if not isinstance(color, str) or HEX_COLOR_RE.match(color) is None:
181
+ errors.append(
182
+ f"{path}[{index}]: {color!r} is not a hex color "
183
+ f"(expected #RGB, #RGBA, #RRGGBB, or #RRGGBBAA)"
184
+ )
185
+
186
+
187
+ def validate_enhanced_prompt(document: Any) -> List[str]:
188
+ """Return schema errors for a rewritten prompt; an empty list means valid.
189
+
190
+ Schema demanded by assets/t2i_rewriter_system_prompt.txt: exactly two
191
+ top-level keys ``canvas_settings`` (exactly ``aspect_ratio``,
192
+ ``ambient_lighting``, ``image_style``) and ``layers`` (each layer exactly
193
+ ``description``, ``coordinates``, ``hierarchy_and_relation``,
194
+ ``color_specs``); ``coordinates`` is a string "cx: 0.500, cy: 0.500,
195
+ w: 1.000, h: 1.000" with values in [0, 1]; ``color_specs`` is a list of
196
+ hex colors. ``layers`` must hold at least one visible layer -- an empty
197
+ list means the rewrite failed even though it is type-correct.
198
+ """
199
+ if not isinstance(document, dict):
200
+ return [f"top level: expected a JSON object, got {type(document).__name__}"]
201
+ errors: List[str] = []
202
+ _check_exact_keys(document, ("canvas_settings", "layers"), "top level", errors)
203
+
204
+ if "canvas_settings" in document:
205
+ canvas = document["canvas_settings"]
206
+ if not isinstance(canvas, dict):
207
+ errors.append(
208
+ f"canvas_settings: expected a JSON object, got {type(canvas).__name__}"
209
+ )
210
+ else:
211
+ _check_exact_keys(canvas, CANVAS_SETTINGS_KEYS, "canvas_settings", errors)
212
+ for key in CANVAS_SETTINGS_KEYS:
213
+ if key in canvas:
214
+ _check_non_empty_string(
215
+ canvas[key], f"canvas_settings.{key}", errors
216
+ )
217
+
218
+ if "layers" in document:
219
+ layers = document["layers"]
220
+ if not isinstance(layers, list):
221
+ errors.append(f"layers: expected a list, got {type(layers).__name__}")
222
+ elif not layers:
223
+ errors.append("layers: expected at least one visible layer")
224
+ else:
225
+ for index, layer in enumerate(layers):
226
+ path = f"layers[{index}]"
227
+ if not isinstance(layer, dict):
228
+ errors.append(
229
+ f"{path}: expected a JSON object, got {type(layer).__name__}"
230
+ )
231
+ continue
232
+ _check_exact_keys(layer, LAYER_KEYS, path, errors)
233
+ for key in ("description", "hierarchy_and_relation"):
234
+ if key in layer:
235
+ _check_non_empty_string(layer[key], f"{path}.{key}", errors)
236
+ if "coordinates" in layer:
237
+ _check_coordinates(
238
+ layer["coordinates"], f"{path}.coordinates", errors
239
+ )
240
+ if "color_specs" in layer:
241
+ _check_color_specs(layer["color_specs"], f"{path}.color_specs", errors)
242
+ return errors
243
+
244
+
245
+ def _chat_completion(
246
+ base_url: str,
247
+ model: str,
248
+ messages: List[Dict[str, str]],
249
+ *,
250
+ temperature: float,
251
+ max_tokens: int,
252
+ api_key: Optional[str],
253
+ timeout: float,
254
+ ) -> Tuple[str, Optional[str]]:
255
+ """POST one chat completion; return (content, finish_reason)."""
256
+ url = base_url.rstrip("/") + "/chat/completions"
257
+ payload = json.dumps(
258
+ {
259
+ "model": model,
260
+ "messages": messages,
261
+ "temperature": temperature,
262
+ "max_tokens": max_tokens,
263
+ "stream": False,
264
+ }
265
+ ).encode("utf-8")
266
+ headers = {"Content-Type": "application/json"}
267
+ if api_key:
268
+ headers["Authorization"] = f"Bearer {api_key}"
269
+ request = urllib.request.Request(url, data=payload, headers=headers, method="POST")
270
+ try:
271
+ with urllib.request.urlopen(request, timeout=timeout) as response:
272
+ body = response.read().decode("utf-8", errors="replace")
273
+ except urllib.error.HTTPError as error:
274
+ detail = error.read().decode("utf-8", errors="replace")
275
+ raise PromptEnhancementError(
276
+ f"HTTP {error.code} from {url}: {detail[:2000]}"
277
+ ) from error
278
+ except urllib.error.URLError as error:
279
+ raise PromptEnhancementError(f"cannot reach {url}: {error.reason}") from error
280
+ except OSError as error: # includes socket timeouts during the read
281
+ raise PromptEnhancementError(f"request to {url} failed: {error}") from error
282
+
283
+ try:
284
+ envelope = json.loads(body)
285
+ choice = envelope["choices"][0]
286
+ content = choice["message"]["content"]
287
+ except (json.JSONDecodeError, KeyError, IndexError, TypeError) as error:
288
+ raise PromptEnhancementError(
289
+ f"malformed chat completion response from {url}: {body[:500]}"
290
+ ) from error
291
+ finish_reason = choice.get("finish_reason")
292
+ if not isinstance(content, str) or not content.strip():
293
+ raise PromptEnhancementError(
294
+ f"empty completion content from {url} (finish_reason={finish_reason!r})"
295
+ )
296
+ return content, finish_reason
297
+
298
+
299
+ def enhance(
300
+ caption: str,
301
+ base_url: str,
302
+ model: str,
303
+ api_key: Optional[str] = None,
304
+ timeout: float = DEFAULT_TIMEOUT_SECONDS,
305
+ temperature: float = DEFAULT_TEMPERATURE,
306
+ max_tokens: int = DEFAULT_MAX_TOKENS,
307
+ ) -> Dict[str, Any]:
308
+ """Return the validated structured rewrite of ``caption``.
309
+
310
+ Sends the verbatim rewriter system prompt plus the caption to
311
+ ``{base_url}/chat/completions``. On a parse or schema failure, retries
312
+ exactly once with the validation errors appended to the user turn; if
313
+ that also fails, raises PromptEnhancementError carrying the errors.
314
+ """
315
+ system_prompt = load_system_prompt()
316
+ messages = [
317
+ {"role": "system", "content": system_prompt},
318
+ {"role": "user", "content": caption},
319
+ ]
320
+ request_kwargs = {
321
+ "temperature": temperature,
322
+ "max_tokens": max_tokens,
323
+ "api_key": api_key,
324
+ "timeout": timeout,
325
+ }
326
+ errors: List[str] = []
327
+ content = ""
328
+ for attempt in (1, 2):
329
+ content, finish_reason = _chat_completion(
330
+ base_url, model, messages, **request_kwargs
331
+ )
332
+ document: Optional[Dict[str, Any]] = None
333
+ try:
334
+ document = extract_json_object(content)
335
+ except ValueError as error:
336
+ errors = [str(error)]
337
+ if document is not None:
338
+ errors = validate_enhanced_prompt(document)
339
+ if not errors:
340
+ assert document is not None # errors empty implies extraction succeeded
341
+ return document
342
+ if finish_reason == "length":
343
+ errors.append(
344
+ "the reply was cut off (finish_reason='length'): the complete "
345
+ f"JSON object must fit within max_tokens={max_tokens}"
346
+ )
347
+ print(f"pe_ling: attempt {attempt}/2 failed validation:", file=sys.stderr)
348
+ for error in errors:
349
+ print(f"pe_ling: - {error}", file=sys.stderr)
350
+ if attempt == 1:
351
+ retry_content = (
352
+ f"{caption}\n\n"
353
+ "Your previous reply failed schema validation:\n"
354
+ + "".join(f"- {error}\n" for error in errors)
355
+ + "\n"
356
+ + REPAIR_INSTRUCTION
357
+ )
358
+ messages = [
359
+ {"role": "system", "content": system_prompt},
360
+ {"role": "user", "content": retry_content},
361
+ ]
362
+ raise PromptEnhancementError(
363
+ "prompt enhancement failed schema validation after 2 attempts:\n"
364
+ + "".join(f" - {error}\n" for error in errors).rstrip(),
365
+ errors=errors,
366
+ reply=content,
367
+ )
368
+
369
+
370
+ def main() -> None:
371
+ parser = argparse.ArgumentParser(
372
+ description=(
373
+ "Enhance a Ming-Image text-to-image caption into the validated "
374
+ "structured JSON prompt via an OpenAI-compatible Ling-3.0-flash-VL "
375
+ "endpoint."
376
+ )
377
+ )
378
+ parser.add_argument("caption", help="free-form design caption to enhance")
379
+ parser.add_argument(
380
+ "--out",
381
+ type=Path,
382
+ help="write the validated JSON here (default: stdout, summary on stderr)",
383
+ )
384
+ parser.add_argument(
385
+ "--base-url",
386
+ default=DEFAULT_BASE_URL,
387
+ help=f"OpenAI-compatible base URL (default: {DEFAULT_BASE_URL})",
388
+ )
389
+ parser.add_argument(
390
+ "--model",
391
+ default=DEFAULT_MODEL,
392
+ help=f"chat model id served at the endpoint (default: {DEFAULT_MODEL})",
393
+ )
394
+ parser.add_argument(
395
+ "--api-key",
396
+ default=os.environ.get(API_KEY_ENV),
397
+ help=f"Bearer token for gated endpoints; defaults to ${API_KEY_ENV} when set",
398
+ )
399
+ parser.add_argument(
400
+ "--timeout",
401
+ type=float,
402
+ default=DEFAULT_TIMEOUT_SECONDS,
403
+ help=f"per-request timeout in seconds (default: {DEFAULT_TIMEOUT_SECONDS})",
404
+ )
405
+ parser.add_argument(
406
+ "--temperature",
407
+ type=float,
408
+ default=DEFAULT_TEMPERATURE,
409
+ help=f"sampling temperature (default: {DEFAULT_TEMPERATURE})",
410
+ )
411
+ parser.add_argument(
412
+ "--max-tokens",
413
+ type=int,
414
+ default=DEFAULT_MAX_TOKENS,
415
+ help=f"completion token budget (default: {DEFAULT_MAX_TOKENS})",
416
+ )
417
+ args = parser.parse_args()
418
+
419
+ started = time.perf_counter()
420
+ try:
421
+ document = enhance(
422
+ args.caption,
423
+ args.base_url,
424
+ args.model,
425
+ api_key=args.api_key,
426
+ timeout=args.timeout,
427
+ temperature=args.temperature,
428
+ max_tokens=args.max_tokens,
429
+ )
430
+ except PromptEnhancementError as error:
431
+ print(f"pe_ling: {error}", file=sys.stderr)
432
+ raise SystemExit(1)
433
+ elapsed = time.perf_counter() - started
434
+ layer_count = len(document["layers"])
435
+ payload = json.dumps(document, indent=2, ensure_ascii=False) + "\n"
436
+ if args.out is not None:
437
+ args.out.parent.mkdir(parents=True, exist_ok=True)
438
+ args.out.write_text(payload, encoding="utf-8")
439
+ print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s) -> {args.out}")
440
+ else:
441
+ sys.stdout.write(payload)
442
+ print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s)", file=sys.stderr)
443
+
444
+
445
+ if __name__ == "__main__":
446
+ main()
code/processing_bailingmm2.py ADDED
@@ -0,0 +1,512 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2024 The HuggingFace Inc. team.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ """Processor class for BailingMM2."""
16
+
17
+ import sys
18
+ from typing import List, Union, Dict, Optional
19
+
20
+ import torch
21
+ import PIL
22
+ from PIL import Image
23
+
24
+ if sys.version_info >= (3, 11):
25
+ from typing import Unpack
26
+ else:
27
+ from typing_extensions import Unpack
28
+
29
+ from transformers.feature_extraction_utils import BatchFeature
30
+ from transformers.image_utils import ImageInput
31
+ from transformers.processing_utils import (
32
+ ProcessingKwargs,
33
+ ProcessorMixin,
34
+ )
35
+ from transformers.tokenization_utils_base import PreTokenizedInput, TextInput
36
+
37
+ from bailingmm_utils import process_vision_info, VideoInput, process_ratio, process_reference_vision_info, get_default_image_gen_hw
38
+ import torchvision
39
+ import math
40
+
41
+ DEFAULT_IMAGE_PATCH_TOKEN = "<imagePatch>"
42
+ DEFAULT_IM_START_TOKEN = "<image>"
43
+ DEFAULT_IM_END_TOKEN = "</image>"
44
+ DEFAULT_VID_START_TOKEN = "<video>"
45
+ DEFAULT_VID_END_TOKEN = "</video>"
46
+ DEFAULT_GEN_IMAGE_PATCH_TOKEN = "<gen_imagePatch>"
47
+ DEFAULT_GEN_IM_START_TOKEN = "<gen_image>"
48
+ DEFAULT_GEN_IM_END_TOKEN = "</gen_image>"
49
+ PLACEHOLDER_IMAGE_TOKEN_IN_TEXT = "<imageHere>"
50
+ DEFAULT_END_OF_CHUNK_TOKEN = "<end_of_chunk>"
51
+
52
+ DEFAULT_FRAME_PATCH_TOKEN = "<framePatch>"
53
+ DEFAULT_TEXT_TOKEN = '<text>'
54
+ DEFAULT_ASR_TOKEN = '<asr>'
55
+ DEFAULT_TTS_TOKEN = '<tts>'
56
+
57
+ USER_PREFIX = "<role>HUMAN</role>"
58
+ ASSISTANT_PREFIX = "<role>ASSISTANT</role>"
59
+
60
+ SYSTEM_PROMPT_LINGV2_FLASH_NOTHINK = "<role>SYSTEM</role>你是一个友好的AI助手。\n\ndetailed thinking off"
61
+ SYSTEM_PROMPT_LINGV2_FLASH_THINK = "<role>SYSTEM</role>你是一个友好的AI助手。\n\ndetailed thinking on"
62
+
63
+ def check_single_quotes(s):
64
+ count = s.count("'")
65
+ if count % 2 != 0:
66
+ return False
67
+
68
+ positions = [i for i, char in enumerate(s) if char == "'"]
69
+ for i in range(0, len(positions), 2):
70
+ start = positions[i]
71
+ end = positions[i+1]
72
+ substr = s[start+1:end]
73
+ chinese_count = 0
74
+ for char in substr:
75
+ if '\u4e00' <= char <= '\u9fff':
76
+ chinese_count += 1
77
+ other_count = len(substr) - chinese_count
78
+ total = 3 * chinese_count + other_count
79
+ if total >= 20:
80
+ return False
81
+ return True
82
+
83
+ def get_text_from_prompt(prompt):
84
+ if "'" in prompt and check_single_quotes(prompt):
85
+ prompt = prompt.replace("'", '"')
86
+
87
+ patterns = [r'\"(.*?)\"', r'‘(.*?)’', r'“(.*?)”']
88
+
89
+
90
+ import re
91
+ texts = []
92
+ patterns = [r'\"(.*?)\"', r'‘(.*?)’', r'“(.*?)”']
93
+ for pattern in patterns:
94
+ texts.extend(re.findall(pattern, prompt))
95
+
96
+ if len(texts) == 1:
97
+ assert texts[0] in prompt
98
+ is_remove = False
99
+ remove_keywords = ["remove", "delete", "erase"]
100
+ text_start = min([j for j in [prompt.find(i) for i in ['"', '‘', '“']] if j >= 0])
101
+ for kw in remove_keywords:
102
+ if kw in prompt.lower():
103
+ if prompt.lower().find(kw) < text_start:
104
+ is_remove = True
105
+ break
106
+
107
+ if is_remove:
108
+ texts = []
109
+
110
+ text = " ".join(texts[-1:])
111
+ if len(text) > 0:
112
+ text = f'Text "{text}"'
113
+ text += ". "
114
+
115
+ return text
116
+
117
+ def crop_to_aspect_max(img: Image.Image, target_ratio: float) -> Image.Image:
118
+ """
119
+ Center-crop a PIL.Image to the largest area fitting the target aspect ratio
120
+ (width/height) without resizing. Uses torchvision CenterCrop.
121
+
122
+ Args:
123
+ img: PIL.Image.Image input image
124
+ target_ratio: float target aspect ratio (width/height), must be positive
125
+
126
+ Returns:
127
+ the center-cropped PIL.Image.Image
128
+ """
129
+ if not isinstance(img, Image.Image):
130
+ raise TypeError("img must be a PIL.Image.Image")
131
+ if not math.isfinite(target_ratio) or target_ratio <= 0:
132
+ raise ValueError("target_ratio must be a positive, finite number")
133
+
134
+ W, H = img.size
135
+ if W <= 0 or H <= 0:
136
+ raise ValueError("image size is invalid")
137
+
138
+ orig_ratio = W / H
139
+
140
+ if orig_ratio >= target_ratio:
141
+ # image is wider than the target: use full height, crop left/right
142
+ new_h = H
143
+ new_w = int(math.floor(target_ratio * H))
144
+ new_w = max(1, min(new_w, W)) # guard against extreme ratios producing invalid sizes
145
+ else:
146
+ # image is narrower than the target: use full width, crop top/bottom
147
+ new_w = W
148
+ new_h = int(math.floor(W / target_ratio))
149
+ new_h = max(1, min(new_h, H))
150
+
151
+ crop = torchvision.transforms.CenterCrop((new_h, new_w)) # size is (h, w)
152
+ return crop(img)
153
+
154
+ def transform_reference_images(
155
+ images,
156
+ image_gen_aspect_ratio=None,
157
+ image_gen_resolution=512,
158
+ image_gen_input_channels=None,
159
+ ):
160
+ if image_gen_input_channels not in (3, 4):
161
+ raise ValueError(
162
+ "image_gen_input_channels must be explicitly set to 3 or 4 "
163
+ "from the checkpoint capability contract"
164
+ )
165
+ image_mode = "RGB" if image_gen_input_channels == 3 else "RGBA"
166
+ images = [image.convert(image_mode) for image in images]
167
+
168
+ ref_pil = images[0]
169
+ if image_gen_aspect_ratio is not None:
170
+ ref_pil = crop_to_aspect_max(ref_pil, image_gen_aspect_ratio)
171
+
172
+ ori_h = ref_pil.size[1]
173
+ ori_w = ref_pil.size[0]
174
+ closest_size, _ = process_ratio(ori_h=ori_h, ori_w=ori_w, highres=image_gen_resolution)
175
+
176
+ ref_pils = [torchvision.transforms.functional.resize(i, closest_size, interpolation=torchvision.transforms.InterpolationMode.BILINEAR) for i in images]
177
+
178
+ ref_tensor = torch.cat([
179
+ ((torchvision.transforms.functional.to_tensor(i) - 0.5) * 2.0).unsqueeze(0)
180
+ for i in ref_pils
181
+ ], dim=0)
182
+
183
+ return ref_tensor, ref_pil.size[1], ref_pil.size[0]
184
+
185
+ class BailingMM2ProcessorKwargs(ProcessingKwargs, total=False):
186
+ # see processing_utils.ProcessingKwargs documentation for usage.
187
+ _defaults = {
188
+ "text_kwargs": {"padding": True, "padding_side": "right"},
189
+ "image_kwargs": {},
190
+ "video_kwargs": {},
191
+ }
192
+
193
+ class BailingMM2Processor(ProcessorMixin):
194
+ r"""
195
+ Constructs a BailingMM2 processor which wraps a bailingmm2 image processor, bailing audio processor and a LLaMa tokenizer into a single processor.
196
+ Args:
197
+ image_processor ([`BailingMM2ImageProcessor`], *optional*):
198
+ The image processor is a required input.
199
+ tokenizer ([`LlamaTokenizerFast`], *optional*):
200
+ The tokenizer is a required input.
201
+ chat_template (`str`, *optional*): A Jinja template which will be used to convert lists of messages
202
+ in a chat into a tokenizable string.
203
+ image_token (`str`, *optional*, defaults to `"<image>"`):
204
+ Special token used to denote image location.
205
+ video_token (`str`, *optional*, defaults to `"<video>"`):
206
+ Special token used to denote video location.
207
+ """
208
+
209
+ attributes = ["image_processor", "tokenizer"]
210
+ optional_attributes = ["chat_template"]
211
+
212
+ image_processor_class = "AutoImageProcessor"
213
+ tokenizer_class = "AutoTokenizer"
214
+
215
+ valid_kwargs = [
216
+ "chat_template",
217
+ "num_image_tokens",
218
+ "image_token",
219
+ "video_token",
220
+ ]
221
+
222
+ def __init__(
223
+ self,
224
+ image_processor=None,
225
+ tokenizer=None,
226
+ chat_template=None,
227
+ image_token="<image>",
228
+ video_token="<video>",
229
+ **kwargs: Unpack[BailingMM2ProcessorKwargs],
230
+ ):
231
+ self.image_token = image_token
232
+ self.video_token = video_token
233
+ if chat_template is None:
234
+ chat_template = tokenizer.chat_template
235
+
236
+ self.gen_terminator = [tokenizer.eos_token_id]
237
+ super().__init__(image_processor, tokenizer, chat_template=chat_template)
238
+
239
+ def __call__(
240
+ self,
241
+ images: ImageInput = None,
242
+ videos: VideoInput = None,
243
+ text: Union[TextInput, PreTokenizedInput, List[TextInput], List[PreTokenizedInput]] = None,
244
+ image_gen_highres = 512,
245
+ image_gen_aspect_ratio = None,
246
+ image_gen_ref_images: Union["PIL.Image.Image", list["PIL.Image.Image"]] = None,
247
+ image_gen_input_channels = None,
248
+ **kwargs,
249
+ ) -> BatchFeature:
250
+ """
251
+ Main method to prepare for the model one or several sequences(s) and image(s). This method forwards the `text`
252
+ and `kwargs` arguments to LlamaTokenizerFast's [`~LlamaTokenizerFast.__call__`] if `text` is not `None` to encode
253
+ the text. To prepare the image(s), this method forwards the `images` and `kwrags` arguments to
254
+ LlavaNextImageProcessor's [`~LlavaNextImageProcessor.__call__`] if `images` is not `None`. Please refer to the doctsring
255
+ of the above two methods for more information.
256
+
257
+ Args:
258
+ images (`PIL.Image.Image`, `np.ndarray`, `torch.Tensor`, `List[PIL.Image.Image]`, `List[np.ndarray]`, `List[torch.Tensor]`):
259
+ The image or batch of images to be prepared. Each image can be a PIL image, NumPy array or torch Tensor.
260
+ tensor. Both channels-first and channels-last formats are supported.
261
+ videos (`np.ndarray`, `torch.Tensor`, `List[np.ndarray]`, `List[torch.Tensor]`):
262
+ The image or batch of videos to be prepared. Each video can be a 4D NumPy array or torch Tensor.
263
+ audios (`Tuple[torch.Tensor, int]`, `List[Tuple[torch.Tensor, int]]`):
264
+ The sequence or batch of audios to be prepared. Each audio can be a 1D torch Tensor (with its sampling rate).
265
+ text (`str`, `List[str]`, `List[List[str]]`):
266
+ The sequence or batch of sequences to be encoded. Each sequence can be a string or a list of strings
267
+ (pretokenized string). If the sequences are provided as a list of strings (pretokenized), you must set
268
+ `is_split_into_words=True` (to lift the ambiguity with a batch of sequences).
269
+
270
+ Returns:
271
+ [`BatchFeature`]: A [`BatchFeature`] with the following fields:
272
+
273
+ - **input_ids** -- List of token ids to be fed to a model. Returned when `text` is not `None`.
274
+ - **attention_mask** -- List of indices specifying which tokens should be attended to by the model (when
275
+ `return_attention_mask=True` or if *"attention_mask"* is in `self.model_input_names` and if `text` is not
276
+ `None`).
277
+ - **pixel_values** -- Pixel values to be fed to a model. Returned when `images` is not `None`.
278
+ - **image_num_patches** -- Patch number to be fed to a model. Returned when `images` is not `None`.
279
+ - **image_sizes** -- Size of each image that will be used to unpad an image. Returned when `images` is not `None`.
280
+ - **pixel_values_videos** -- Pixel values of a video input to be fed to a model. Returned when `videos` is not `None`.
281
+ - **pixel_values_audios** -- Pixel values of an audio input to be fed to a model. Returned when `audios` is not `None`.
282
+
283
+ """
284
+ output_kwargs = self._merge_kwargs(
285
+ BailingMM2ProcessorKwargs,
286
+ tokenizer_init_kwargs=self.tokenizer.init_kwargs,
287
+ **kwargs,
288
+ )
289
+ if isinstance(text, str):
290
+ text = [text]
291
+ elif not isinstance(text, list) and not isinstance(text[0], str):
292
+ raise ValueError("Invalid input text. Please provide a string, or a list of strings")
293
+
294
+ image_inputs = {}
295
+ video_inputs = {}
296
+ image_gen_inputs = {}
297
+
298
+ text_in_text = [get_text_from_prompt(i) for i in text]
299
+
300
+
301
+ default_image_gen_height, default_image_gen_width = get_default_image_gen_hw(image_gen_highres, image_gen_aspect_ratio)
302
+
303
+ image_gen_inputs.update({
304
+ "image_gen_text": text_in_text,
305
+ "image_gen_highres": image_gen_highres,
306
+ "image_gen_height": torch.LongTensor([default_image_gen_height] * len(text)),
307
+ "image_gen_width": torch.LongTensor([default_image_gen_width] * len(text)),
308
+ })
309
+
310
+ if images is not None:
311
+ image_inputs = self.image_processor(images=images, videos=None, **output_kwargs["images_kwargs"])
312
+ image_grid_thw = image_inputs["image_grid_thw"]
313
+
314
+ text = self._expand_image_tokens(text, image_grid_thw)
315
+
316
+ # image_gen_pixel_values_reference, image_gen_height, image_gen_width = None, 512, 512
317
+ if image_gen_ref_images is not None:
318
+ if isinstance(image_gen_ref_images, PIL.Image.Image):
319
+ image_gen_ref_images = [image_gen_ref_images]
320
+ elif not isinstance(image_gen_ref_images, list) and not isinstance(image_gen_ref_images[0], PIL.Image.Image):
321
+ raise ValueError("Invalid input image_gen_ref_images. Please provide a PIL.Image.Image, or a list of PIL.Image.Image")
322
+
323
+ assert len(image_gen_ref_images) == len(text) # same batch_size
324
+
325
+ image_gen_pixel_values_reference, image_gen_height_list, image_gen_width_list = transform_reference_images(
326
+ image_gen_ref_images,
327
+ image_gen_aspect_ratio,
328
+ image_gen_highres,
329
+ image_gen_input_channels,
330
+ )
331
+
332
+ image_gen_inputs.update({
333
+ "image_gen_pixel_values_reference": image_gen_pixel_values_reference,
334
+ "image_gen_height": torch.LongTensor([image_gen_height_list] * len(text)),
335
+ "image_gen_width": torch.LongTensor([image_gen_width_list] * len(text)),
336
+ #"image_gen_height": torch.LongTensor([ori_h]),
337
+ #"image_gen_width": torch.LongTensor([ori_w]),
338
+ })
339
+
340
+ if videos is not None:
341
+ video_inputs = self.image_processor(images=None, videos=videos, do_resize=False, **output_kwargs["videos_kwargs"])
342
+ video_grid_thw = video_inputs["video_grid_thw"]
343
+ text = self._expand_video_tokens(text, video_grid_thw)
344
+
345
+ # Padding side can be in TextKwargs but is not accepted by the tokenizer
346
+ _ = output_kwargs["text_kwargs"].pop("padding_side", None)
347
+ text_inputs = self.tokenizer(text, **output_kwargs["text_kwargs"])
348
+
349
+ return BatchFeature(data={**text_inputs, **image_inputs, **video_inputs, **image_gen_inputs})
350
+
351
+ def apply_system_template(self, sys_prompt_exp=None, use_cot_system_prompt=False):
352
+ if use_cot_system_prompt:
353
+ sys_prompt = SYSTEM_PROMPT_LINGV2_FLASH_THINK
354
+ else:
355
+ sys_prompt = SYSTEM_PROMPT_LINGV2_FLASH_NOTHINK
356
+ if sys_prompt_exp is not None:
357
+ sys_prompt = sys_prompt.replace("你是一个友好的AI助手。", sys_prompt_exp)
358
+
359
+ return sys_prompt
360
+
361
+ def apply_chat_template(
362
+ self,
363
+ conversation: Union[List[Dict[str, str]]],
364
+ sys_prompt_exp: Optional[str] = None,
365
+ use_cot_system_prompt: Optional[bool] = False,
366
+ **kwargs,
367
+ ) -> str:
368
+ """
369
+ Similar to the `apply_chat_template` method on tokenizers, this method applies a Jinja template to input
370
+ conversations to turn them into a single tokenizable string.
371
+
372
+ Args:
373
+ conversation (`List[Dict, str, str]`):
374
+ The conversation to format.
375
+ sys_prompt_exp (`Optional[str]`, *optional*):
376
+ The system prompt. If not provided, the processor's sysyetm template is used.
377
+ **kwargs:
378
+ Additional keyword arguments
379
+ """
380
+ text = ""
381
+ sys_prompt = self.apply_system_template(sys_prompt_exp, use_cot_system_prompt)
382
+ text = sys_prompt + self.tokenizer.eos_token
383
+
384
+ for idx, message in enumerate(conversation):
385
+ assert message["role"] in ["HUMAN", "ASSISTANT"]
386
+ if idx == len(conversation) - 1:
387
+ assert message["role"] == "HUMAN"
388
+
389
+ if message["role"] == "HUMAN":
390
+ text += USER_PREFIX
391
+ elif message["role"] == "ASSISTANT":
392
+ text += ASSISTANT_PREFIX
393
+
394
+ image_counts = str(message["content"]).count("<image>")
395
+ video_counts = str(message["content"]).count("<video>")
396
+
397
+ for content in message["content"]:
398
+ if content["type"] == "image":
399
+ num_images = 1 if isinstance(content["image"], (str, Image.Image)) else len(content["image"])
400
+ if image_counts < num_images:
401
+ image_placeholder = "<IMAGE>\n" * (num_images - image_counts)
402
+ text += image_placeholder.rstrip("\n")
403
+ # only one video supported now
404
+ elif content["type"] == "video":
405
+ assert video_counts <= 1, "Video count must be at most 1!"
406
+ if video_counts == 0:
407
+ text += "<VIDEO>"
408
+ elif content["type"] == "audio":
409
+ raise ValueError("audio input is not supported by Ming Image inference")
410
+ elif content["type"] == "text":
411
+ text += content['text']
412
+ text += self.tokenizer.eos_token
413
+ text += ASSISTANT_PREFIX
414
+
415
+ return text
416
+
417
+ def process_vision_info(
418
+ self,
419
+ conversations,
420
+ ):
421
+ return process_vision_info(conversations)
422
+
423
+ def process_reference_vision_info(
424
+ self,
425
+ conversations,
426
+ ):
427
+ return process_reference_vision_info(conversations)
428
+
429
+ def _expand_image_tokens(
430
+ self,
431
+ text: List[TextInput],
432
+ image_grid_thw: Union[List[int], int],
433
+ special_token: str = "<IMAGE>",
434
+ ):
435
+ prompt_strings = []
436
+ image_index = 0
437
+ num_query_token = torch.prod(image_grid_thw, dim=1) // 4
438
+ for sample in text:
439
+ num_images = sample.count(special_token)
440
+ if num_images > 0:
441
+ for i in range(image_index, num_images + image_index):
442
+ img_text = DEFAULT_IM_START_TOKEN + num_query_token[i] * DEFAULT_IMAGE_PATCH_TOKEN + DEFAULT_IM_END_TOKEN + "\n"
443
+ sample = sample.replace(special_token, img_text, 1)
444
+ image_index += num_images
445
+ prompt_strings.append(sample)
446
+ text = [sample for sample in prompt_strings]
447
+ return text
448
+
449
+ def _expand_video_tokens(
450
+ self,
451
+ text: List[TextInput],
452
+ video_grid_thw: Union[List[int], int],
453
+ special_token: str = "<VIDEO>",
454
+ ):
455
+ prompt_strings = []
456
+ video_index = 0
457
+ num_query_token = torch.prod(video_grid_thw, dim=1) // 4
458
+ for sample in text:
459
+ num_videos = sample.count(special_token)
460
+ if num_videos > 0:
461
+ for i in range(video_index, num_videos + video_index):
462
+ video_text = num_query_token[i] * DEFAULT_FRAME_PATCH_TOKEN
463
+ video_text = DEFAULT_VID_START_TOKEN + video_text + DEFAULT_VID_END_TOKEN + "\n"
464
+ sample = sample.replace(special_token, video_text, 1)
465
+ video_index += num_videos
466
+ prompt_strings.append(sample)
467
+ text = [sample for sample in prompt_strings]
468
+ return text
469
+
470
+ # Copied from transformers.models.clip.processing_clip.CLIPProcessor.batch_decode with CLIP->Llama
471
+ def batch_decode(self, *args, **kwargs):
472
+ """
473
+ This method forwards all its arguments to LlamaTokenizerFast's [`~PreTrainedTokenizer.batch_decode`]. Please
474
+ refer to the docstring of this method for more information.
475
+ """
476
+ return self.tokenizer.batch_decode(*args, **kwargs)
477
+
478
+ # Copied from transformers.models.clip.processing_clip.CLIPProcessor.decode with CLIP->Llama
479
+ def decode(self, *args, **kwargs):
480
+ """
481
+ This method forwards all its arguments to LlamaTokenizerFast's [`~PreTrainedTokenizer.decode`]. Please refer to
482
+ the docstring of this method for more information.
483
+ """
484
+ return self.tokenizer.decode(*args, **kwargs)
485
+
486
+ @property
487
+ def model_input_names(self):
488
+ tokenizer_input_names = self.tokenizer.model_input_names
489
+ image_processor_input_names = self.image_processor.model_input_names
490
+ return list(
491
+ dict.fromkeys(
492
+ tokenizer_input_names + image_processor_input_names))
493
+
494
+
495
+ def load_bailingmm2_processor(data_directory):
496
+ """Build the processor from data files with this repository's classes.
497
+
498
+ Checkpoint component directories (``mllm/``) carry only data files
499
+ (``preprocessor_config.json``, ``tokenizer_config.json``,
500
+ ``special_tokens_map.json``, ``tokenizer.json``). Loading them through
501
+ ``AutoProcessor`` with ``trust_remote_code=True`` fails because
502
+ Transformers requires the Python implementation inside the loaded
503
+ directory, while the implementation intentionally lives only in this
504
+ repository. Construct the components explicitly instead.
505
+ """
506
+ from image_processing_bailingmm2 import BailingMM2ImageProcessor
507
+ from tokenization_bailing import BailingTokenizer
508
+
509
+ data_directory = str(data_directory)
510
+ tokenizer = BailingTokenizer.from_pretrained(data_directory)
511
+ image_processor = BailingMM2ImageProcessor.from_pretrained(data_directory)
512
+ return BailingMM2Processor(image_processor=image_processor, tokenizer=tokenizer)
code/quant/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Weight-only INT8 for the Ming-Image MLLM: quantize_stream.py writes it, load_int8.py loads it."""
code/quant/int8_linear.py ADDED
@@ -0,0 +1,221 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Weight-only symmetric per-output-channel INT8 linear.
2
+
3
+ Scales stay float32 across dtype casts. ``module.to(dtype=torch.bfloat16)``
4
+ (and ``.bfloat16()`` / ``.half()`` / ``.to(device, dtype)``) must not touch them;
5
+ device moves still do. The int8 weight codes are likewise dtype-stable.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import torch
11
+ import torch.nn.functional as F
12
+ from torch import nn
13
+
14
+ # Leaf names of Linear modules whose 2-D weights are quantized.
15
+ # Exact match: the routers are `gate` / `image_gate` / `audio_gate`, NOT `gate_proj`.
16
+ QUANT_LEAVES = frozenset(
17
+ {"query_key_value", "dense", "gate_proj", "up_proj", "down_proj"}
18
+ )
19
+
20
+ QUANT_RULE = (
21
+ "Quantize ONLY 2-D .weight tensors under model.model.layers. whose owning "
22
+ "module's leaf name is exactly one of query_key_value, dense, gate_proj, "
23
+ "up_proj, down_proj. Everything else stays byte-identical BF16: embeddings, "
24
+ "lm_head, all norms, the vision tower, linear_proj, and the three routers "
25
+ "(modules named gate, image_gate, audio_gate — leaf match, not a substring). "
26
+ "Per-output-channel symmetric: scale = absmax/127, "
27
+ "q = clamp(round(w/scale), -127, 127). All-zero rows: scale 1.0, q 0."
28
+ )
29
+
30
+ # Real checkpoint keys look like `model.model.layers.N...`. A bare
31
+ # `layers.N...` name is the same stack with the root prefix omitted (tests).
32
+ _DECODER_LAYER_PREFIXES = ((), ("model", "model"))
33
+
34
+
35
+ def _weight_leaf(tensor_name: str) -> str | None:
36
+ """Owning module's leaf name if `tensor_name` ends in `.weight`, else None."""
37
+ if not isinstance(tensor_name, str) or not tensor_name.endswith(".weight"):
38
+ return None
39
+ module = tensor_name[: -len(".weight")]
40
+ if not module:
41
+ return None
42
+ return module.rsplit(".", 1)[-1]
43
+
44
+
45
+ def _under_decoder_layers(tensor_name: str) -> bool:
46
+ """True when the tensor lives under the MLLM decoder `model.model.layers` stack.
47
+
48
+ `layers` must be its own path component, followed by a layer index. The
49
+ components before it must be empty or end in `model.model` — so a vision
50
+ tower that happens to contain the substring "layers" is not selected, and
51
+ `gate` is never selected just because `gate_proj` contains those letters.
52
+ """
53
+ parts = tensor_name.split(".")
54
+ for i, part in enumerate(parts):
55
+ if part != "layers":
56
+ continue
57
+ if i + 1 >= len(parts) or not parts[i + 1].isdigit():
58
+ continue
59
+ prefix = tuple(parts[:i])
60
+ if prefix in _DECODER_LAYER_PREFIXES:
61
+ return True
62
+ if len(prefix) >= 2 and prefix[-2:] == ("model", "model"):
63
+ return True
64
+ return False
65
+
66
+
67
+ def quant_rule_leaf(tensor_name: str) -> str | None:
68
+ """Leaf name if the name matches the quantize rule, ignoring rank.
69
+
70
+ Returns None when the tensor is not a candidate. A candidate whose rank is
71
+ not 2 is a hard error for the stream (see quantize_stream); ``is_quantizable``
72
+ itself returns False for that case.
73
+ """
74
+ leaf = _weight_leaf(tensor_name)
75
+ if leaf not in QUANT_LEAVES:
76
+ return None
77
+ if not _under_decoder_layers(tensor_name):
78
+ return None
79
+ return leaf
80
+
81
+
82
+ def is_quantizable(tensor_name: str, shape) -> bool:
83
+ """True only for 2-D quantize-rule weights. See ``QUANT_RULE``."""
84
+ if quant_rule_leaf(tensor_name) is None:
85
+ return False
86
+ try:
87
+ rank = len(shape)
88
+ except TypeError:
89
+ return False
90
+ return rank == 2
91
+
92
+
93
+ def quantize_weight(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
94
+ """Per-output-channel symmetric int8.
95
+
96
+ ``scale = absmax(row) / 127``, ``q = clamp(round(w / scale), -127, 127)``.
97
+ An all-zero row gets scale 1.0 and q 0 (no div-by-zero, no NaN/Inf).
98
+ """
99
+ if weight.ndim != 2:
100
+ raise ValueError(
101
+ f"quantize_weight expects a 2-D weight, got shape {tuple(weight.shape)}"
102
+ )
103
+ wf = weight.detach().to(dtype=torch.float32)
104
+ absmax = wf.abs().amax(dim=1)
105
+ scale = absmax / 127.0
106
+ zero = scale == 0
107
+ # All-zero rows would divide by 0. Force scale 1 and q 0 instead of NaN.
108
+ scale = torch.where(zero, torch.ones_like(scale), scale)
109
+ q = torch.round(wf / scale[:, None]).clamp(-127, 127).to(dtype=torch.int8)
110
+ q = torch.where(zero[:, None], torch.zeros_like(q), q)
111
+ return q.contiguous(), scale.to(dtype=torch.float32).contiguous()
112
+
113
+
114
+ def _scale_name(weight_name: str) -> str:
115
+ if not weight_name.endswith(".weight"):
116
+ raise ValueError(f"not a weight tensor name: {weight_name}")
117
+ return weight_name[: -len("weight")] + "scale"
118
+
119
+
120
+ class Int8Linear(nn.Module):
121
+ """``F.linear`` on a weight dequantized from int8 + per-row float32 scale.
122
+
123
+ ``weight`` is int8 ``[out, in]``, ``scale`` is float32 ``[out]``, ``bias``
124
+ (optional) keeps the source dtype. All three are buffers.
125
+ """
126
+
127
+ def __init__(self, weight: torch.Tensor, scale: torch.Tensor, bias: torch.Tensor | None):
128
+ super().__init__()
129
+ if weight.dtype != torch.int8 or weight.ndim != 2:
130
+ raise ValueError(
131
+ f"weight must be int8 [out, in], got dtype={weight.dtype} shape={tuple(weight.shape)}"
132
+ )
133
+ if scale.dtype != torch.float32 or tuple(scale.shape) != (weight.shape[0],):
134
+ raise ValueError(
135
+ f"scale must be float32 [{weight.shape[0]}], got dtype={scale.dtype} shape={tuple(scale.shape)}"
136
+ )
137
+ if bias is not None:
138
+ if bias.ndim != 1 or bias.shape[0] != weight.shape[0]:
139
+ raise ValueError(
140
+ f"bias must be [{weight.shape[0]}], got shape={tuple(bias.shape)}"
141
+ )
142
+ self.in_features = int(weight.shape[1])
143
+ self.out_features = int(weight.shape[0])
144
+ self.register_buffer("weight", weight)
145
+ self.register_buffer("scale", scale)
146
+ self.register_buffer("bias", bias)
147
+
148
+ def _apply(self, fn, *args, **kwargs):
149
+ # Pull dtype-stable buffers out before Module._apply. Putting them back
150
+ # with only a device move (never fn's dtype cast) keeps scale float32
151
+ # and weight int8. Bias is left in the dict so it follows the cast.
152
+ saved: dict[str, torch.Tensor] = {}
153
+ for name in ("weight", "scale"):
154
+ buf = self._buffers.get(name, None)
155
+ if buf is not None:
156
+ saved[name] = buf
157
+ self._buffers[name] = None
158
+ try:
159
+ out = super()._apply(fn, *args, **kwargs)
160
+ finally:
161
+ for name, buf in saved.items():
162
+ self._buffers[name] = _move_device_keep_dtype(buf, fn)
163
+ return out
164
+
165
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
166
+ # One dequant in fp32, one cast to the activation dtype, then linear.
167
+ w = (self.weight.float() * self.scale[:, None]).to(dtype=x.dtype)
168
+ return F.linear(x, w, self.bias)
169
+
170
+ @classmethod
171
+ def from_linear(cls, linear: nn.Linear) -> "Int8Linear":
172
+ if not isinstance(linear, nn.Linear):
173
+ raise TypeError(f"from_linear expects nn.Linear, got {type(linear).__name__}")
174
+ q, scale = quantize_weight(linear.weight.data)
175
+ if linear.bias is None:
176
+ bias = None
177
+ else:
178
+ bias = linear.bias.detach().clone()
179
+ return cls(q, scale, bias)
180
+
181
+ @classmethod
182
+ def shell(
183
+ cls,
184
+ in_features: int,
185
+ out_features: int,
186
+ bias: bool,
187
+ bias_dtype: torch.dtype,
188
+ device,
189
+ ) -> "Int8Linear":
190
+ """Empty buffers (for ``meta``). Does not read or write any weight values."""
191
+ dev = torch.device(device) if not isinstance(device, torch.device) else device
192
+ weight = torch.empty((out_features, in_features), dtype=torch.int8, device=dev)
193
+ scale = torch.empty((out_features,), dtype=torch.float32, device=dev)
194
+ if bias:
195
+ bias_t: torch.Tensor | None = torch.empty(
196
+ (out_features,), dtype=bias_dtype, device=dev
197
+ )
198
+ else:
199
+ bias_t = None
200
+ return cls(weight, scale, bias_t)
201
+
202
+ def extra_repr(self) -> str:
203
+ return (
204
+ f"in_features={self.in_features}, out_features={self.out_features}, "
205
+ f"bias={self.bias is not None}"
206
+ )
207
+
208
+
209
+ def _move_device_keep_dtype(buf: torch.Tensor, fn) -> torch.Tensor:
210
+ """Apply only the device change implied by ``fn``, preserving ``buf``'s dtype and values.
211
+
212
+ Probed with a 0-element tensor so a dtype cast cannot round the real scale.
213
+ """
214
+ try:
215
+ probe = torch.empty((), dtype=buf.dtype, device=buf.device)
216
+ moved = fn(probe)
217
+ except Exception:
218
+ return buf
219
+ if not torch.is_tensor(moved) or moved.device == buf.device:
220
+ return buf
221
+ return buf.to(device=moved.device)
code/quant/quantize_stream.py ADDED
@@ -0,0 +1,556 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Stream a Ming MLLM directory to weight-only INT8 shards.
2
+
3
+ Never builds the model: it buffers at most one output shard (<= 5 GB) of tensors at a time. Measured on
4
+ the real 34.0 GB checkpoint (AMD Strix Halo, 2026-09-23): 266 s wall, peak RSS 17.8 GiB.
5
+ CLI: ``python quantize_stream.py SRC_MLLM_DIR DST_DIR [--exclude MODULE_REGEX]`` (matching modules stay BF16).
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+ import math
12
+ import re
13
+ import os
14
+ import shutil
15
+ import sys
16
+ from dataclasses import dataclass
17
+ from pathlib import Path
18
+
19
+ import torch
20
+ from safetensors import safe_open
21
+ from safetensors.torch import save_file
22
+
23
+ try: # imported as the `quant` package
24
+ from .int8_linear import QUANT_RULE, quant_rule_leaf, quantize_weight
25
+ except ImportError: # run as a script: python quant/quantize_stream.py SRC DST
26
+ from int8_linear import QUANT_RULE, quant_rule_leaf, quantize_weight
27
+
28
+ # Decimal GB, same unit Hugging Face uses for max_shard_size="5GB".
29
+ MAX_SHARD_BYTES = 5 * 10**9
30
+
31
+ _DTYPE_BYTES = {
32
+ "BOOL": 1,
33
+ "U8": 1,
34
+ "I8": 1,
35
+ "F8_E4M3": 1,
36
+ "F8_E5M2": 1,
37
+ "F8_E8M0": 1,
38
+ "U16": 2,
39
+ "I16": 2,
40
+ "F16": 2,
41
+ "BF16": 2,
42
+ "U32": 4,
43
+ "I32": 4,
44
+ "F32": 4,
45
+ "U64": 8,
46
+ "I64": 8,
47
+ "F64": 8,
48
+ }
49
+
50
+ INDEX_NAME = "model.safetensors.index.json"
51
+ MANIFEST_NAME = "int8_manifest.json"
52
+
53
+
54
+ class QuantizeError(Exception):
55
+ """User-facing checkpoint error. main() prints it and returns 1."""
56
+
57
+
58
+ def _die(msg: str) -> None:
59
+ raise QuantizeError(msg)
60
+
61
+
62
+ def _normalize_dtype(dtype_name) -> str:
63
+ text = str(dtype_name).upper()
64
+ if "." in text:
65
+ text = text.rsplit(".", 1)[-1]
66
+ aliases = {
67
+ "BFLOAT16": "BF16",
68
+ "FLOAT16": "F16",
69
+ "FLOAT32": "F32",
70
+ "FLOAT64": "F64",
71
+ "FLOAT8_E4M3FN": "F8_E4M3",
72
+ "FLOAT8_E5M2": "F8_E5M2",
73
+ "INT8": "I8",
74
+ "INT16": "I16",
75
+ "INT32": "I32",
76
+ "INT64": "I64",
77
+ "UINT8": "U8",
78
+ }
79
+ return aliases.get(text, text)
80
+
81
+
82
+ def _dtype_nbytes(dtype_name: str) -> int:
83
+ try:
84
+ return _DTYPE_BYTES[dtype_name]
85
+ except KeyError:
86
+ _die(f"unsupported safetensors dtype {dtype_name!r}")
87
+ raise # unreachable; satisfies type checkers
88
+
89
+
90
+ def _numel(shape: tuple[int, ...]) -> int:
91
+ n = 1
92
+ for d in shape:
93
+ n *= int(d)
94
+ return n
95
+
96
+
97
+ def _load_index(path: Path) -> dict:
98
+ if not path.is_file():
99
+ _die(f"missing index: {path}")
100
+
101
+ def _pairs(pairs):
102
+ keys = [k for k, _ in pairs]
103
+ dupes = sorted({k for k in keys if keys.count(k) > 1})
104
+ if dupes:
105
+ _die(f"duplicate key(s) in {path}: {dupes}")
106
+ return dict(pairs)
107
+
108
+ try:
109
+ raw = path.read_text(encoding="utf-8")
110
+ index = json.loads(raw, object_pairs_hook=_pairs)
111
+ except QuantizeError:
112
+ raise
113
+ except (OSError, json.JSONDecodeError) as exc:
114
+ _die(f"cannot read index {path}: {exc}")
115
+ if not isinstance(index, dict) or not isinstance(index.get("weight_map"), dict):
116
+ _die(f"index {path} has no weight_map object")
117
+ if not index["weight_map"]:
118
+ _die(f"index {path} weight_map is empty")
119
+ return index
120
+
121
+
122
+ def _check_dst_clean(dst: Path) -> None:
123
+ if not dst.exists():
124
+ return
125
+ if not dst.is_dir():
126
+ _die(f"destination is not a directory: {dst}")
127
+ found = sorted(p.relative_to(dst).as_posix() for p in dst.rglob("*.safetensors"))
128
+ if found:
129
+ _die(f"destination already contains safetensors: {found}")
130
+
131
+
132
+ def _reject_nested(src: Path, dst: Path) -> None:
133
+ src_r = src.resolve()
134
+ dst_r = dst.resolve()
135
+ if src_r == dst_r or src_r in dst_r.parents or dst_r in src_r.parents:
136
+ _die(f"SRC and DST must be distinct and not nested: {src} vs {dst}")
137
+
138
+
139
+ def _shard_path(src: Path, shard_name: str) -> Path:
140
+ rel = Path(shard_name)
141
+ if rel.is_absolute() or ".." in rel.parts:
142
+ _die(f"unsafe shard path in index: {shard_name}")
143
+ path = src / rel
144
+ if not path.is_file():
145
+ _die(f"index lists missing shard: {shard_name}")
146
+ return path
147
+
148
+
149
+ @dataclass
150
+ class Item:
151
+ src_shard: str
152
+ name: str
153
+ kind: str # "copy" or "quant"
154
+ shape: tuple[int, ...]
155
+ src_dtype: str
156
+ src_bytes: int
157
+ out_bytes: int
158
+ group: int = -1
159
+
160
+
161
+ def _scale_name(weight_name: str) -> str:
162
+ return weight_name[: -len("weight")] + "scale"
163
+
164
+
165
+ def _plan(src: Path, index: dict, exclude: str | None = None) -> list[Item]:
166
+ """Metadata-only pass. Reads shapes and dtypes, not tensor bodies."""
167
+ weight_map: dict[str, str] = index["weight_map"]
168
+ shard_order: list[str] = []
169
+ seen_shards: set[str] = set()
170
+ for shard in weight_map.values():
171
+ if shard not in seen_shards:
172
+ seen_shards.add(shard)
173
+ shard_order.append(shard)
174
+
175
+ index_names_by_shard: dict[str, set[str]] = {s: set() for s in shard_order}
176
+ for name, shard in weight_map.items():
177
+ if shard not in index_names_by_shard:
178
+ _die(f"weight_map value {shard!r} for {name} was not collected")
179
+ index_names_by_shard[shard].add(name)
180
+
181
+ items: list[Item] = []
182
+ seen_names: dict[str, str] = {}
183
+ for shard in shard_order:
184
+ path = _shard_path(src, shard)
185
+ with safe_open(str(path), framework="pt", device="cpu") as handle:
186
+ file_names = list(handle.keys())
187
+ file_set = set(file_names)
188
+ if len(file_set) != len(file_names):
189
+ _die(f"shard {shard} header lists a tensor name twice")
190
+ missing = sorted(index_names_by_shard[shard] - file_set)
191
+ extra = sorted(file_set - index_names_by_shard[shard])
192
+ if missing:
193
+ _die(f"index lists tensors missing from {shard}: {missing}")
194
+ if extra:
195
+ _die(f"{shard} contains tensors absent from the index: {extra}")
196
+ for name in file_names:
197
+ if name in seen_names:
198
+ _die(
199
+ f"tensor name appears twice: {name} "
200
+ f"({seen_names[name]} and {shard})"
201
+ )
202
+ seen_names[name] = shard
203
+ sl = handle.get_slice(name)
204
+ if not hasattr(sl, "get_dtype") or not hasattr(sl, "get_shape"):
205
+ _die(
206
+ "safetensors safe_open slice is missing get_shape/get_dtype; "
207
+ "cannot plan shards without loading tensor bodies"
208
+ )
209
+ shape = tuple(int(d) for d in sl.get_shape())
210
+ dtype_name = _normalize_dtype(sl.get_dtype())
211
+ src_bytes = _numel(shape) * _dtype_nbytes(dtype_name)
212
+ leaf = quant_rule_leaf(name)
213
+ if leaf is not None and exclude and re.search(exclude, name[: -len(".weight")]):
214
+ leaf = None # kept BF16 by --exclude
215
+ if leaf is not None and len(shape) != 2:
216
+ _die(
217
+ f"tensor {name} matches the quantize rule but is not 2-D "
218
+ f"(shape={list(shape)}, dtype={dtype_name})"
219
+ )
220
+ if leaf is not None:
221
+ out_bytes = _numel(shape) * 1 + shape[0] * 4 # int8 weight + fp32 scale
222
+ items.append(
223
+ Item(shard, name, "quant", shape, dtype_name, src_bytes, out_bytes)
224
+ )
225
+ else:
226
+ items.append(
227
+ Item(shard, name, "copy", shape, dtype_name, src_bytes, src_bytes)
228
+ )
229
+
230
+ index_names = set(weight_map)
231
+ planned = {it.name for it in items}
232
+ if planned != index_names:
233
+ _die(
234
+ "index / shard mismatch after scan: "
235
+ f"only_in_index={sorted(index_names - planned)[:8]} "
236
+ f"only_in_shards={sorted(planned - index_names)[:8]}"
237
+ )
238
+
239
+ produced = set(planned)
240
+ for it in items:
241
+ if it.kind != "quant":
242
+ continue
243
+ sname = _scale_name(it.name)
244
+ if sname in produced:
245
+ _die(f"scale name collides with an existing tensor: {sname}")
246
+ produced.add(sname)
247
+ return items
248
+
249
+
250
+ def _assign_groups(items: list[Item], max_shard_bytes: int) -> list[list[Item]]:
251
+ if max_shard_bytes <= 0:
252
+ _die(f"max_shard_bytes must be positive, got {max_shard_bytes}")
253
+ groups: list[list[Item]] = []
254
+ cur: list[Item] = []
255
+ cur_bytes = 0
256
+ for it in items:
257
+ if cur and cur_bytes + it.out_bytes > max_shard_bytes:
258
+ groups.append(cur)
259
+ cur = []
260
+ cur_bytes = 0
261
+ if cur_bytes == 0 and it.out_bytes > max_shard_bytes:
262
+ print(
263
+ f"warning: {it.name} contributes {it.out_bytes} bytes, "
264
+ f"over the {max_shard_bytes}-byte shard target; writing it alone",
265
+ file=sys.stderr,
266
+ flush=True,
267
+ )
268
+ it.group = len(groups)
269
+ cur.append(it)
270
+ cur_bytes += it.out_bytes
271
+ if cur:
272
+ groups.append(cur)
273
+ return groups
274
+
275
+
276
+ def _relative_frobenius(weight: torch.Tensor, q: torch.Tensor, scale: torch.Tensor) -> float:
277
+ w = weight.detach().to(dtype=torch.float64)
278
+ deq = q.detach().to(dtype=torch.float64) * scale.detach().to(dtype=torch.float64)[:, None]
279
+ denom = torch.linalg.matrix_norm(w, ord="fro")
280
+ numer = torch.linalg.matrix_norm(w - deq, ord="fro")
281
+ d = denom.item()
282
+ n = numer.item()
283
+ if d == 0.0:
284
+ return 0.0 if n == 0.0 else math.inf
285
+ return n / d
286
+
287
+
288
+ def _percentile_linear(values: list[float], pct: float) -> float:
289
+ """NumPy-style linear percentile. Empty → 0."""
290
+ if not values:
291
+ return 0.0
292
+ ordered = sorted(values)
293
+ if len(ordered) == 1:
294
+ return ordered[0]
295
+ rank = (len(ordered) - 1) * (pct / 100.0)
296
+ lo = math.floor(rank)
297
+ hi = math.ceil(rank)
298
+ if lo == hi:
299
+ return ordered[lo]
300
+ w = rank - lo
301
+ return ordered[lo] * (1.0 - w) + ordered[hi] * w
302
+
303
+
304
+ def _copy_sidecars(src: Path, dst: Path) -> list[str]:
305
+ copied: list[str] = []
306
+ for dirpath, _dirnames, filenames in os.walk(src):
307
+ rel = Path(dirpath).relative_to(src)
308
+ out_dir = dst / rel
309
+ out_dir.mkdir(parents=True, exist_ok=True)
310
+ for filename in filenames:
311
+ if filename.endswith(".safetensors"):
312
+ continue
313
+ if filename == INDEX_NAME and rel == Path("."):
314
+ continue
315
+ src_file = Path(dirpath) / filename
316
+ dst_file = out_dir / filename
317
+ shutil.copy2(src_file, dst_file)
318
+ copied.append((rel / filename).as_posix())
319
+ return copied
320
+
321
+
322
+ def _write_shards(
323
+ src: Path,
324
+ dst: Path,
325
+ items: list[Item],
326
+ groups: list[list[Item]],
327
+ ) -> tuple[dict[str, str], int, int, list[tuple[str, float]], list[Path]]:
328
+ n_out = len(groups)
329
+ weight_map: dict[str, str] = {}
330
+ bytes_in = 0
331
+ bytes_out = 0
332
+ errors: list[tuple[str, float]] = []
333
+ written: list[Path] = []
334
+
335
+ n_src = len({it.src_shard for it in items})
336
+ src_seen = 0
337
+ open_name: str | None = None
338
+ handle = None
339
+ buf: dict[str, torch.Tensor] = {}
340
+ buf_q = 0
341
+ buf_c = 0
342
+ current_group = 0
343
+
344
+ def flush() -> None:
345
+ nonlocal buf, buf_q, buf_c, current_group
346
+ if not buf:
347
+ return
348
+ fname = f"model-{current_group + 1:05d}-of-{n_out:05d}.safetensors"
349
+ path = dst / fname
350
+ for key, tensor in buf.items():
351
+ if not tensor.is_contiguous():
352
+ buf[key] = tensor.contiguous()
353
+ save_file(buf, str(path))
354
+ shard_bytes = 0
355
+ for key, tensor in buf.items():
356
+ weight_map[key] = fname
357
+ shard_bytes += tensor.numel() * tensor.element_size()
358
+ written.append(path)
359
+ print(
360
+ f"wrote {fname}: tensors={len(buf)} quantized={buf_q} copied={buf_c} "
361
+ f"bytes={shard_bytes}",
362
+ flush=True,
363
+ )
364
+ buf = {}
365
+ buf_q = 0
366
+ buf_c = 0
367
+ current_group += 1
368
+
369
+ try:
370
+ for it in items:
371
+ if it.src_shard != open_name:
372
+ if handle is not None:
373
+ handle.__exit__(None, None, None)
374
+ handle = None
375
+ path = _shard_path(src, it.src_shard)
376
+ handle = safe_open(str(path), framework="pt", device="cpu")
377
+ handle.__enter__()
378
+ open_name = it.src_shard
379
+ src_seen += 1
380
+ n_here = sum(1 for x in items if x.src_shard == it.src_shard)
381
+ print(
382
+ f"reading source shard {src_seen}/{n_src} {it.src_shard} ({n_here} tensors)",
383
+ flush=True,
384
+ )
385
+ assert handle is not None
386
+ tensor = handle.get_tensor(it.name)
387
+ got = tensor.numel() * tensor.element_size()
388
+ if got != it.src_bytes:
389
+ _die(
390
+ f"{it.name} byte size {got} != planned {it.src_bytes} "
391
+ f"(dtype={tensor.dtype}, shape={tuple(tensor.shape)})"
392
+ )
393
+ bytes_in += got
394
+ if it.kind == "quant":
395
+ if not tensor.is_floating_point():
396
+ _die(
397
+ f"{it.name} matches the quantize rule but dtype is {tensor.dtype}, "
398
+ "expected a floating dtype"
399
+ )
400
+ if tuple(tensor.shape) != it.shape:
401
+ _die(f"{it.name} shape changed between passes: {tuple(tensor.shape)} vs {it.shape}")
402
+ q, scale = quantize_weight(tensor)
403
+ err = _relative_frobenius(tensor, q, scale)
404
+ if math.isnan(err) or math.isinf(err):
405
+ _die(f"non-finite relative error for {it.name}: {err}")
406
+ errors.append((it.name, err))
407
+ del tensor
408
+ sname = _scale_name(it.name)
409
+ buf[it.name] = q
410
+ buf[sname] = scale
411
+ produced = q.numel() * q.element_size() + scale.numel() * scale.element_size()
412
+ if produced != it.out_bytes:
413
+ _die(f"{it.name} output bytes {produced} != planned {it.out_bytes}")
414
+ buf_q += 1
415
+ else:
416
+ if not tensor.is_contiguous():
417
+ tensor = tensor.contiguous()
418
+ buf[it.name] = tensor
419
+ buf_c += 1
420
+ bytes_out += it.out_bytes
421
+ # Flush when this item closes its planned output shard.
422
+ group_items = groups[it.group]
423
+ if it is group_items[-1]:
424
+ flush()
425
+ finally:
426
+ if handle is not None:
427
+ handle.__exit__(None, None, None)
428
+
429
+ if buf:
430
+ _die("internal error: output buffer not flushed")
431
+ if current_group != n_out:
432
+ _die(f"internal error: wrote {current_group} shards, planned {n_out}")
433
+ return weight_map, bytes_in, bytes_out, errors, written
434
+
435
+
436
+ def _summary(
437
+ errors: list[tuple[str, float]],
438
+ n_quant: int,
439
+ n_copy: int,
440
+ bytes_in: int,
441
+ bytes_out: int,
442
+ ) -> dict:
443
+ vals = [e for _, e in errors]
444
+ if errors:
445
+ worst_name, worst_err = min(
446
+ errors,
447
+ key=lambda pair: (-pair[1], pair[0]),
448
+ )
449
+ else:
450
+ worst_name, worst_err = None, 0.0
451
+ mean = (sum(vals) / len(vals)) if vals else 0.0
452
+ return {
453
+ "tensors_quantized": n_quant,
454
+ "tensors_copied": n_copy,
455
+ "bytes_in": bytes_in,
456
+ "bytes_out": bytes_out,
457
+ "mean_relative_error": mean,
458
+ "p99_relative_error": _percentile_linear(vals, 99.0),
459
+ "max_relative_error": worst_err if vals else 0.0,
460
+ "worst_tensor": worst_name,
461
+ }
462
+
463
+
464
+ def run(src: Path, dst: Path, max_shard_bytes: int = MAX_SHARD_BYTES, exclude: str | None = None) -> dict:
465
+ src = src.resolve()
466
+ dst = dst.resolve()
467
+ if not src.is_dir():
468
+ _die(f"SRC is not a directory: {src}")
469
+ _reject_nested(src, dst)
470
+ _check_dst_clean(dst)
471
+ index = _load_index(src / INDEX_NAME)
472
+ items = _plan(src, index, exclude)
473
+ groups = _assign_groups(items, max_shard_bytes)
474
+ dst.mkdir(parents=True, exist_ok=True)
475
+
476
+ written: list[Path] = []
477
+ try:
478
+ weight_map, bytes_in, bytes_out, errors, written = _write_shards(src, dst, items, groups)
479
+ copied = _copy_sidecars(src, dst)
480
+ n_quant = sum(1 for it in items if it.kind == "quant")
481
+ n_copy = sum(1 for it in items if it.kind == "copy")
482
+ measured = _summary(errors, n_quant, n_copy, bytes_in, bytes_out)
483
+ if measured["bytes_in"] != bytes_in or measured["bytes_out"] != bytes_out:
484
+ _die("internal error: summary byte counters diverged")
485
+ # Recompute the on-disk total from the tensors we recorded. weight_map
486
+ # values are what we just saved; bytes_out is that sum.
487
+ out_index = {"metadata": {"total_size": bytes_out}, "weight_map": weight_map}
488
+ (dst / INDEX_NAME).write_text(
489
+ json.dumps(out_index, indent=2) + "\n", encoding="utf-8"
490
+ )
491
+ modules = sorted(
492
+ it.name[: -len(".weight")] for it in items if it.kind == "quant"
493
+ )
494
+ manifest = {
495
+ "format": "ming-int8-wo-v1",
496
+ "scheme": "weight-only int8, per-output-channel symmetric, fp32 scales",
497
+ "rule": QUANT_RULE + (f" Additionally kept BF16: modules matching /{exclude}/." if exclude else ""),
498
+ "exclude": exclude,
499
+ "quantized_modules": modules,
500
+ "source_total_size": bytes_in,
501
+ "total_size": bytes_out,
502
+ "measured": measured,
503
+ }
504
+ (dst / MANIFEST_NAME).write_text(
505
+ json.dumps(manifest, indent=2, allow_nan=False) + "\n", encoding="utf-8"
506
+ )
507
+ except Exception:
508
+ for path in written:
509
+ try:
510
+ path.unlink()
511
+ except OSError:
512
+ pass
513
+ raise
514
+
515
+ print(f"copied {len(copied)} non-safetensors file(s)", flush=True)
516
+ m = measured
517
+ print(
518
+ "summary: "
519
+ f"quantized={m['tensors_quantized']} copied={m['tensors_copied']} "
520
+ f"bytes_in={m['bytes_in']} bytes_out={m['bytes_out']} "
521
+ f"mean_rel={m['mean_relative_error']:.8g} "
522
+ f"p99_rel={m['p99_relative_error']:.8g} "
523
+ f"max_rel={m['max_relative_error']:.8g} "
524
+ f"worst={m['worst_tensor']}",
525
+ flush=True,
526
+ )
527
+ return manifest
528
+
529
+
530
+ def main(argv: list[str] | None = None) -> int:
531
+ args = list(sys.argv[1:] if argv is None else argv)
532
+ exclude = None
533
+ if "--exclude" in args:
534
+ i = args.index("--exclude")
535
+ if i + 1 >= len(args):
536
+ print("--exclude needs a regex", file=sys.stderr)
537
+ return 2
538
+ exclude = args[i + 1]
539
+ re.compile(exclude)
540
+ del args[i : i + 2]
541
+ if len(args) != 2:
542
+ print(
543
+ "usage: python quantize_stream.py SRC_MLLM_DIR DST_DIR [--exclude MODULE_REGEX]",
544
+ file=sys.stderr,
545
+ )
546
+ return 2
547
+ try:
548
+ run(Path(args[0]), Path(args[1]), max_shard_bytes=MAX_SHARD_BYTES, exclude=exclude)
549
+ except QuantizeError as exc:
550
+ print(f"error: {exc}", file=sys.stderr)
551
+ return 1
552
+ return 0
553
+
554
+
555
+ if __name__ == "__main__":
556
+ sys.exit(main())
code/qwen2_5_vit.py ADDED
@@ -0,0 +1,508 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2025 The Qwen Team and The HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
5
+ # and OPT implementations in this library. It has been modified from its
6
+ # original forms to accommodate minor architectural differences compared
7
+ # to GPT-NeoX and OPT used by the Meta AI team that trained the model.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ """PyTorch Qwen2_5_ViT model."""
21
+
22
+ import math
23
+
24
+ import os
25
+ import torch
26
+ import torch.nn as nn
27
+ import torch.nn.functional as F
28
+
29
+ from transformers.activations import ACT2FN
30
+ from transformers.modeling_utils import PreTrainedModel
31
+ from transformers.utils import (
32
+ is_flash_attn_2_available,
33
+ logging,
34
+ )
35
+
36
+ from typing import Union
37
+
38
+ from transformers.configuration_utils import PretrainedConfig
39
+
40
+ if is_flash_attn_2_available():
41
+ from flash_attn import flash_attn_varlen_func
42
+ from flash_attn.layers.rotary import apply_rotary_emb
43
+
44
+ else:
45
+ flash_attn_varlen_func = None
46
+ apply_rotary_emb = None
47
+
48
+ logger = logging.get_logger(__name__)
49
+
50
+ class Qwen2_5_VLVisionConfig(PretrainedConfig):
51
+ model_type = "qwen2_5_vit"
52
+
53
+ def __init__(
54
+ self,
55
+ depth=32,
56
+ hidden_size=3584,
57
+ hidden_act="silu",
58
+ intermediate_size=3420,
59
+ num_heads=16,
60
+ in_channels=3,
61
+ patch_size=14,
62
+ spatial_merge_size=2,
63
+ temporal_patch_size=2,
64
+ tokens_per_second=4,
65
+ window_size=112,
66
+ out_hidden_size=3584,
67
+ fullatt_block_indexes=[7, 15, 23, 31],
68
+ _attn_implementation="flash_attention_2",
69
+ **kwargs,
70
+ ):
71
+ super().__init__(**kwargs)
72
+ self.depth = depth
73
+ self.hidden_size = hidden_size
74
+ self.hidden_act = hidden_act
75
+ self.intermediate_size = intermediate_size
76
+ self.num_heads = num_heads
77
+ self.in_channels = in_channels
78
+ self.patch_size = patch_size
79
+ self.spatial_merge_size = spatial_merge_size
80
+ self.temporal_patch_size = temporal_patch_size
81
+ self.tokens_per_second = tokens_per_second
82
+ self.window_size = window_size
83
+ self.fullatt_block_indexes = fullatt_block_indexes
84
+ self.out_hidden_size = out_hidden_size
85
+ self._attn_implementation = _attn_implementation
86
+
87
+ @classmethod
88
+ def from_pretrained(cls, pretrained_model_name_or_path: Union[str, os.PathLike], **kwargs) -> "PretrainedConfig":
89
+ cls._set_token_in_kwargs(kwargs)
90
+
91
+ config_dict, kwargs = cls.get_config_dict(pretrained_model_name_or_path, **kwargs)
92
+
93
+ if 'vision_config' in config_dict:
94
+ config_dict = config_dict['vision_config']
95
+
96
+ if "model_type" in config_dict and hasattr(cls, "model_type") and config_dict["model_type"] != cls.model_type:
97
+ logger.warning(
98
+ f"You are using a model of type {config_dict['model_type']} to instantiate a model of type "
99
+ f"{cls.model_type}. This is not supported for all configurations of models and can yield errors."
100
+ )
101
+
102
+ return cls.from_dict(config_dict, **kwargs)
103
+
104
+ class Qwen2_5_VLMLP(nn.Module):
105
+ def __init__(self, config, bias: bool = False):
106
+ super().__init__()
107
+ self.hidden_size = config.hidden_size
108
+ self.intermediate_size = config.intermediate_size
109
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=bias)
110
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=bias)
111
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=bias)
112
+ self.act_fn = ACT2FN[config.hidden_act]
113
+
114
+ def forward(self, hidden_state):
115
+ return self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state))
116
+
117
+ class Qwen2_5_VisionPatchEmbed(nn.Module):
118
+ def __init__(
119
+ self,
120
+ patch_size: int = 14,
121
+ temporal_patch_size: int = 2,
122
+ in_channels: int = 3,
123
+ embed_dim: int = 1152,
124
+ ) -> None:
125
+ super().__init__()
126
+ self.patch_size = patch_size
127
+ self.temporal_patch_size = temporal_patch_size
128
+ self.in_channels = in_channels
129
+ self.embed_dim = embed_dim
130
+
131
+ kernel_size = [temporal_patch_size, patch_size, patch_size]
132
+ self.proj = nn.Conv3d(in_channels, embed_dim, kernel_size=kernel_size, stride=kernel_size, bias=False)
133
+
134
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
135
+ target_dtype = self.proj.weight.dtype
136
+ hidden_states = hidden_states.view(
137
+ -1, self.in_channels, self.temporal_patch_size, self.patch_size, self.patch_size
138
+ )
139
+ hidden_states = self.proj(hidden_states.to(dtype=target_dtype)).view(-1, self.embed_dim)
140
+ return hidden_states
141
+
142
+ class Qwen2_5_VisionRotaryEmbedding(nn.Module):
143
+ def __init__(self, dim: int, theta: float = 10000.0) -> None:
144
+ super().__init__()
145
+ inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim))
146
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
147
+ self.theta = theta
148
+ self.dim = dim
149
+
150
+ def forward(self, seqlen: int) -> torch.Tensor:
151
+ seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype)
152
+ freqs = torch.outer(seq, self.inv_freq)
153
+ return freqs
154
+
155
+ def reset_parameters(self):
156
+ # recompute inv_freq (for dynamic adjustment)
157
+ new_inv_freq = 1.0 / (self.theta ** (torch.arange(0, self.dim, 2, dtype=torch.float) / self.dim))
158
+ self.inv_freq.copy_(new_inv_freq)
159
+
160
+ class Qwen2RMSNorm(nn.Module):
161
+ def __init__(self, hidden_size, eps=1e-6):
162
+ """
163
+ Qwen2RMSNorm is equivalent to T5LayerNorm.
164
+
165
+ Replaces transformer_engine.pytorch.RMSNorm: ROCm has no transformer-engine.
166
+ te.RMSNorm defaults (zero_centered_gamma=False) are standard RMSNorm, and the
167
+ checkpoint stores this affine as `weight`.
168
+ """
169
+ super().__init__()
170
+ self.weight = nn.Parameter(torch.ones(hidden_size))
171
+ self.variance_epsilon = eps
172
+
173
+ def forward(self, hidden_states):
174
+ input_dtype = hidden_states.dtype
175
+ hidden_states = hidden_states.to(torch.float32)
176
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
177
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
178
+ return self.weight * hidden_states.to(input_dtype)
179
+
180
+ def extra_repr(self):
181
+ return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
182
+
183
+ class Qwen2_5_VLPatchMerger(nn.Module):
184
+ def __init__(self, dim: int, context_dim: int, spatial_merge_size: int = 2) -> None:
185
+ super().__init__()
186
+ self.hidden_size = context_dim * (spatial_merge_size ** 2)
187
+ self.ln_q = Qwen2RMSNorm(context_dim, eps=1e-6)
188
+ self.mlp = nn.Sequential(
189
+ nn.Linear(self.hidden_size, self.hidden_size),
190
+ nn.GELU(),
191
+ nn.Linear(self.hidden_size, dim),
192
+ )
193
+
194
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
195
+ x = self.mlp(self.ln_q(x).view(-1, self.hidden_size))
196
+ return x
197
+
198
+ def apply_rotary_pos_emb_flashatt(tensor: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
199
+ tensor_ = tensor.float()
200
+ cos = freqs.cos().float()
201
+ sin = freqs.sin().float()
202
+ output = apply_rotary_emb(tensor_, cos, sin).type_as(tensor)
203
+ return output
204
+
205
+ class Qwen2_5_VLVisionFlashAttention2(nn.Module):
206
+ def __init__(self, dim: int, num_heads: int = 16) -> None:
207
+ super().__init__()
208
+ self.num_heads = num_heads
209
+ self.qkv = nn.Linear(dim, dim * 3, bias=True)
210
+ self.proj = nn.Linear(dim, dim)
211
+
212
+ def forward(
213
+ self,
214
+ hidden_states: torch.Tensor,
215
+ cu_seqlens: torch.Tensor,
216
+ rotary_pos_emb: torch.Tensor = None,
217
+ ) -> torch.Tensor:
218
+ seq_length = hidden_states.shape[0]
219
+ q, k, v = self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0)
220
+ q = apply_rotary_pos_emb_flashatt(q.unsqueeze(0), rotary_pos_emb).squeeze(0)
221
+ k = apply_rotary_pos_emb_flashatt(k.unsqueeze(0), rotary_pos_emb).squeeze(0)
222
+
223
+ max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max().item()
224
+ attn_output = flash_attn_varlen_func(q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen).reshape(
225
+ seq_length, -1
226
+ )
227
+ attn_output = self.proj(attn_output)
228
+ return attn_output
229
+
230
+ def rotate_half(x):
231
+ """Rotates half the hidden dims of the input."""
232
+ x1 = x[..., : x.shape[-1] // 2]
233
+ x2 = x[..., x.shape[-1] // 2:]
234
+ return torch.cat((-x2, x1), dim=-1)
235
+
236
+ def apply_rotary_pos_emb_vision(tensor: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
237
+ orig_dtype = tensor.dtype
238
+ tensor = tensor.float()
239
+ cos = freqs.cos()
240
+ sin = freqs.sin()
241
+ cos = cos.unsqueeze(1).repeat(1, 1, 2).unsqueeze(0).float()
242
+ sin = sin.unsqueeze(1).repeat(1, 1, 2).unsqueeze(0).float()
243
+ output = (tensor * cos) + (rotate_half(tensor) * sin)
244
+ output = output.to(orig_dtype)
245
+ return output
246
+
247
+ class Qwen2_5_VLVisionAttention(nn.Module):
248
+ def __init__(self, dim: int, num_heads: int = 16) -> None:
249
+ super().__init__()
250
+ self.num_heads = num_heads
251
+ self.head_dim = dim // num_heads
252
+ self.qkv = nn.Linear(dim, dim * 3, bias=True)
253
+ self.proj = nn.Linear(dim, dim)
254
+
255
+ def forward(
256
+ self, hidden_states: torch.Tensor, cu_seqlens: torch.Tensor, rotary_pos_emb: torch.Tensor = None
257
+ ) -> torch.Tensor:
258
+ seq_length = hidden_states.shape[0]
259
+ q, k, v = self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0)
260
+ q = apply_rotary_pos_emb_vision(q.unsqueeze(0), rotary_pos_emb).squeeze(0)
261
+ k = apply_rotary_pos_emb_vision(k.unsqueeze(0), rotary_pos_emb).squeeze(0)
262
+
263
+ attention_mask = torch.full(
264
+ [1, seq_length, seq_length], torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype
265
+ )
266
+ for i in range(1, len(cu_seqlens)):
267
+ attention_mask[..., cu_seqlens[i - 1]: cu_seqlens[i], cu_seqlens[i - 1]: cu_seqlens[i]] = 0
268
+
269
+ q = q.transpose(0, 1)
270
+ k = k.transpose(0, 1)
271
+ v = v.transpose(0, 1)
272
+ attn_weights = torch.matmul(q, k.transpose(1, 2)) / math.sqrt(self.head_dim)
273
+ attn_weights = attn_weights + attention_mask
274
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(q.dtype)
275
+ attn_output = torch.matmul(attn_weights, v)
276
+ attn_output = attn_output.transpose(0, 1)
277
+ attn_output = attn_output.reshape(seq_length, -1)
278
+ attn_output = self.proj(attn_output)
279
+ return attn_output
280
+
281
+ class Qwen2_5_VLVisionSdpaAttention(nn.Module):
282
+ def __init__(self, dim: int, num_heads: int = 16) -> None:
283
+ super().__init__()
284
+ self.num_heads = num_heads
285
+ self.qkv = nn.Linear(dim, dim * 3, bias=True)
286
+ self.proj = nn.Linear(dim, dim)
287
+
288
+ def forward(
289
+ self, hidden_states: torch.Tensor, cu_seqlens: torch.Tensor, rotary_pos_emb: torch.Tensor = None
290
+ ) -> torch.Tensor:
291
+ seq_length = hidden_states.shape[0]
292
+ q, k, v = self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0)
293
+ q = apply_rotary_pos_emb_vision(q.unsqueeze(0), rotary_pos_emb).squeeze(0)
294
+ k = apply_rotary_pos_emb_vision(k.unsqueeze(0), rotary_pos_emb).squeeze(0)
295
+
296
+ attention_mask = torch.zeros([1, seq_length, seq_length], device=q.device, dtype=torch.bool)
297
+ for i in range(1, len(cu_seqlens)):
298
+ attention_mask[..., cu_seqlens[i - 1]: cu_seqlens[i], cu_seqlens[i - 1]: cu_seqlens[i]] = True
299
+ q = q.transpose(0, 1)
300
+ k = k.transpose(0, 1)
301
+ v = v.transpose(0, 1)
302
+ attn_output = F.scaled_dot_product_attention(q, k, v, attention_mask, dropout_p=0.0)
303
+ attn_output = attn_output.transpose(0, 1)
304
+ attn_output = attn_output.reshape(seq_length, -1)
305
+ attn_output = self.proj(attn_output)
306
+ return attn_output
307
+
308
+ QWEN2_5_VL_VISION_ATTENTION_CLASSES = {
309
+ "eager": Qwen2_5_VLVisionAttention,
310
+ "flash_attention_2": Qwen2_5_VLVisionFlashAttention2,
311
+ "sdpa": Qwen2_5_VLVisionSdpaAttention,
312
+ }
313
+
314
+ class Qwen2_5_VLVisionBlock(nn.Module):
315
+ def __init__(self, config, attn_implementation: str = "sdpa") -> None:
316
+ super().__init__()
317
+ self.norm1 = Qwen2RMSNorm(config.hidden_size, eps=1e-6)
318
+ self.norm2 = Qwen2RMSNorm(config.hidden_size, eps=1e-6)
319
+ self.attn = QWEN2_5_VL_VISION_ATTENTION_CLASSES[attn_implementation](
320
+ config.hidden_size, num_heads=config.num_heads
321
+ )
322
+ self.mlp = Qwen2_5_VLMLP(config, bias=True)
323
+
324
+ def forward(self, hidden_states, cu_seqlens, rotary_pos_emb) -> torch.Tensor:
325
+ hidden_states = hidden_states + self.attn(
326
+ self.norm1(hidden_states),
327
+ cu_seqlens=cu_seqlens,
328
+ rotary_pos_emb=rotary_pos_emb,
329
+ )
330
+ hidden_states = hidden_states + self.mlp(self.norm2(hidden_states))
331
+ return hidden_states
332
+
333
+ class Qwen2_5_VisionTransformer(PreTrainedModel):
334
+ config_class = Qwen2_5_VLVisionConfig
335
+ _no_split_modules = ["Qwen2_5_VLVisionBlock"]
336
+ _supports_flash_attn_2 = True
337
+ _supports_sdpa = True
338
+
339
+ def __init__(self, config, *inputs, **kwargs) -> None:
340
+ super().__init__(config, *inputs, **kwargs)
341
+ self.spatial_merge_size = config.spatial_merge_size
342
+ self.patch_size = config.patch_size
343
+ self.fullatt_block_indexes = config.fullatt_block_indexes
344
+ self.window_size = config.window_size
345
+ self.spatial_merge_unit = self.spatial_merge_size * self.spatial_merge_size
346
+
347
+ self.patch_embed = Qwen2_5_VisionPatchEmbed(
348
+ patch_size=config.patch_size,
349
+ temporal_patch_size=config.temporal_patch_size,
350
+ in_channels=config.in_channels,
351
+ embed_dim=config.hidden_size,
352
+ )
353
+
354
+ head_dim = config.hidden_size // config.num_heads
355
+ self.rotary_pos_emb = Qwen2_5_VisionRotaryEmbedding(head_dim // 2)
356
+
357
+ self.blocks = nn.ModuleList(
358
+ [Qwen2_5_VLVisionBlock(config, config._attn_implementation) for _ in range(config.depth)]
359
+ )
360
+ self.merger = Qwen2_5_VLPatchMerger(
361
+ dim=config.out_hidden_size,
362
+ context_dim=config.hidden_size,
363
+ spatial_merge_size=config.spatial_merge_size,
364
+ )
365
+ self.gradient_checkpointing = False
366
+ self.image_emb_dim = 8192
367
+
368
+ def get_dtype(self) -> torch.dtype:
369
+ return self.blocks[0].mlp.down_proj.weight.dtype
370
+
371
+ def rot_pos_emb(self, grid_thw):
372
+ pos_ids = []
373
+ for t, h, w in grid_thw:
374
+ hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
375
+ hpos_ids = hpos_ids.reshape(
376
+ h // self.spatial_merge_size,
377
+ self.spatial_merge_size,
378
+ w // self.spatial_merge_size,
379
+ self.spatial_merge_size,
380
+ )
381
+ hpos_ids = hpos_ids.permute(0, 2, 1, 3)
382
+ hpos_ids = hpos_ids.flatten()
383
+
384
+ wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
385
+ wpos_ids = wpos_ids.reshape(
386
+ h // self.spatial_merge_size,
387
+ self.spatial_merge_size,
388
+ w // self.spatial_merge_size,
389
+ self.spatial_merge_size,
390
+ )
391
+ wpos_ids = wpos_ids.permute(0, 2, 1, 3)
392
+ wpos_ids = wpos_ids.flatten()
393
+ pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))
394
+ pos_ids = torch.cat(pos_ids, dim=0)
395
+ max_grid_size = grid_thw[:, 1:].max()
396
+ rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size)
397
+ rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1)
398
+ return rotary_pos_emb
399
+
400
+ def get_window_index(self, grid_thw):
401
+ window_index: list = []
402
+ cu_window_seqlens: list = [0]
403
+ window_index_id = 0
404
+ vit_merger_window_size = self.window_size // self.spatial_merge_size // self.patch_size
405
+
406
+ for grid_t, grid_h, grid_w in grid_thw:
407
+ llm_grid_h, llm_grid_w = (
408
+ grid_h // self.spatial_merge_size,
409
+ grid_w // self.spatial_merge_size,
410
+ )
411
+ index = torch.arange(grid_t * llm_grid_h * llm_grid_w).reshape(grid_t, llm_grid_h, llm_grid_w)
412
+ pad_h = vit_merger_window_size - llm_grid_h % vit_merger_window_size
413
+ pad_w = vit_merger_window_size - llm_grid_w % vit_merger_window_size
414
+ num_windows_h = (llm_grid_h + pad_h) // vit_merger_window_size
415
+ num_windows_w = (llm_grid_w + pad_w) // vit_merger_window_size
416
+ index_padded = F.pad(index, (0, pad_w, 0, pad_h), "constant", -100)
417
+ index_padded = index_padded.reshape(
418
+ grid_t,
419
+ num_windows_h,
420
+ vit_merger_window_size,
421
+ num_windows_w,
422
+ vit_merger_window_size,
423
+ )
424
+ index_padded = index_padded.permute(0, 1, 3, 2, 4).reshape(
425
+ grid_t,
426
+ num_windows_h * num_windows_w,
427
+ vit_merger_window_size,
428
+ vit_merger_window_size,
429
+ )
430
+ seqlens = (index_padded != -100).sum([2, 3]).reshape(-1)
431
+ index_padded = index_padded.reshape(-1)
432
+ index_new = index_padded[index_padded != -100]
433
+ window_index.append(index_new + window_index_id)
434
+ cu_seqlens_tmp = seqlens.cumsum(0) * self.spatial_merge_unit + cu_window_seqlens[-1]
435
+ cu_window_seqlens.extend(cu_seqlens_tmp.tolist())
436
+ window_index_id += (grid_t * llm_grid_h * llm_grid_w).item()
437
+ window_index = torch.cat(window_index, dim=0)
438
+
439
+ return window_index, cu_window_seqlens
440
+
441
+ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor, is_list=False, **kwargs) -> torch.Tensor:
442
+ """
443
+ Args:
444
+ hidden_states (`torch.Tensor` of shape `(batch_size, seq_len, hidden_size)`):
445
+ The final hidden states of the model.
446
+ grid_thw (`torch.Tensor` of shape `(num_images_or_videos, 3)`):
447
+ The temporal, height and width of feature shape of each image in LLM.
448
+
449
+ Returns:
450
+ `torch.Tensor`: hidden_states.
451
+ """
452
+ hidden_states = hidden_states.type(self.get_dtype())
453
+ if is_list:
454
+ image_grid_thw_array = []
455
+ for image_grid_thw_ in grid_thw:
456
+ image_grid_thw_array.extend(image_grid_thw_)
457
+ grid_thw = torch.tensor(image_grid_thw_array, device=hidden_states.device)
458
+
459
+ hidden_states = self.patch_embed(hidden_states)
460
+ rotary_pos_emb = self.rot_pos_emb(grid_thw)
461
+ window_index, cu_window_seqlens = self.get_window_index(grid_thw)
462
+ cu_window_seqlens = torch.tensor(
463
+ cu_window_seqlens,
464
+ device=hidden_states.device,
465
+ dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32,
466
+ )
467
+ cu_window_seqlens = torch.unique_consecutive(cu_window_seqlens)
468
+
469
+ seq_len, _ = hidden_states.size()
470
+ hidden_states = hidden_states.reshape(seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1)
471
+ hidden_states = hidden_states[window_index, :, :]
472
+ hidden_states = hidden_states.reshape(seq_len, -1)
473
+ rotary_pos_emb = rotary_pos_emb.reshape(seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1)
474
+ rotary_pos_emb = rotary_pos_emb[window_index, :, :]
475
+ rotary_pos_emb = rotary_pos_emb.reshape(seq_len, -1)
476
+
477
+ cu_seqlens = torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]).cumsum(
478
+ dim=0,
479
+ # Select dtype based on the following factors:
480
+ # - FA2 requires that cu_seqlens_q must have dtype int32
481
+ # - torch.onnx.export requires that cu_seqlens_q must have same dtype as grid_thw
482
+ # See https://github.com/huggingface/transformers/pull/34852 for more information
483
+ dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32,
484
+ )
485
+ cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0)
486
+
487
+ for layer_num, blk in enumerate(self.blocks):
488
+ if layer_num in self.fullatt_block_indexes:
489
+ cu_seqlens_now = cu_seqlens
490
+ else:
491
+ cu_seqlens_now = cu_window_seqlens
492
+ if self.gradient_checkpointing and self.training:
493
+ hidden_states = self._gradient_checkpointing_func(
494
+ blk.__call__, hidden_states, cu_seqlens_now, rotary_pos_emb
495
+ )
496
+ else:
497
+ hidden_states = blk(
498
+ hidden_states,
499
+ cu_seqlens=cu_seqlens_now,
500
+ rotary_pos_emb=rotary_pos_emb,
501
+ )
502
+
503
+ hidden_states = self.merger(hidden_states)
504
+
505
+ reverse_indices = torch.argsort(window_index)
506
+ hidden_states = hidden_states[reverse_indices, :]
507
+
508
+ return hidden_states
code/requirements-rocm.txt ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ROCm port of requirements.txt for AMD gfx1151 (ROCm 7.13).
2
+ # Omitted on purpose — do not add them back:
3
+ # torch, torchvision: the target interpreter already has a working ROCm
4
+ # build (torch 2.10.0, torch.version.hip 7.13.99004). Reinstalling the
5
+ # upstream CUDA pins would replace it.
6
+ # transformer-engine: NVIDIA CUDA-only; ROCm has no TE. Qwen2RMSNorm in
7
+ # qwen2_5_vit.py is pure PyTorch, and the unused TE import is gone.
8
+ # Create the venv with system site packages so that ROCm torch is inherited:
9
+ # python3 -m venv --system-site-packages .venv
10
+ # .venv/bin/pip install -r requirements-rocm.txt
11
+ #
12
+ # Also omitted / relaxed versus upstream, because the ROCm interpreter is Python 3.13
13
+ # and its torch is built against numpy 2.x:
14
+ # numpy upstream 1.23.1 has no Python 3.13 wheels, and downgrading would
15
+ # break the inherited torch. Inherit the system numpy (validated 2.2.4).
16
+ # Pillow upstream 10.4.0 has no Python 3.13 wheels. Inherit (validated 11.1.0).
17
+ # safetensors inherit the system build (validated 0.8.0).
18
+ #
19
+ # Validated on halo (gfx1151, ROCm 7.13, Python 3.13.5) on 2026-09-22:
20
+ # torch 2.10.0 (hip 7.13.99004) | numpy 2.2.4 | Pillow 11.1.0 | safetensors 0.8.0
21
+ # transformers 4.57.1 | diffusers 0.36.0 | accelerate 1.13.0 | tokenizers 0.22.2
22
+ # huggingface-hub 0.34.0 | peft 0.17.0
23
+ transformers==4.57.1
24
+ diffusers==0.36.0
25
+ accelerate==1.13.0
26
+ tokenizers==0.22.2
27
+ huggingface-hub==0.34.0
28
+ peft==0.17.0
29
+ requests==2.32.3
30
+ tqdm==4.67.1
31
+ typing-extensions==4.15.0
32
+
33
+ # Optional FlashAttention 2 backend (validated: flash-attn==2.7.3). The CLI
34
+ # default is eager attention; --attn-implementation flash_attention_2 needs
35
+ # this package.
36
+ # flash-attn==2.7.3
code/requirements.txt ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Dependency set validated end-to-end on 2026-09-16. Exact pins are
2
+ # intentional: this is the environment known to run both published checkpoint
3
+ # families, not a broader compatibility matrix.
4
+ torch==2.4.0
5
+ torchvision==0.19.0
6
+ transformers==4.57.1
7
+ diffusers==0.36.0
8
+ accelerate==1.13.0
9
+ transformer-engine[pytorch]==1.11.0
10
+ safetensors==0.7.0
11
+ tokenizers==0.22.2
12
+ huggingface-hub==0.34.0
13
+ peft==0.17.0
14
+ numpy==1.23.1
15
+ Pillow==10.4.0
16
+ requests==2.32.3
17
+ tqdm==4.67.1
18
+ typing-extensions==4.15.0
19
+
20
+ # Optional FlashAttention 2 backend (validated: flash-attn==2.7.3). The CLI
21
+ # default is eager attention; --attn-implementation flash_attention_2 needs
22
+ # this package.
23
+ # flash-attn==2.7.3
code/rocm.patch ADDED
The diff for this file is too large to render. See raw diff
 
code/tests/assets/smoke_input.png ADDED
code/tokenization_bailing.py ADDED
@@ -0,0 +1,1024 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # coding=utf-8
3
+ # Copyright (c) Ant Group. All rights reserved.
4
+
5
+ import itertools
6
+ from typing import Any, Dict, List, Optional, Union
7
+
8
+ import torch
9
+ from transformers import PreTrainedTokenizerFast
10
+ from transformers.tokenization_utils_base import AddedToken, BatchEncoding
11
+ from transformers.utils import TensorType, logging
12
+
13
+ logger = logging.get_logger(__name__)
14
+
15
+
16
+ def is_system(msg):
17
+ return msg['role'].lower() == 'system'
18
+
19
+
20
+ def is_user(msg):
21
+ return msg['role'].lower() in ['human', 'user']
22
+
23
+
24
+ def is_assistant(msg):
25
+ return msg['role'].lower() == 'assistant'
26
+
27
+
28
+ def _convert_to_conversation(query, system=None):
29
+ conversation = []
30
+ if system:
31
+ conversation.append({"role": "SYSTEM", "content": system})
32
+ if isinstance(query, str):
33
+ conversation.append({"role": "HUMAN", "content": query})
34
+ elif isinstance(query, List):
35
+ conversation.extend(query)
36
+ elif isinstance(query, Dict):
37
+ if "messages" in query:
38
+ conversation.extend(query["messages"])
39
+ if "system_message" in query and len(conversation) > 0 and not is_system(conversation[0]):
40
+ conversation.insert(0, {"role": "SYSTEM", "content": query["system_message"]})
41
+ else:
42
+ conversation.append(query)
43
+ return conversation
44
+
45
+
46
+ class BailingTokenizer(PreTrainedTokenizerFast):
47
+ is_bailing_tokenizer = True
48
+ model_input_names = ["input_ids", "attention_mask"]
49
+ slow_tokenizer_class = None
50
+
51
+ # add gmask_token
52
+ SPECIAL_TOKENS_ATTRIBUTES = [
53
+ "bos_token",
54
+ "eos_token",
55
+ "unk_token",
56
+ "sep_token",
57
+ "pad_token",
58
+ "cls_token",
59
+ "mask_token",
60
+ "gmask_token",
61
+ "additional_special_tokens",
62
+ ]
63
+
64
+ def __init__(
65
+ self,
66
+ vocab_file=None,
67
+ merges_file=None,
68
+ tokenizer_file=None,
69
+ clean_up_tokenization_spaces=False,
70
+ bos_token="<|startoftext|>",
71
+ eos_token="<|endoftext|>",
72
+ cls_token="[CLS]",
73
+ pad_token="<|endoftext|>",
74
+ gmask_token="[gMASK]",
75
+ add_bos_token=False,
76
+ add_eos_token=False,
77
+ **kwargs,
78
+ ):
79
+ self.add_bos_token = add_bos_token
80
+
81
+ self._gmask_token = (
82
+ AddedToken(gmask_token, lstrip=False, rstrip=False, normalized=False)
83
+ if isinstance(gmask_token, str)
84
+ else gmask_token
85
+ )
86
+
87
+ self._sop_token = (
88
+ AddedToken(bos_token, lstrip=False, rstrip=False, normalized=False)
89
+ if isinstance(bos_token, str)
90
+ else bos_token
91
+ )
92
+
93
+ self._eop_token = (
94
+ AddedToken(eos_token, lstrip=False, rstrip=False, normalized=False)
95
+ if isinstance(eos_token, str)
96
+ else eos_token
97
+ )
98
+
99
+ super().__init__(
100
+ vocab_file=vocab_file,
101
+ merges_file=merges_file,
102
+ tokenizer_file=tokenizer_file,
103
+ clean_up_tokenization_spaces=clean_up_tokenization_spaces,
104
+ bos_token=bos_token,
105
+ eos_token=eos_token,
106
+ cls_token=cls_token,
107
+ pad_token=pad_token,
108
+ gmask_token=gmask_token,
109
+ add_bos_token=add_bos_token,
110
+ add_eos_token=add_eos_token,
111
+ **kwargs,
112
+ )
113
+
114
+ self.check_special_tokens()
115
+
116
+ def check_special_tokens(self):
117
+ '''
118
+ eos_token, cls_token, mask_token
119
+ special tokens should init, check special token is not None
120
+ '''
121
+ for name, special_token in zip(
122
+ ['eos', 'bos', 'cls', 'gmask'],
123
+ [self.eos_token, self.bos_token, self.cls_token, self.gmask_token],
124
+ ):
125
+ assert special_token is not None, f'should init special token [{name}] in tokenizer_config.json'
126
+
127
+ @property
128
+ def gmask_token(self) -> Optional[str]:
129
+ if self._gmask_token is None:
130
+ if self.verbose:
131
+ logger.error("Using gmask_token, but it is not set yet.")
132
+ return None
133
+ return str(self._gmask_token)
134
+
135
+ @gmask_token.setter
136
+ def gmask_token(self, value):
137
+ if not isinstance(value, (str, AddedToken)) and value is not None:
138
+ raise ValueError("Cannot set a non-string value as the gmask token")
139
+ self._gmask_token = value
140
+
141
+ @property
142
+ def gmask_token_id(self) -> Optional[int]:
143
+ if self._gmask_token is None:
144
+ return None
145
+ return self.convert_tokens_to_ids(self.gmask_token)
146
+
147
+ @property
148
+ def sop_token(self) -> Optional[str]:
149
+ if self._sop_token is None:
150
+ if self.verbose:
151
+ logger.error("Using sop_token, but it is not set yet.")
152
+ return None
153
+ return str(self._sop_token)
154
+
155
+ @sop_token.setter
156
+ def sop_token(self, value):
157
+ if not isinstance(value, (str, AddedToken)) and value is not None:
158
+ raise ValueError("Cannot set a non-string value as the sop token")
159
+ self._sop_token = value
160
+
161
+ @property
162
+ def sop_token_id(self) -> Optional[int]:
163
+ if self._sop_token is None:
164
+ return None
165
+ return self.convert_tokens_to_ids(self.sop_token)
166
+
167
+ @property
168
+ def eop_token(self) -> Optional[str]:
169
+ if self._eop_token is None:
170
+ if self.verbose:
171
+ logger.error("Using eop_token, but it is not set yet.")
172
+ return None
173
+ return str(self._eop_token)
174
+
175
+ @eop_token.setter
176
+ def eop_token(self, value):
177
+ if not isinstance(value, (str, AddedToken)) and value is not None:
178
+ raise ValueError("Cannot set a non-string value as the eop token")
179
+ self._eop_token = value
180
+
181
+ @property
182
+ def eop_token_id(self) -> Optional[int]:
183
+ if self._eop_token is None:
184
+ return None
185
+ return self.convert_tokens_to_ids(self.eop_token)
186
+
187
+ @property
188
+ def vocab_size(self):
189
+ return len(self.get_vocab())
190
+
191
+ def _chat_from_json(self, chat, chat_format="antglm_chat", system=None):
192
+ raise NotImplementedError(
193
+ "the legacy antglm_chat rendering path was removed; a checkpoint "
194
+ "tokenizer must provide a chat_template in tokenizer_config.json"
195
+ )
196
+
197
+ def apply_chat_template(
198
+ self,
199
+ conversation: Union[List[Dict[str, str]], List[List[Dict[str, str]]]],
200
+ tools: Optional[List[Dict]] = None,
201
+ documents: Optional[List[Dict[str, str]]] = None,
202
+ chat_template: Optional[str] = None,
203
+ add_generation_prompt: bool = False,
204
+ system: str = None, # only used for legacy chatml
205
+ tokenize=False,
206
+ padding: bool = False,
207
+ truncation: bool = False,
208
+ max_length: Optional[int] = None,
209
+ return_tensors: Optional[Union[str, TensorType]] = None,
210
+ return_dict: bool = False,
211
+ return_assistant_tokens_mask: bool = False,
212
+ tokenizer_kwargs: Optional[Dict[str, Any]] = None,
213
+ **kwargs,
214
+ ):
215
+ if hasattr(self, "chat_template") and self.chat_template:
216
+ if isinstance(conversation, Dict) and "messages" in conversation:
217
+ conversation = conversation["messages"]
218
+ # use transformers built-in method
219
+ return super().apply_chat_template(
220
+ conversation=conversation,
221
+ tools=tools,
222
+ documents=documents,
223
+ chat_template=chat_template,
224
+ add_generation_prompt=add_generation_prompt,
225
+ tokenize=tokenize,
226
+ padding=padding,
227
+ truncation=truncation,
228
+ return_tensors=return_tensors,
229
+ return_dict=return_dict,
230
+ return_assistant_tokens_mask=return_assistant_tokens_mask,
231
+ tokenizer_kwargs=tokenizer_kwargs,
232
+ )
233
+
234
+ # The legacy antglm_chat rendering path was removed. A checkpoint
235
+ # tokenizer must provide a chat_template in tokenizer_config.json.
236
+ raise ValueError(
237
+ "BailingMM2 inference requires a chat_template in the checkpoint "
238
+ "tokenizer_config.json"
239
+ )
240
+
241
+ def _build_position_ids(
242
+ self,
243
+ mask_pos: int,
244
+ bos_pos: int,
245
+ max_output_length: int,
246
+ rotary_type: Optional[str] = "none",
247
+ **kwargs,
248
+ ) -> List[List[int]]:
249
+ window_size = kwargs.get("window_size", 1024) - 1
250
+ block_position_ids = [0] * bos_pos
251
+
252
+ # location of the mask, used to construct output position ids later
253
+ if "1d" in rotary_type:
254
+ position_ids = list(range(bos_pos)) + list(range(mask_pos + 1, mask_pos + max_output_length + 2))
255
+ block_position_ids = block_position_ids + list(range(1, max_output_length + 2))
256
+ elif "2d" in rotary_type:
257
+ # a bos_id is appended to input_ids
258
+ position_ids = list(range(bos_pos))
259
+ position_ids = position_ids + [mask_pos] * (1 + max_output_length)
260
+ block_position_ids = block_position_ids + list(range(1, max_output_length + 2))
261
+ else:
262
+ # build position ids
263
+ position_ids = []
264
+ repeat_times = bos_pos // window_size
265
+ for _ in range(repeat_times):
266
+ position_ids += list(range(window_size))
267
+ position_ids += list(range(bos_pos - window_size * repeat_times))
268
+ # need consider additional bos_id after input_ids
269
+ mask_pos = position_ids[-1]
270
+ position_ids += [mask_pos] * (max_output_length + 1)
271
+
272
+ block_repeat_times = max_output_length // (window_size - 1)
273
+ additional_block_position_ids = []
274
+ for _ in range(block_repeat_times):
275
+ additional_block_position_ids += list(range(1, window_size))
276
+ additional_block_position_ids += list(
277
+ range(1, max_output_length + 2 - (window_size - 1) * block_repeat_times)
278
+ )
279
+ block_position_ids = block_position_ids + additional_block_position_ids
280
+
281
+ position_ids = [position_ids, block_position_ids]
282
+ return position_ids
283
+
284
+ def _build_inputs_for_generation(
285
+ self,
286
+ input_ids: List[int],
287
+ max_input_length=None,
288
+ left_truncate=True,
289
+ max_output_length=1024,
290
+ rotary_type="none",
291
+ unidirectional_attention: bool = True,
292
+ attention_dtype=None,
293
+ **kwargs,
294
+ ):
295
+ if max_input_length and len(input_ids) > max_input_length:
296
+ if left_truncate:
297
+ input_ids = input_ids[-max_input_length:]
298
+ else:
299
+ input_ids = input_ids[:max_input_length]
300
+
301
+ is_left_padding = input_ids[0] == self.eos_token_id
302
+ if not unidirectional_attention:
303
+ if input_ids[0] != self.cls_token_id:
304
+ input_ids = [self.cls_token_id] + input_ids
305
+
306
+ if self.gmask_token_id not in set(input_ids):
307
+ input_ids = input_ids + [self.gmask_token_id]
308
+
309
+ mask_pos = input_ids.index(self.gmask_token_id)
310
+ sep = len(input_ids)
311
+ else:
312
+ if self.add_bos_token:
313
+ input_ids = input_ids + [self.bos_token_id]
314
+ if self.eos_token_id in input_ids:
315
+ mask_pos = input_ids.index(self.eos_token_id) - 1
316
+ else:
317
+ mask_pos = len(input_ids) - 1
318
+ sep = len(input_ids) - 1
319
+ else:
320
+ sep = len(input_ids)
321
+ if self.eos_token_id in input_ids:
322
+ if is_left_padding:
323
+ ori_input_ids = input_ids
324
+ input_ids = input_ids[::-1]
325
+ mask_pos = input_ids.index(self.eos_token_id) - 1
326
+ mask_pos = max(0, mask_pos) # for empty sequence
327
+ if is_left_padding:
328
+ input_ids = ori_input_ids
329
+ mask_pos = sep - 1 - mask_pos # the first non-eos token
330
+
331
+ else:
332
+ mask_pos = len(input_ids) - 1
333
+
334
+ position_ids = self._build_position_ids(mask_pos, sep, max_output_length, rotary_type, **kwargs)
335
+
336
+ if is_left_padding:
337
+ position_ids[0] = [max(0, i - mask_pos) for i in range(len(position_ids[0]))]
338
+
339
+ # a bos_id is appended to input_ids
340
+ total_length = sep + max_output_length
341
+ if self.add_bos_token:
342
+ total_length += 1
343
+
344
+ def build_mask_matrix(seq_length, sep, mask_pos, unidirectional_attention):
345
+ # use bool attention masks for long sequences to save memory
346
+ if unidirectional_attention:
347
+ attention_mask = torch.ones([seq_length, seq_length], dtype=attention_dtype)
348
+ attention_mask = torch.tril(attention_mask)
349
+ if is_left_padding:
350
+ attention_mask[:, :mask_pos] = 0
351
+ else:
352
+ attention_mask[:, mask_pos + 1 : sep] = 0
353
+ else:
354
+ attention_mask = torch.zeros([seq_length, seq_length], dtype=attention_dtype)
355
+ attention_mask[:, : mask_pos + 1] = 1
356
+ for i in range(sep, total_length):
357
+ attention_mask[i, sep : i + 1] = 1
358
+ return attention_mask
359
+
360
+ if self.add_bos_token:
361
+ attention_mask = build_mask_matrix(total_length, sep + 1, mask_pos, unidirectional_attention)
362
+ else:
363
+ attention_mask = build_mask_matrix(total_length, sep, mask_pos, unidirectional_attention)
364
+ attention_mask = torch.unsqueeze(attention_mask, dim=0)
365
+ attention_mask = torch.unsqueeze(attention_mask, dim=1)
366
+ if attention_dtype is None:
367
+ attention_mask = attention_mask.long()
368
+ inputs = {
369
+ "input_ids": torch.Tensor([input_ids]).long(),
370
+ "position_ids": torch.Tensor([position_ids]).long(),
371
+ "attention_mask": attention_mask,
372
+ }
373
+ return BatchEncoding(inputs)
374
+
375
+ def build_inputs_for_generation(
376
+ self,
377
+ input_ids: Union[List[int], List[List[int]], torch.Tensor],
378
+ max_input_length=None,
379
+ left_truncate=True,
380
+ max_output_length=1024,
381
+ rotary_type="1d",
382
+ unidirectional_attention=True,
383
+ attention_dtype=None,
384
+ **kwargs,
385
+ ):
386
+ if isinstance(input_ids, torch.Tensor):
387
+ input_ids = input_ids.tolist()
388
+
389
+ if isinstance(input_ids[0], list):
390
+ input_ids_list = []
391
+ position_ids_list = []
392
+ attention_mask_list = []
393
+ for _input_ids in input_ids:
394
+ inputs = self._build_inputs_for_generation(
395
+ _input_ids,
396
+ max_input_length=max_input_length,
397
+ left_truncate=left_truncate,
398
+ max_output_length=max_output_length,
399
+ rotary_type=rotary_type,
400
+ unidirectional_attention=unidirectional_attention,
401
+ attention_dtype=attention_dtype,
402
+ **kwargs,
403
+ )
404
+ input_ids_list.append(inputs['input_ids'])
405
+ position_ids_list.append(inputs['position_ids'])
406
+ attention_mask_list.append(inputs["attention_mask"])
407
+
408
+ max_ids_length = max([input.size(1) for input in input_ids_list])
409
+
410
+ for i in range(len(input_ids)):
411
+ cur_ids_length = input_ids_list[i].size(1)
412
+ if cur_ids_length < max_ids_length:
413
+ # pad input ids
414
+ pad_input_ids = input_ids_list[i].new_zeros((1, max_ids_length - cur_ids_length))
415
+ input_ids_list[i] = torch.cat([pad_input_ids, input_ids_list[i]], dim=-1)
416
+
417
+ # pad postition ids with left pad
418
+ # 0, 1, 2, 3, 4 ... -> 0, ..., 0, 1, 2, 3, 4, ...
419
+ pad_position_ids = input_ids_list[i].new_zeros((1, 2, max_ids_length - cur_ids_length))
420
+ position_ids_list[i] = torch.cat([pad_position_ids, position_ids_list[i]], dim=-1)
421
+
422
+ # pad generation attention mask with left and bottom pad
423
+ new_attention_mask = input_ids_list[i].new_zeros(
424
+ 1,
425
+ 1,
426
+ max_ids_length + max_output_length,
427
+ max_ids_length + max_output_length,
428
+ )
429
+ new_attention_mask[
430
+ :,
431
+ :,
432
+ max_ids_length - cur_ids_length :,
433
+ max_ids_length - cur_ids_length :,
434
+ ] = attention_mask_list[i]
435
+ attention_mask_list[i] = new_attention_mask.contiguous()
436
+
437
+ input_ids_list = torch.cat(input_ids_list, dim=0)
438
+ position_ids_list = torch.cat(position_ids_list, dim=0)
439
+ attention_mask_list = torch.cat(attention_mask_list, dim=0)
440
+
441
+ inputs = {
442
+ "input_ids": input_ids_list,
443
+ "position_ids": position_ids_list,
444
+ "attention_mask": attention_mask_list,
445
+ }
446
+
447
+ return BatchEncoding(inputs)
448
+ else:
449
+ return self._build_inputs_for_generation(
450
+ input_ids,
451
+ max_input_length=max_input_length,
452
+ left_truncate=left_truncate,
453
+ max_output_length=max_output_length,
454
+ rotary_type=rotary_type,
455
+ unidirectional_attention=unidirectional_attention,
456
+ **kwargs,
457
+ )
458
+
459
+ def _build_inputs_for_train(
460
+ self,
461
+ inputs: Union[str, List[str]],
462
+ outputs: Union[str, List[str]],
463
+ new_conversation_offset: List[int] = None,
464
+ max_length: int = 2048,
465
+ rotary_type: str = "1d",
466
+ left_truncate: bool = True,
467
+ unidirectional_attention: bool = True,
468
+ isolation_position_ids: bool = False,
469
+ padding: bool = True,
470
+ use_fa2: bool = True,
471
+ use_packed: bool = True,
472
+ use_baichuan_packed: bool = False,
473
+ skip_truncated_turn: bool = False,
474
+ return_attention_mask: bool = True,
475
+ ):
476
+ r"""
477
+ Build tensor input for model training. If inputs and outputs are list, will pack them.
478
+
479
+ Args:
480
+ inputs (str, List[str], List[Dict], List[List[Dict]]): the input prompts.
481
+ outputs (str, List[str]): the output responses.
482
+ max_length (int, Optional): the maximum length of the final input ids for training. Default: 2048
483
+ rotary_type (str, Optional): the rotary type of position embedding. Default: 1d
484
+ left_truncate (bool, Optional): whether truncate the inputs from left. Default: True
485
+ use_fa2 (bool, Optional): whether to build attention mask under flash attention 2.
486
+ new_conversation_offset (List[int], Optional): marks turns that start a brand-new conversation; [0, 1] means inputs[0]/outputs[0] is one turn, inputs[1]/outputs[1] another.
487
+ """
488
+ if use_packed and use_baichuan_packed and unidirectional_attention:
489
+ return self._build_baichuan_inputs_for_train(
490
+ inputs,
491
+ outputs,
492
+ new_conversation_offset,
493
+ max_length,
494
+ rotary_type,
495
+ left_truncate,
496
+ skip_truncated_turn,
497
+ use_fa2,
498
+ padding,
499
+ )
500
+ if isinstance(inputs, str):
501
+ inputs = [inputs]
502
+ if isinstance(outputs, str):
503
+ outputs = [outputs]
504
+
505
+ assert len(inputs) == len(outputs)
506
+
507
+ input_ids = [self(item)['input_ids'] for item in inputs]
508
+ output_ids = [self(item)['input_ids'] for item in outputs]
509
+
510
+ packed_input_ids = []
511
+ packed_output_ids = []
512
+ if new_conversation_offset is None:
513
+ new_conversation_offset = list(range(0, len(inputs)))
514
+ assert 0 in new_conversation_offset, f"0 must be present, check new_conversation_offset: {new_conversation_offset}"
515
+ current_len = 0
516
+
517
+ for idx, (input, output) in enumerate(zip(input_ids, output_ids)):
518
+ num_special_tokens = 0
519
+ if not unidirectional_attention:
520
+ if idx in new_conversation_offset:
521
+ # cls and gmask
522
+ num_special_tokens += 2
523
+ else:
524
+ # only gmask
525
+ num_special_tokens += 1
526
+ else:
527
+ # sop and eos
528
+ if self.add_bos_token:
529
+ num_special_tokens += 2
530
+ else:
531
+ num_special_tokens += 1
532
+
533
+ # truncate
534
+ if len(input) + len(output) + current_len > max_length - num_special_tokens:
535
+ if not use_packed or use_fa2 and unidirectional_attention:
536
+ attention_mask = torch.tensor(0)
537
+ elif use_fa2:
538
+ attention_mask = -1 * torch.ones([2, max_length])
539
+ else:
540
+ attention_mask = torch.tril(torch.ones([max_length, max_length]))
541
+ # return an empty sample that does not participate in training
542
+ default_return = {
543
+ 'input_ids': (torch.ones(max_length) * self.eos_token_id).long(),
544
+ 'position_ids': torch.zeros(2, max_length).long(),
545
+ 'attention_mask': (attention_mask.long()),
546
+ 'labels': (torch.ones(max_length) * -100).long(),
547
+ }
548
+ # if no truncation is needed, return directly
549
+ if skip_truncated_turn:
550
+ if current_len == 0:
551
+ return default_return
552
+ else:
553
+ break
554
+ left_len = max_length - num_special_tokens - current_len
555
+ # truncate only the prompt when truncation is required
556
+ if left_len - len(output) > 0:
557
+ if left_truncate:
558
+ input = input[-(left_len - len(output)) :]
559
+ else:
560
+ input = input[: left_len - len(output)]
561
+ else:
562
+ # the response exceeds left_len, return directly
563
+ if current_len == 0:
564
+ return default_return
565
+ else:
566
+ break
567
+ if unidirectional_attention:
568
+ packed_input_ids.append(list(input))
569
+ else:
570
+ if num_special_tokens == 4:
571
+ packed_input_ids.append([self.cls_token_id] + list(input) + [self.gmask_token_id])
572
+ else:
573
+ packed_input_ids.append(list(input) + [self.gmask_token_id])
574
+
575
+ packed_output_ids.append(list(output) + [self.eos_token_id])
576
+ current_len += len(input) + len(output) + num_special_tokens
577
+
578
+ assert current_len <= max_length
579
+
580
+ if use_packed:
581
+ # pack mode
582
+ def build_mask_matrix(seq_length, sep):
583
+ # https://github.com/pytorch/pytorch/issues/101932, fix triu/tril bf16 support
584
+ m = torch.ones((1, seq_length, seq_length))
585
+ mask = torch.arange(1, m.shape[-1] + 1).reshape(1, -1, 1).to(m.device)
586
+ ids = torch.arange(1, m.shape[-1] + 1).reshape(1, 1, -1).expand(1, m.shape[-1], -1).to(m.device)
587
+ m = (ids <= mask).type_as(m)
588
+
589
+ m[0, :, : int(sep)] = 1
590
+ m = m.squeeze(0)
591
+ return m
592
+
593
+ tokens = []
594
+ attention_mask_list = []
595
+ input_length_list = []
596
+ position_id_list = []
597
+ block_position_id_list = []
598
+ for input, output in zip(packed_input_ids, packed_output_ids):
599
+ if self.add_bos_token:
600
+ data = input + [self.sop_token_id] + output
601
+ mask_pos = len(input) - 1
602
+ else:
603
+ data = input + output
604
+ mask_pos = len(input) - 2
605
+ if return_attention_mask:
606
+ if unidirectional_attention:
607
+ attention_mask = build_mask_matrix(len(data), 0)
608
+ else:
609
+ attention_mask = build_mask_matrix(len(data), len(input))
610
+ attention_mask = attention_mask.squeeze((0, 1))
611
+
612
+ attention_mask_list.append(attention_mask)
613
+ input_length_list.append(len(input))
614
+ tokens += data
615
+
616
+ sop_pos = mask_pos + 1
617
+ position_ids, block_position_ids = self._build_position_ids(
618
+ mask_pos=mask_pos, bos_pos=sop_pos, max_output_length=len(output), rotary_type=rotary_type
619
+ )
620
+
621
+ position_id_list.append(position_ids)
622
+ block_position_id_list.append(block_position_ids)
623
+
624
+ labels = []
625
+ for i in range(len(packed_input_ids)):
626
+ if self.add_bos_token:
627
+ labels += [-100] * len(packed_input_ids[i]) + packed_output_ids[i] + [-100]
628
+ else:
629
+ labels += [-100] * (len(packed_input_ids[i]) - 1) + packed_output_ids[i] + [-100]
630
+
631
+ total_len = 0
632
+ if use_fa2:
633
+ pack_attention_mask = -1 * torch.ones([2, current_len])
634
+ else:
635
+ pack_attention_mask = torch.tril(torch.ones([current_len, current_len]))
636
+
637
+ pack_position_ids = []
638
+ pack_block_position_ids = []
639
+ total_len = 0
640
+ max_index = 0
641
+ for i in range(len(position_id_list)):
642
+
643
+ if use_fa2:
644
+ pack_attention_mask[0][i] = total_len
645
+ pack_attention_mask[1][i] = total_len + input_length_list[i]
646
+ else:
647
+ pack_attention_mask[
648
+ total_len : total_len + attention_mask.shape[0],
649
+ total_len : total_len + attention_mask.shape[0],
650
+ ] = attention_mask
651
+ position_ids = [pid + max_index for pid in position_id_list[i]]
652
+ block_position_ids = block_position_id_list[i]
653
+ pack_position_ids.extend(position_ids)
654
+ pack_block_position_ids.extend(block_position_ids)
655
+ if not isolation_position_ids:
656
+ max_index = pack_position_ids[-1] + 1
657
+ total_len += len(position_id_list[i])
658
+ position_ids = [pack_position_ids, pack_block_position_ids]
659
+ else:
660
+ # single-input mode
661
+ # in true multi-turn, one sample can span several turns; find the end position of the first turn
662
+ if len(new_conversation_offset) > 1:
663
+ end_idx = new_conversation_offset[1]
664
+ else:
665
+ end_idx = 1
666
+ input, output = list(itertools.chain(*packed_input_ids[:end_idx])), list(
667
+ itertools.chain(*packed_output_ids[:end_idx])
668
+ )
669
+ if self.add_bos_token:
670
+ tokens = input + [self.sop_token_id] + output
671
+ else:
672
+ tokens = input + output
673
+
674
+ if self.add_bos_token:
675
+ labels = [-100] * len(input) + output + [-100]
676
+ position_ids = self._build_position_ids(
677
+ mask_pos=len(input) - 1, bos_pos=len(input), max_output_length=len(output), rotary_type=rotary_type
678
+ )
679
+ else:
680
+ labels = [-100] * (len(input) - 1) + output + [-100]
681
+ position_ids = self._build_position_ids(
682
+ mask_pos=len(input) - 2,
683
+ bos_pos=len(input) - 1,
684
+ max_output_length=len(output),
685
+ rotary_type=rotary_type,
686
+ )
687
+ attention_mask = len(input)
688
+ assert current_len == len(tokens)
689
+
690
+ # pad up to the maximum length
691
+ if max_length > 0 and len(tokens) < max_length and padding:
692
+ pad_length = max_length - len(tokens)
693
+ tokens += [self.pad_token_id] * pad_length
694
+ labels.extend([-100] * pad_length)
695
+ position_ids[0] += [0] * pad_length
696
+ position_ids[1] += [0] * pad_length
697
+
698
+ if use_packed:
699
+ if use_fa2:
700
+ new_attention_mask = -1 * torch.ones([2, max_length])
701
+ new_attention_mask[:, :current_len] = pack_attention_mask
702
+ else:
703
+ new_attention_mask = torch.tril(torch.ones([max_length, max_length]))
704
+ new_attention_mask[:current_len, :current_len] = pack_attention_mask
705
+ pack_attention_mask = new_attention_mask.contiguous()
706
+
707
+ assert len(tokens) == len(labels)
708
+
709
+ if max_length > 0 and padding:
710
+ assert len(tokens) == max_length
711
+
712
+ if use_fa2 and unidirectional_attention:
713
+ # pack_attention_mask = torch.zeros([1], dtype=torch.long)
714
+ pack_attention_mask = torch.tensor(0)
715
+
716
+ if use_packed:
717
+ if not use_fa2:
718
+ attention_mask = pack_attention_mask.unsqueeze(0).long()
719
+ else:
720
+ attention_mask = pack_attention_mask
721
+ else:
722
+ attention_mask = torch.tensor(attention_mask).long()
723
+ return {
724
+ 'input_ids': torch.tensor(tokens).long(),
725
+ 'position_ids': torch.tensor(position_ids).long(),
726
+ 'attention_mask': attention_mask,
727
+ 'labels': torch.tensor(labels).long(),
728
+ }
729
+
730
+ def _build_baichuan_inputs_for_train(
731
+ self,
732
+ inputs: Union[str, List[str]],
733
+ outputs: Union[str, List[str]],
734
+ new_conversation_offset: List[int] = None,
735
+ max_length: int = 2048,
736
+ rotary_type: str = "1d",
737
+ left_truncate: bool = True,
738
+ skip_truncated_turn: bool = True,
739
+ use_fa2: bool = True,
740
+ padding: bool = True,
741
+ ):
742
+ '''
743
+ input: <role> HUMAN </role> u1 <role> ASSISTANT </role> a11 a12 <role> HUMAN </role> u2 <role> ASSISTANT </role> a21 a22 <|endoftext|> <role> HUMAN </role> u1 <role> ASSISTANT </role> a11 a12 <role> HUMAN </role> u2 <role> ASSISTANT </role> a21 a22 <|endoftext|>
744
+ output: x x x x x x a11 a12 <|endoftext|> x x x x x x a21 a22 <|endoftext|> x x x x x x x a11 a12 <|endoftext|> x x x x x x a21 a22 <|endoftext|> x
745
+ Only for true multi-turn + pack training of a unidirectional model; requires use_true_multiturn=True
746
+ '''
747
+ if isinstance(inputs, str):
748
+ inputs = [inputs]
749
+ if isinstance(outputs, str):
750
+ outputs = [outputs]
751
+ assert len(inputs) == len(outputs)
752
+
753
+ input_ids = [self(item)['input_ids'] for item in inputs]
754
+ output_ids = [self(item)['input_ids'] for item in outputs]
755
+
756
+ packed_input_ids = []
757
+ packed_output_ids = []
758
+
759
+ if new_conversation_offset is None:
760
+ new_conversation_offset = list(range(0, len(inputs)))
761
+ assert 0 in new_conversation_offset, f"0 must be present, check new_conversation_offset: {new_conversation_offset}"
762
+ current_len = 0
763
+
764
+ for idx, (input, output) in enumerate(zip(input_ids, output_ids)):
765
+ num_special_tokens = 0
766
+ if idx != 0 and idx in new_conversation_offset:
767
+ # append eos to input_ids; skip it only for the first sample
768
+ num_special_tokens += 1
769
+
770
+ # truncate
771
+ if len(input) + len(output) + current_len > max_length - num_special_tokens:
772
+ if use_fa2:
773
+ attention_mask = torch.tensor(0)
774
+ else:
775
+ attention_mask = torch.tril(torch.ones([max_length, max_length]))
776
+ # return an empty sample that does not participate in training
777
+ default_return = {
778
+ 'input_ids': (torch.ones(max_length) * self.eos_token_id).long(),
779
+ 'position_ids': torch.zeros(2, max_length).long(),
780
+ 'attention_mask': (attention_mask.long()),
781
+ 'labels': (torch.ones(max_length) * -100).long(),
782
+ }
783
+
784
+ # if no truncation is needed, return directly
785
+ if skip_truncated_turn:
786
+ if current_len == 0:
787
+ return default_return
788
+ else:
789
+ break
790
+ left_len = max_length - num_special_tokens - current_len
791
+ # truncate only the prompt when truncation is required
792
+ if left_len - len(output) > 0:
793
+ if left_truncate:
794
+ input = input[-(left_len - len(output)) :]
795
+ else:
796
+ input = input[: left_len - len(output)]
797
+ else:
798
+ # the response exceeds left_len, return directly
799
+ if current_len == 0:
800
+ return default_return
801
+ else:
802
+ break
803
+ # input_ids are concatenated here
804
+ if num_special_tokens == 1:
805
+ packed_input_ids.append([self.eos_token_id] + list(input))
806
+ else:
807
+ packed_input_ids.append(list(input))
808
+ packed_output_ids.append(list(output))
809
+ current_len += len(input) + len(output) + num_special_tokens
810
+ assert current_len <= max_length
811
+
812
+ def build_mask_matrix(seq_length, sep):
813
+ # https://github.com/pytorch/pytorch/issues/101932, fix triu/tril bf16 support
814
+ m = torch.ones((1, seq_length, seq_length))
815
+ mask = torch.arange(1, m.shape[-1] + 1).reshape(1, -1, 1).to(m.device)
816
+ ids = torch.arange(1, m.shape[-1] + 1).reshape(1, 1, -1).expand(1, m.shape[-1], -1).to(m.device)
817
+ m = (ids <= mask).type_as(m)
818
+
819
+ m[0, :, : int(sep)] = 1
820
+ m = m.squeeze(0)
821
+ return m
822
+
823
+ tokens = []
824
+ attention_mask_list = []
825
+ position_id_list = []
826
+ block_position_id_list = []
827
+ token_lens = []
828
+ for input, output in zip(packed_input_ids, packed_output_ids):
829
+ data = input + output
830
+ if not use_fa2:
831
+ attention_mask = build_mask_matrix(len(data), 0)
832
+ attention_mask_list.append(attention_mask)
833
+ tokens += data
834
+ token_lens.append(len(data))
835
+
836
+ position_ids, block_position_ids = self._build_position_ids(
837
+ mask_pos=len(input) - 2, bos_pos=len(input) - 1, max_output_length=len(output), rotary_type=rotary_type
838
+ )
839
+
840
+ position_id_list.append(position_ids)
841
+ block_position_id_list.append(block_position_ids)
842
+
843
+ labels = []
844
+ for i in range(len(packed_input_ids)):
845
+ labels += [-100] * (len(packed_input_ids[i]) - 1) + packed_output_ids[i] + [self.eos_token_id]
846
+
847
+ total_len = 0
848
+ if use_fa2:
849
+ pack_attention_mask = torch.Tensor([[0], [1]])
850
+ else:
851
+ pack_attention_mask = torch.tril(torch.ones([max_length, max_length]))
852
+
853
+ pack_position_ids = []
854
+ pack_block_position_ids = []
855
+ total_len = 0
856
+ max_index = 0
857
+ for i in range(len(token_lens)):
858
+ if not use_fa2:
859
+ attention_mask = attention_mask_list[i]
860
+ pack_attention_mask[
861
+ total_len : total_len + attention_mask.shape[0], total_len : total_len + attention_mask.shape[0]
862
+ ] = attention_mask
863
+ position_ids = [pid + max_index for pid in position_id_list[i]]
864
+ block_position_ids = block_position_id_list[i]
865
+ pack_position_ids.extend(position_ids)
866
+ pack_block_position_ids.extend(block_position_ids)
867
+ max_index = pack_position_ids[-1] + 1
868
+ total_len += token_lens[i]
869
+ position_ids = [pack_position_ids, pack_block_position_ids]
870
+
871
+ if max_length > 0 and len(tokens) < max_length and padding:
872
+ pad_length = max_length - len(tokens)
873
+ tokens += [self.pad_token_id] * pad_length
874
+ labels.extend([-100] * pad_length)
875
+ position_ids[0] += [0] * pad_length
876
+ position_ids[1] += [0] * pad_length
877
+
878
+ assert len(tokens) == len(labels)
879
+
880
+ if not use_fa2:
881
+ attention_mask = pack_attention_mask.unsqueeze(0).long()
882
+ else:
883
+ attention_mask = torch.tensor(0)
884
+ return {
885
+ 'input_ids': torch.tensor(tokens).long(),
886
+ 'position_ids': torch.tensor(position_ids).long(),
887
+ 'attention_mask': attention_mask,
888
+ 'labels': torch.tensor(labels).long(),
889
+ }
890
+
891
+ def build_inputs_for_train(
892
+ self,
893
+ data: Union[Dict, List[Dict]],
894
+ new_conversation_offset: List[int] = None,
895
+ chat_format="antglm_chat",
896
+ is_chat_format=True, # whether the input is already formatted as chat text
897
+ use_true_multiturn=False,
898
+ max_length: int = 2048,
899
+ rotary_type: str = "1d",
900
+ left_truncate: bool = True,
901
+ unidirectional_attention: bool = True,
902
+ isolation_position_ids: bool = False,
903
+ padding: bool = True,
904
+ use_fa2: bool = True,
905
+ use_packed: bool = True,
906
+ use_baichuan_packed: bool = False,
907
+ skip_truncated_turn: bool = False,
908
+ return_attention_mask: bool = True,
909
+ ):
910
+ r"""
911
+ Build tensor input for model training. If inputs and outputs are list, will pack them.
912
+
913
+ Args:
914
+ inputs (str, List[str], List[Dict], List[List[Dict]]): the input prompts.
915
+ outputs (str, List[str]): the output responses.
916
+ new_conversation_offset (List[int]): the offset index of the new conversation turn.
917
+ is_chat_format (bool): whether the input is already chatml format
918
+ max_length (int, Optional): the maximum length of the final input ids for training. Default: 2048
919
+ rotary_type (str, Optional): the rotary type of position embedding. Default: 1d
920
+ left_truncate (bool, Optional): whether truncate the inputs from left. Default: True
921
+ use_fa2 (bool, Optional): whether to build attention mask under flash attention 2.
922
+ """
923
+ if isinstance(data, List):
924
+ # chatml list
925
+ _inputs = []
926
+ _outputs = []
927
+ new_conversation_offset = []
928
+ for _input in data:
929
+ if use_true_multiturn:
930
+ chat = self._chat_from_json(_input, chat_format=chat_format)
931
+ chat_data = chat.prompt_pack
932
+ new_conversation_offset.append(len(_inputs))
933
+ _inputs.extend(chat_data['input'])
934
+ _outputs.extend(chat_data['output'])
935
+ else:
936
+ _conversation = _convert_to_conversation(_input)
937
+ assert is_assistant(_conversation[-1])
938
+
939
+ _inputs.append(
940
+ self.apply_chat_template(_conversation[:-1], tokenize=False, add_generation_prompt=True)
941
+ )
942
+ _outputs.append(_conversation[-1]['content'])
943
+
944
+ return self._build_inputs_for_train(
945
+ inputs=_inputs,
946
+ outputs=_outputs,
947
+ new_conversation_offset=new_conversation_offset,
948
+ max_length=max_length,
949
+ rotary_type=rotary_type,
950
+ left_truncate=left_truncate,
951
+ unidirectional_attention=unidirectional_attention,
952
+ isolation_position_ids=isolation_position_ids,
953
+ padding=padding,
954
+ use_fa2=use_fa2,
955
+ use_packed=use_packed,
956
+ use_baichuan_packed=use_baichuan_packed,
957
+ skip_truncated_turn=skip_truncated_turn,
958
+ return_attention_mask=return_attention_mask,
959
+ )
960
+ elif isinstance(data, Dict):
961
+ if 'messages' in data:
962
+ # chatml format
963
+ if use_true_multiturn:
964
+ chat = self._chat_from_json(data, chat_format=chat_format)
965
+ chat_data = chat.prompt_pack
966
+ else:
967
+ _conversation = _convert_to_conversation(data)
968
+ assert is_assistant(_conversation[-1])
969
+
970
+ chat_data = {
971
+ "input": self.apply_chat_template(
972
+ _conversation[:-1], tokenize=False, add_generation_prompt=True
973
+ ),
974
+ "output": _conversation[-1]['content'],
975
+ }
976
+
977
+ return self._build_inputs_for_train(
978
+ inputs=chat_data['input'],
979
+ outputs=chat_data['output'],
980
+ max_length=max_length,
981
+ rotary_type=rotary_type,
982
+ left_truncate=left_truncate,
983
+ unidirectional_attention=unidirectional_attention,
984
+ isolation_position_ids=isolation_position_ids,
985
+ padding=padding,
986
+ use_fa2=use_fa2,
987
+ use_packed=use_packed,
988
+ use_baichuan_packed=use_baichuan_packed,
989
+ skip_truncated_turn=skip_truncated_turn,
990
+ return_attention_mask=return_attention_mask,
991
+ )
992
+ else:
993
+ inputs = data['input']
994
+ outputs = data['output']
995
+
996
+ if isinstance(inputs, str):
997
+ inputs = [inputs]
998
+ if isinstance(outputs, str):
999
+ outputs = [outputs]
1000
+
1001
+ if not is_chat_format and chat_format:
1002
+ inputs = [
1003
+ self.apply_chat_template(
1004
+ [{"role": "HUMAN", "content": item}], tokenize=False, chat_format=chat_format
1005
+ )
1006
+ for item in inputs
1007
+ ]
1008
+
1009
+ return self._build_inputs_for_train(
1010
+ inputs=inputs,
1011
+ outputs=outputs,
1012
+ new_conversation_offset=new_conversation_offset,
1013
+ max_length=max_length,
1014
+ rotary_type=rotary_type,
1015
+ left_truncate=left_truncate,
1016
+ unidirectional_attention=unidirectional_attention,
1017
+ isolation_position_ids=isolation_position_ids,
1018
+ padding=padding,
1019
+ use_fa2=use_fa2,
1020
+ use_packed=use_packed,
1021
+ use_baichuan_packed=use_baichuan_packed,
1022
+ skip_truncated_turn=skip_truncated_turn,
1023
+ return_attention_mask=return_attention_mask,
1024
+ )
connector/config.json ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "Qwen2ForCausalLM"
4
+ ],
5
+ "attention_dropout": 0.0,
6
+ "bos_token_id": 151643,
7
+ "dtype": "bfloat16",
8
+ "eos_token_id": 151645,
9
+ "hidden_act": "silu",
10
+ "hidden_size": 1536,
11
+ "initializer_range": 0.02,
12
+ "intermediate_size": 8960,
13
+ "layer_types": [
14
+ "full_attention",
15
+ "full_attention",
16
+ "full_attention",
17
+ "full_attention",
18
+ "full_attention",
19
+ "full_attention",
20
+ "full_attention",
21
+ "full_attention",
22
+ "full_attention",
23
+ "full_attention",
24
+ "full_attention",
25
+ "full_attention",
26
+ "full_attention",
27
+ "full_attention",
28
+ "full_attention",
29
+ "full_attention",
30
+ "full_attention",
31
+ "full_attention",
32
+ "full_attention",
33
+ "full_attention",
34
+ "full_attention",
35
+ "full_attention",
36
+ "full_attention",
37
+ "full_attention",
38
+ "full_attention",
39
+ "full_attention",
40
+ "full_attention",
41
+ "full_attention"
42
+ ],
43
+ "max_position_embeddings": 32768,
44
+ "max_window_layers": 21,
45
+ "model_type": "qwen2",
46
+ "num_attention_heads": 12,
47
+ "num_hidden_layers": 28,
48
+ "num_key_value_heads": 2,
49
+ "rms_norm_eps": 1e-06,
50
+ "rope_scaling": null,
51
+ "rope_theta": 1000000.0,
52
+ "sliding_window": null,
53
+ "tie_word_embeddings": true,
54
+ "transformers_version": "4.57.1",
55
+ "use_cache": true,
56
+ "use_sliding_window": false,
57
+ "vocab_size": 151936
58
+ }
connector/generation_config.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token_id": 151643,
3
+ "do_sample": true,
4
+ "eos_token_id": [
5
+ 151645,
6
+ 151643
7
+ ],
8
+ "pad_token_id": 151643,
9
+ "repetition_penalty": 1.1,
10
+ "temperature": 0.7,
11
+ "top_k": 20,
12
+ "top_p": 0.8,
13
+ "transformers_version": "4.57.1"
14
+ }
mllm/config.json ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "BailingMM2NativeForConditionalGeneration"
4
+ ],
5
+ "audio_config": null,
6
+ "llm_config": {
7
+ "_name_or_path": "",
8
+ "add_cross_attention": false,
9
+ "architectures": [
10
+ "BailingMoeV2ForCausalLM"
11
+ ],
12
+ "attention_dropout": 0.0,
13
+ "bad_words_ids": null,
14
+ "begin_suppress_tokens": null,
15
+ "bos_token_id": null,
16
+ "chunk_size_feed_forward": 0,
17
+ "cross_attention_hidden_size": null,
18
+ "decoder_start_token_id": null,
19
+ "diversity_penalty": 0.0,
20
+ "do_sample": false,
21
+ "early_stopping": false,
22
+ "embedding_dropout": 0.0,
23
+ "encoder_no_repeat_ngram_size": 0,
24
+ "eos_token_id": 156895,
25
+ "exponential_decay_length_penalty": null,
26
+ "finetuning_task": null,
27
+ "first_k_dense_replace": 1,
28
+ "forced_bos_token_id": null,
29
+ "forced_eos_token_id": null,
30
+ "head_dim": 128,
31
+ "hidden_act": "silu",
32
+ "hidden_size": 2048,
33
+ "id2label": {
34
+ "0": "LABEL_0",
35
+ "1": "LABEL_1"
36
+ },
37
+ "image_end_token": 157159,
38
+ "image_patch_token": 157157,
39
+ "image_start_token": 157158,
40
+ "video_end_token": 157161,
41
+ "video_patch_token": 157175,
42
+ "video_start_token": 157160,
43
+ "initializer_range": 0.006,
44
+ "intermediate_size": 5120,
45
+ "is_decoder": false,
46
+ "is_encoder_decoder": false,
47
+ "label2id": {
48
+ "LABEL_0": 0,
49
+ "LABEL_1": 1
50
+ },
51
+ "length_penalty": 1.0,
52
+ "max_length": 20,
53
+ "max_position_embeddings": 32768,
54
+ "max_window_layers": 28,
55
+ "min_length": 0,
56
+ "model_type": "bailing_moe_v2",
57
+ "moe_intermediate_size": 512,
58
+ "moe_router_topk_scaling_factor": 2.5,
59
+ "n_group": 8,
60
+ "no_repeat_ngram_size": 0,
61
+ "norm_head": false,
62
+ "norm_softmax": false,
63
+ "norm_topk_prob": true,
64
+ "num_attention_heads": 16,
65
+ "num_beam_groups": 1,
66
+ "num_beams": 1,
67
+ "num_experts": 256,
68
+ "num_experts_per_tok": 8,
69
+ "num_hidden_layers": 20,
70
+ "num_key_value_heads": 4,
71
+ "num_return_sequences": 1,
72
+ "num_shared_experts": 1,
73
+ "output_attentions": false,
74
+ "output_dropout": 0.0,
75
+ "output_hidden_states": false,
76
+ "output_router_logits": false,
77
+ "output_scores": false,
78
+ "pad_token_id": 156892,
79
+ "partial_rotary_factor": 0.5,
80
+ "prefix": null,
81
+ "pretraining_tp": 1,
82
+ "problem_type": null,
83
+ "pruned_heads": {},
84
+ "remove_invalid_values": false,
85
+ "repetition_penalty": 1.0,
86
+ "return_dict": true,
87
+ "return_dict_in_generate": false,
88
+ "rms_norm_eps": 1e-06,
89
+ "rope_scaling": {
90
+ "factor": null,
91
+ "type": "video_rope"
92
+ },
93
+ "rope_theta": 600000,
94
+ "routed_scaling_factor": 2.5,
95
+ "router_type": "MultiRouter",
96
+ "sep_token_id": null,
97
+ "sliding_window": 4096,
98
+ "spatial_merge_size": 2,
99
+ "suppress_tokens": null,
100
+ "task_specific_params": null,
101
+ "temperature": 1.0,
102
+ "tf_legacy_loss": false,
103
+ "tie_encoder_decoder": false,
104
+ "tie_word_embeddings": false,
105
+ "tokenizer_class": null,
106
+ "tokens_per_second": 2,
107
+ "top_k": 50,
108
+ "top_p": 1.0,
109
+ "topk_group": 4,
110
+ "torch_dtype": "bfloat16",
111
+ "torchscript": false,
112
+ "typical_p": 1.0,
113
+ "use_bfloat16": false,
114
+ "use_bias": false,
115
+ "use_cache": true,
116
+ "use_expert_bias": true,
117
+ "use_qkv_bias": false,
118
+ "use_sliding_window": false,
119
+ "vocab_size": 157184
120
+ },
121
+ "mlp_depth": 2,
122
+ "model_type": "bailingmm_moe_v2_lite",
123
+ "torch_dtype": "bfloat16",
124
+ "transformers_version": "4.53.0.dev0",
125
+ "vision_config": {
126
+ "_name_or_path": "",
127
+ "add_cross_attention": false,
128
+ "architectures": [
129
+ "Qwen2_5_VisionTransformer"
130
+ ],
131
+ "bad_words_ids": null,
132
+ "begin_suppress_tokens": null,
133
+ "bos_token_id": null,
134
+ "chunk_size_feed_forward": 0,
135
+ "cross_attention_hidden_size": null,
136
+ "decoder_start_token_id": null,
137
+ "depth": 32,
138
+ "diversity_penalty": 0.0,
139
+ "do_sample": false,
140
+ "early_stopping": false,
141
+ "encoder_no_repeat_ngram_size": 0,
142
+ "eos_token_id": null,
143
+ "exponential_decay_length_penalty": null,
144
+ "finetuning_task": null,
145
+ "forced_bos_token_id": null,
146
+ "forced_eos_token_id": null,
147
+ "fullatt_block_indexes": [
148
+ 7,
149
+ 15,
150
+ 23,
151
+ 31
152
+ ],
153
+ "hidden_act": "silu",
154
+ "hidden_size": 1280,
155
+ "id2label": {
156
+ "0": "LABEL_0",
157
+ "1": "LABEL_1"
158
+ },
159
+ "in_channels": 3,
160
+ "in_chans": 3,
161
+ "intermediate_size": 3456,
162
+ "is_decoder": false,
163
+ "is_encoder_decoder": false,
164
+ "label2id": {
165
+ "LABEL_0": 0,
166
+ "LABEL_1": 1
167
+ },
168
+ "length_penalty": 1.0,
169
+ "max_length": 20,
170
+ "min_length": 0,
171
+ "model_type": "qwen2_5_vit",
172
+ "no_repeat_ngram_size": 0,
173
+ "num_beam_groups": 1,
174
+ "num_beams": 1,
175
+ "num_heads": 16,
176
+ "num_return_sequences": 1,
177
+ "out_hidden_size": 8192,
178
+ "output_attentions": false,
179
+ "output_hidden_states": false,
180
+ "output_scores": false,
181
+ "pad_token_id": null,
182
+ "patch_size": 14,
183
+ "prefix": null,
184
+ "problem_type": null,
185
+ "pruned_heads": {},
186
+ "remove_invalid_values": false,
187
+ "repetition_penalty": 1.0,
188
+ "return_dict": true,
189
+ "return_dict_in_generate": false,
190
+ "sep_token_id": null,
191
+ "spatial_merge_size": 2,
192
+ "spatial_patch_size": 14,
193
+ "suppress_tokens": null,
194
+ "task_specific_params": null,
195
+ "temperature": 1.0,
196
+ "temporal_patch_size": 2,
197
+ "tf_legacy_loss": false,
198
+ "tie_encoder_decoder": false,
199
+ "tie_word_embeddings": true,
200
+ "tokenizer_class": null,
201
+ "tokens_per_second": 2,
202
+ "top_k": 50,
203
+ "top_p": 1.0,
204
+ "torch_dtype": "bfloat16",
205
+ "torchscript": false,
206
+ "typical_p": 1.0,
207
+ "use_bfloat16": false,
208
+ "window_size": 112
209
+ }
210
+ }
mllm/int8_manifest.json ADDED
The diff for this file is too large to render. See raw diff
 
mllm/model.safetensors.index.json ADDED
The diff for this file is too large to render. See raw diff
 
mllm/preprocessor_config.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "min_pixels": 451584,
3
+ "max_pixels": 451584,
4
+ "patch_size": 14,
5
+ "temporal_patch_size": 2,
6
+ "merge_size": 2,
7
+ "image_mean": [
8
+ 0.48145466,
9
+ 0.4578275,
10
+ 0.40821073
11
+ ],
12
+ "image_std": [
13
+ 0.26862954,
14
+ 0.26130258,
15
+ 0.27577711
16
+ ],
17
+ "image_token": "<image>",
18
+ "video_token": "<video>",
19
+ "image_processor_type": "BailingMM2ImageProcessor",
20
+ "return_attention_mask": true,
21
+ "padding_side": "right",
22
+ "padding_value": 0.0
23
+ }
mllm/special_tokens_map.json ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<|arithmetic_end|>",
4
+ "<role>",
5
+ "<|number_end|>",
6
+ "<|arithmetic_start|>",
7
+ "</role>",
8
+ "<|number_start|>",
9
+ "<|role_end|>",
10
+ "<tool_call>",
11
+ "</tool_call>",
12
+ "<tool_response>",
13
+ "</tool_response>"
14
+ ],
15
+ "bos_token": {
16
+ "content": "<|startoftext|>",
17
+ "lstrip": false,
18
+ "normalized": false,
19
+ "rstrip": false,
20
+ "single_word": false
21
+ },
22
+ "cls_token": {
23
+ "content": "[CLS]",
24
+ "lstrip": false,
25
+ "normalized": false,
26
+ "rstrip": false,
27
+ "single_word": false
28
+ },
29
+ "eos_token": {
30
+ "content": "<|role_end|>",
31
+ "lstrip": false,
32
+ "normalized": false,
33
+ "rstrip": false,
34
+ "single_word": false
35
+ },
36
+ "gmask_token": {
37
+ "content": "[gMASK]",
38
+ "lstrip": false,
39
+ "normalized": false,
40
+ "rstrip": false,
41
+ "single_word": false
42
+ },
43
+ "pad_token": "<|role_end|>"
44
+ }
mllm/tokenizer_config.json ADDED
@@ -0,0 +1,2334 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_bos_token": false,
3
+ "add_eos_token": false,
4
+ "added_tokens_decoder": {
5
+ "156891": {
6
+ "content": "<|startoftext|>",
7
+ "lstrip": false,
8
+ "normalized": false,
9
+ "rstrip": false,
10
+ "single_word": false,
11
+ "special": true
12
+ },
13
+ "156892": {
14
+ "content": "<|endoftext|>",
15
+ "lstrip": false,
16
+ "normalized": false,
17
+ "rstrip": false,
18
+ "single_word": false,
19
+ "special": true
20
+ },
21
+ "156893": {
22
+ "content": "[CLS]",
23
+ "lstrip": false,
24
+ "normalized": false,
25
+ "rstrip": false,
26
+ "single_word": false,
27
+ "special": true
28
+ },
29
+ "156894": {
30
+ "content": "[gMASK]",
31
+ "lstrip": false,
32
+ "normalized": false,
33
+ "rstrip": false,
34
+ "single_word": false,
35
+ "special": true
36
+ },
37
+ "156895": {
38
+ "content": "<|role_end|>",
39
+ "lstrip": false,
40
+ "normalized": false,
41
+ "rstrip": false,
42
+ "single_word": false,
43
+ "special": true
44
+ },
45
+ "156896": {
46
+ "content": "<tool_call>",
47
+ "lstrip": false,
48
+ "normalized": false,
49
+ "rstrip": false,
50
+ "single_word": false,
51
+ "special": true
52
+ },
53
+ "156897": {
54
+ "content": "</tool_call>",
55
+ "lstrip": false,
56
+ "normalized": false,
57
+ "rstrip": false,
58
+ "single_word": false,
59
+ "special": true
60
+ },
61
+ "156898": {
62
+ "content": "<tool_response>",
63
+ "lstrip": false,
64
+ "normalized": false,
65
+ "rstrip": false,
66
+ "single_word": false,
67
+ "special": true
68
+ },
69
+ "156899": {
70
+ "content": "</tool_response>",
71
+ "lstrip": false,
72
+ "normalized": false,
73
+ "rstrip": false,
74
+ "single_word": false,
75
+ "special": true
76
+ },
77
+ "156900": {
78
+ "content": "<|reserved_token_5|>",
79
+ "lstrip": false,
80
+ "normalized": false,
81
+ "rstrip": false,
82
+ "single_word": false,
83
+ "special": true
84
+ },
85
+ "156901": {
86
+ "content": "<|reserved_token_6|>",
87
+ "lstrip": false,
88
+ "normalized": false,
89
+ "rstrip": false,
90
+ "single_word": false,
91
+ "special": true
92
+ },
93
+ "156902": {
94
+ "content": "<|reserved_token_7|>",
95
+ "lstrip": false,
96
+ "normalized": false,
97
+ "rstrip": false,
98
+ "single_word": false,
99
+ "special": true
100
+ },
101
+ "156903": {
102
+ "content": "<|reserved_token_8|>",
103
+ "lstrip": false,
104
+ "normalized": false,
105
+ "rstrip": false,
106
+ "single_word": false,
107
+ "special": true
108
+ },
109
+ "156904": {
110
+ "content": "<|reserved_token_9|>",
111
+ "lstrip": false,
112
+ "normalized": false,
113
+ "rstrip": false,
114
+ "single_word": false,
115
+ "special": true
116
+ },
117
+ "156905": {
118
+ "content": "<|reserved_token_10|>",
119
+ "lstrip": false,
120
+ "normalized": false,
121
+ "rstrip": false,
122
+ "single_word": false,
123
+ "special": true
124
+ },
125
+ "156906": {
126
+ "content": "<|reserved_token_11|>",
127
+ "lstrip": false,
128
+ "normalized": false,
129
+ "rstrip": false,
130
+ "single_word": false,
131
+ "special": true
132
+ },
133
+ "156907": {
134
+ "content": "<|reserved_token_12|>",
135
+ "lstrip": false,
136
+ "normalized": false,
137
+ "rstrip": false,
138
+ "single_word": false,
139
+ "special": true
140
+ },
141
+ "156908": {
142
+ "content": "<|reserved_token_13|>",
143
+ "lstrip": false,
144
+ "normalized": false,
145
+ "rstrip": false,
146
+ "single_word": false,
147
+ "special": true
148
+ },
149
+ "156909": {
150
+ "content": "<|reserved_token_14|>",
151
+ "lstrip": false,
152
+ "normalized": false,
153
+ "rstrip": false,
154
+ "single_word": false,
155
+ "special": true
156
+ },
157
+ "156910": {
158
+ "content": "<|reserved_token_15|>",
159
+ "lstrip": false,
160
+ "normalized": false,
161
+ "rstrip": false,
162
+ "single_word": false,
163
+ "special": true
164
+ },
165
+ "156911": {
166
+ "content": "<|reserved_token_16|>",
167
+ "lstrip": false,
168
+ "normalized": false,
169
+ "rstrip": false,
170
+ "single_word": false,
171
+ "special": true
172
+ },
173
+ "156912": {
174
+ "content": "<|reserved_token_17|>",
175
+ "lstrip": false,
176
+ "normalized": false,
177
+ "rstrip": false,
178
+ "single_word": false,
179
+ "special": true
180
+ },
181
+ "156913": {
182
+ "content": "<|reserved_token_18|>",
183
+ "lstrip": false,
184
+ "normalized": false,
185
+ "rstrip": false,
186
+ "single_word": false,
187
+ "special": true
188
+ },
189
+ "156914": {
190
+ "content": "<|reserved_token_19|>",
191
+ "lstrip": false,
192
+ "normalized": false,
193
+ "rstrip": false,
194
+ "single_word": false,
195
+ "special": true
196
+ },
197
+ "156915": {
198
+ "content": "<|reserved_token_20|>",
199
+ "lstrip": false,
200
+ "normalized": false,
201
+ "rstrip": false,
202
+ "single_word": false,
203
+ "special": true
204
+ },
205
+ "156916": {
206
+ "content": "<|reserved_token_21|>",
207
+ "lstrip": false,
208
+ "normalized": false,
209
+ "rstrip": false,
210
+ "single_word": false,
211
+ "special": true
212
+ },
213
+ "156917": {
214
+ "content": "<|reserved_token_22|>",
215
+ "lstrip": false,
216
+ "normalized": false,
217
+ "rstrip": false,
218
+ "single_word": false,
219
+ "special": true
220
+ },
221
+ "156918": {
222
+ "content": "<|reserved_token_23|>",
223
+ "lstrip": false,
224
+ "normalized": false,
225
+ "rstrip": false,
226
+ "single_word": false,
227
+ "special": true
228
+ },
229
+ "156919": {
230
+ "content": "<|reserved_token_24|>",
231
+ "lstrip": false,
232
+ "normalized": false,
233
+ "rstrip": false,
234
+ "single_word": false,
235
+ "special": true
236
+ },
237
+ "156920": {
238
+ "content": "<|reserved_token_25|>",
239
+ "lstrip": false,
240
+ "normalized": false,
241
+ "rstrip": false,
242
+ "single_word": false,
243
+ "special": true
244
+ },
245
+ "156921": {
246
+ "content": "<|reserved_token_26|>",
247
+ "lstrip": false,
248
+ "normalized": false,
249
+ "rstrip": false,
250
+ "single_word": false,
251
+ "special": true
252
+ },
253
+ "156922": {
254
+ "content": "<|reserved_token_27|>",
255
+ "lstrip": false,
256
+ "normalized": false,
257
+ "rstrip": false,
258
+ "single_word": false,
259
+ "special": true
260
+ },
261
+ "156923": {
262
+ "content": "<|reserved_token_28|>",
263
+ "lstrip": false,
264
+ "normalized": false,
265
+ "rstrip": false,
266
+ "single_word": false,
267
+ "special": true
268
+ },
269
+ "156924": {
270
+ "content": "<|reserved_token_29|>",
271
+ "lstrip": false,
272
+ "normalized": false,
273
+ "rstrip": false,
274
+ "single_word": false,
275
+ "special": true
276
+ },
277
+ "156925": {
278
+ "content": "<|reserved_token_30|>",
279
+ "lstrip": false,
280
+ "normalized": false,
281
+ "rstrip": false,
282
+ "single_word": false,
283
+ "special": true
284
+ },
285
+ "156926": {
286
+ "content": "<|reserved_token_31|>",
287
+ "lstrip": false,
288
+ "normalized": false,
289
+ "rstrip": false,
290
+ "single_word": false,
291
+ "special": true
292
+ },
293
+ "156927": {
294
+ "content": "<|reserved_token_32|>",
295
+ "lstrip": false,
296
+ "normalized": false,
297
+ "rstrip": false,
298
+ "single_word": false,
299
+ "special": true
300
+ },
301
+ "156928": {
302
+ "content": "<|reserved_token_33|>",
303
+ "lstrip": false,
304
+ "normalized": false,
305
+ "rstrip": false,
306
+ "single_word": false,
307
+ "special": true
308
+ },
309
+ "156929": {
310
+ "content": "<|reserved_token_34|>",
311
+ "lstrip": false,
312
+ "normalized": false,
313
+ "rstrip": false,
314
+ "single_word": false,
315
+ "special": true
316
+ },
317
+ "156930": {
318
+ "content": "<|reserved_token_35|>",
319
+ "lstrip": false,
320
+ "normalized": false,
321
+ "rstrip": false,
322
+ "single_word": false,
323
+ "special": true
324
+ },
325
+ "156931": {
326
+ "content": "<|reserved_token_36|>",
327
+ "lstrip": false,
328
+ "normalized": false,
329
+ "rstrip": false,
330
+ "single_word": false,
331
+ "special": true
332
+ },
333
+ "156932": {
334
+ "content": "<|reserved_token_37|>",
335
+ "lstrip": false,
336
+ "normalized": false,
337
+ "rstrip": false,
338
+ "single_word": false,
339
+ "special": true
340
+ },
341
+ "156933": {
342
+ "content": "<|reserved_token_38|>",
343
+ "lstrip": false,
344
+ "normalized": false,
345
+ "rstrip": false,
346
+ "single_word": false,
347
+ "special": true
348
+ },
349
+ "156934": {
350
+ "content": "<|reserved_token_39|>",
351
+ "lstrip": false,
352
+ "normalized": false,
353
+ "rstrip": false,
354
+ "single_word": false,
355
+ "special": true
356
+ },
357
+ "156935": {
358
+ "content": "<|reserved_token_40|>",
359
+ "lstrip": false,
360
+ "normalized": false,
361
+ "rstrip": false,
362
+ "single_word": false,
363
+ "special": true
364
+ },
365
+ "156936": {
366
+ "content": "<|reserved_token_41|>",
367
+ "lstrip": false,
368
+ "normalized": false,
369
+ "rstrip": false,
370
+ "single_word": false,
371
+ "special": true
372
+ },
373
+ "156937": {
374
+ "content": "<|reserved_token_42|>",
375
+ "lstrip": false,
376
+ "normalized": false,
377
+ "rstrip": false,
378
+ "single_word": false,
379
+ "special": true
380
+ },
381
+ "156938": {
382
+ "content": "<|reserved_token_43|>",
383
+ "lstrip": false,
384
+ "normalized": false,
385
+ "rstrip": false,
386
+ "single_word": false,
387
+ "special": true
388
+ },
389
+ "156939": {
390
+ "content": "<|reserved_token_44|>",
391
+ "lstrip": false,
392
+ "normalized": false,
393
+ "rstrip": false,
394
+ "single_word": false,
395
+ "special": true
396
+ },
397
+ "156940": {
398
+ "content": "<|reserved_token_45|>",
399
+ "lstrip": false,
400
+ "normalized": false,
401
+ "rstrip": false,
402
+ "single_word": false,
403
+ "special": true
404
+ },
405
+ "156941": {
406
+ "content": "<|reserved_token_46|>",
407
+ "lstrip": false,
408
+ "normalized": false,
409
+ "rstrip": false,
410
+ "single_word": false,
411
+ "special": true
412
+ },
413
+ "156942": {
414
+ "content": "<|reserved_token_47|>",
415
+ "lstrip": false,
416
+ "normalized": false,
417
+ "rstrip": false,
418
+ "single_word": false,
419
+ "special": true
420
+ },
421
+ "156943": {
422
+ "content": "<|reserved_token_48|>",
423
+ "lstrip": false,
424
+ "normalized": false,
425
+ "rstrip": false,
426
+ "single_word": false,
427
+ "special": true
428
+ },
429
+ "156944": {
430
+ "content": "<|reserved_token_49|>",
431
+ "lstrip": false,
432
+ "normalized": false,
433
+ "rstrip": false,
434
+ "single_word": false,
435
+ "special": true
436
+ },
437
+ "156945": {
438
+ "content": "<|reserved_token_50|>",
439
+ "lstrip": false,
440
+ "normalized": false,
441
+ "rstrip": false,
442
+ "single_word": false,
443
+ "special": true
444
+ },
445
+ "156946": {
446
+ "content": "<|reserved_token_51|>",
447
+ "lstrip": false,
448
+ "normalized": false,
449
+ "rstrip": false,
450
+ "single_word": false,
451
+ "special": true
452
+ },
453
+ "156947": {
454
+ "content": "<|reserved_token_52|>",
455
+ "lstrip": false,
456
+ "normalized": false,
457
+ "rstrip": false,
458
+ "single_word": false,
459
+ "special": true
460
+ },
461
+ "156948": {
462
+ "content": "<|reserved_token_53|>",
463
+ "lstrip": false,
464
+ "normalized": false,
465
+ "rstrip": false,
466
+ "single_word": false,
467
+ "special": true
468
+ },
469
+ "156949": {
470
+ "content": "<|reserved_token_54|>",
471
+ "lstrip": false,
472
+ "normalized": false,
473
+ "rstrip": false,
474
+ "single_word": false,
475
+ "special": true
476
+ },
477
+ "156950": {
478
+ "content": "<|reserved_token_55|>",
479
+ "lstrip": false,
480
+ "normalized": false,
481
+ "rstrip": false,
482
+ "single_word": false,
483
+ "special": true
484
+ },
485
+ "156951": {
486
+ "content": "<|reserved_token_56|>",
487
+ "lstrip": false,
488
+ "normalized": false,
489
+ "rstrip": false,
490
+ "single_word": false,
491
+ "special": true
492
+ },
493
+ "156952": {
494
+ "content": "<|reserved_token_57|>",
495
+ "lstrip": false,
496
+ "normalized": false,
497
+ "rstrip": false,
498
+ "single_word": false,
499
+ "special": true
500
+ },
501
+ "156953": {
502
+ "content": "<|reserved_token_58|>",
503
+ "lstrip": false,
504
+ "normalized": false,
505
+ "rstrip": false,
506
+ "single_word": false,
507
+ "special": true
508
+ },
509
+ "156954": {
510
+ "content": "<|reserved_token_59|>",
511
+ "lstrip": false,
512
+ "normalized": false,
513
+ "rstrip": false,
514
+ "single_word": false,
515
+ "special": true
516
+ },
517
+ "156955": {
518
+ "content": "<|reserved_token_60|>",
519
+ "lstrip": false,
520
+ "normalized": false,
521
+ "rstrip": false,
522
+ "single_word": false,
523
+ "special": true
524
+ },
525
+ "156956": {
526
+ "content": "<|reserved_token_61|>",
527
+ "lstrip": false,
528
+ "normalized": false,
529
+ "rstrip": false,
530
+ "single_word": false,
531
+ "special": true
532
+ },
533
+ "156957": {
534
+ "content": "<|reserved_token_62|>",
535
+ "lstrip": false,
536
+ "normalized": false,
537
+ "rstrip": false,
538
+ "single_word": false,
539
+ "special": true
540
+ },
541
+ "156958": {
542
+ "content": "<|reserved_token_63|>",
543
+ "lstrip": false,
544
+ "normalized": false,
545
+ "rstrip": false,
546
+ "single_word": false,
547
+ "special": true
548
+ },
549
+ "156959": {
550
+ "content": "<|reserved_token_64|>",
551
+ "lstrip": false,
552
+ "normalized": false,
553
+ "rstrip": false,
554
+ "single_word": false,
555
+ "special": true
556
+ },
557
+ "156960": {
558
+ "content": "<|reserved_token_65|>",
559
+ "lstrip": false,
560
+ "normalized": false,
561
+ "rstrip": false,
562
+ "single_word": false,
563
+ "special": true
564
+ },
565
+ "156961": {
566
+ "content": "<|reserved_token_66|>",
567
+ "lstrip": false,
568
+ "normalized": false,
569
+ "rstrip": false,
570
+ "single_word": false,
571
+ "special": true
572
+ },
573
+ "156962": {
574
+ "content": "<|reserved_token_67|>",
575
+ "lstrip": false,
576
+ "normalized": false,
577
+ "rstrip": false,
578
+ "single_word": false,
579
+ "special": true
580
+ },
581
+ "156963": {
582
+ "content": "<|reserved_token_68|>",
583
+ "lstrip": false,
584
+ "normalized": false,
585
+ "rstrip": false,
586
+ "single_word": false,
587
+ "special": true
588
+ },
589
+ "156964": {
590
+ "content": "<|reserved_token_69|>",
591
+ "lstrip": false,
592
+ "normalized": false,
593
+ "rstrip": false,
594
+ "single_word": false,
595
+ "special": true
596
+ },
597
+ "156965": {
598
+ "content": "<|reserved_token_70|>",
599
+ "lstrip": false,
600
+ "normalized": false,
601
+ "rstrip": false,
602
+ "single_word": false,
603
+ "special": true
604
+ },
605
+ "156966": {
606
+ "content": "<|reserved_token_71|>",
607
+ "lstrip": false,
608
+ "normalized": false,
609
+ "rstrip": false,
610
+ "single_word": false,
611
+ "special": true
612
+ },
613
+ "156967": {
614
+ "content": "<|reserved_token_72|>",
615
+ "lstrip": false,
616
+ "normalized": false,
617
+ "rstrip": false,
618
+ "single_word": false,
619
+ "special": true
620
+ },
621
+ "156968": {
622
+ "content": "<|reserved_token_73|>",
623
+ "lstrip": false,
624
+ "normalized": false,
625
+ "rstrip": false,
626
+ "single_word": false,
627
+ "special": true
628
+ },
629
+ "156969": {
630
+ "content": "<|reserved_token_74|>",
631
+ "lstrip": false,
632
+ "normalized": false,
633
+ "rstrip": false,
634
+ "single_word": false,
635
+ "special": true
636
+ },
637
+ "156970": {
638
+ "content": "<|reserved_token_75|>",
639
+ "lstrip": false,
640
+ "normalized": false,
641
+ "rstrip": false,
642
+ "single_word": false,
643
+ "special": true
644
+ },
645
+ "156971": {
646
+ "content": "<|reserved_token_76|>",
647
+ "lstrip": false,
648
+ "normalized": false,
649
+ "rstrip": false,
650
+ "single_word": false,
651
+ "special": true
652
+ },
653
+ "156972": {
654
+ "content": "<|reserved_token_77|>",
655
+ "lstrip": false,
656
+ "normalized": false,
657
+ "rstrip": false,
658
+ "single_word": false,
659
+ "special": true
660
+ },
661
+ "156973": {
662
+ "content": "<|reserved_token_78|>",
663
+ "lstrip": false,
664
+ "normalized": false,
665
+ "rstrip": false,
666
+ "single_word": false,
667
+ "special": true
668
+ },
669
+ "156974": {
670
+ "content": "<|reserved_token_79|>",
671
+ "lstrip": false,
672
+ "normalized": false,
673
+ "rstrip": false,
674
+ "single_word": false,
675
+ "special": true
676
+ },
677
+ "156975": {
678
+ "content": "<|reserved_token_80|>",
679
+ "lstrip": false,
680
+ "normalized": false,
681
+ "rstrip": false,
682
+ "single_word": false,
683
+ "special": true
684
+ },
685
+ "156976": {
686
+ "content": "<|reserved_token_81|>",
687
+ "lstrip": false,
688
+ "normalized": false,
689
+ "rstrip": false,
690
+ "single_word": false,
691
+ "special": true
692
+ },
693
+ "156977": {
694
+ "content": "<|reserved_token_82|>",
695
+ "lstrip": false,
696
+ "normalized": false,
697
+ "rstrip": false,
698
+ "single_word": false,
699
+ "special": true
700
+ },
701
+ "156978": {
702
+ "content": "<|reserved_token_83|>",
703
+ "lstrip": false,
704
+ "normalized": false,
705
+ "rstrip": false,
706
+ "single_word": false,
707
+ "special": true
708
+ },
709
+ "156979": {
710
+ "content": "<|reserved_token_84|>",
711
+ "lstrip": false,
712
+ "normalized": false,
713
+ "rstrip": false,
714
+ "single_word": false,
715
+ "special": true
716
+ },
717
+ "156980": {
718
+ "content": "<|reserved_token_85|>",
719
+ "lstrip": false,
720
+ "normalized": false,
721
+ "rstrip": false,
722
+ "single_word": false,
723
+ "special": true
724
+ },
725
+ "156981": {
726
+ "content": "<|reserved_token_86|>",
727
+ "lstrip": false,
728
+ "normalized": false,
729
+ "rstrip": false,
730
+ "single_word": false,
731
+ "special": true
732
+ },
733
+ "156982": {
734
+ "content": "<|reserved_token_87|>",
735
+ "lstrip": false,
736
+ "normalized": false,
737
+ "rstrip": false,
738
+ "single_word": false,
739
+ "special": true
740
+ },
741
+ "156983": {
742
+ "content": "<|reserved_token_88|>",
743
+ "lstrip": false,
744
+ "normalized": false,
745
+ "rstrip": false,
746
+ "single_word": false,
747
+ "special": true
748
+ },
749
+ "156984": {
750
+ "content": "<|reserved_token_89|>",
751
+ "lstrip": false,
752
+ "normalized": false,
753
+ "rstrip": false,
754
+ "single_word": false,
755
+ "special": true
756
+ },
757
+ "156985": {
758
+ "content": "<|reserved_token_90|>",
759
+ "lstrip": false,
760
+ "normalized": false,
761
+ "rstrip": false,
762
+ "single_word": false,
763
+ "special": true
764
+ },
765
+ "156986": {
766
+ "content": "<|reserved_token_91|>",
767
+ "lstrip": false,
768
+ "normalized": false,
769
+ "rstrip": false,
770
+ "single_word": false,
771
+ "special": true
772
+ },
773
+ "156987": {
774
+ "content": "<|reserved_token_92|>",
775
+ "lstrip": false,
776
+ "normalized": false,
777
+ "rstrip": false,
778
+ "single_word": false,
779
+ "special": true
780
+ },
781
+ "156988": {
782
+ "content": "<|reserved_token_93|>",
783
+ "lstrip": false,
784
+ "normalized": false,
785
+ "rstrip": false,
786
+ "single_word": false,
787
+ "special": true
788
+ },
789
+ "156989": {
790
+ "content": "<|reserved_token_94|>",
791
+ "lstrip": false,
792
+ "normalized": false,
793
+ "rstrip": false,
794
+ "single_word": false,
795
+ "special": true
796
+ },
797
+ "156990": {
798
+ "content": "<|reserved_token_95|>",
799
+ "lstrip": false,
800
+ "normalized": false,
801
+ "rstrip": false,
802
+ "single_word": false,
803
+ "special": true
804
+ },
805
+ "156991": {
806
+ "content": "<|reserved_token_96|>",
807
+ "lstrip": false,
808
+ "normalized": false,
809
+ "rstrip": false,
810
+ "single_word": false,
811
+ "special": true
812
+ },
813
+ "156992": {
814
+ "content": "<|reserved_token_97|>",
815
+ "lstrip": false,
816
+ "normalized": false,
817
+ "rstrip": false,
818
+ "single_word": false,
819
+ "special": true
820
+ },
821
+ "156993": {
822
+ "content": "<|reserved_token_98|>",
823
+ "lstrip": false,
824
+ "normalized": false,
825
+ "rstrip": false,
826
+ "single_word": false,
827
+ "special": true
828
+ },
829
+ "156994": {
830
+ "content": "<|reserved_token_99|>",
831
+ "lstrip": false,
832
+ "normalized": false,
833
+ "rstrip": false,
834
+ "single_word": false,
835
+ "special": true
836
+ },
837
+ "156995": {
838
+ "content": "<|reserved_token_100|>",
839
+ "lstrip": false,
840
+ "normalized": false,
841
+ "rstrip": false,
842
+ "single_word": false,
843
+ "special": true
844
+ },
845
+ "156996": {
846
+ "content": "<|reserved_token_101|>",
847
+ "lstrip": false,
848
+ "normalized": false,
849
+ "rstrip": false,
850
+ "single_word": false,
851
+ "special": true
852
+ },
853
+ "156997": {
854
+ "content": "<|reserved_token_102|>",
855
+ "lstrip": false,
856
+ "normalized": false,
857
+ "rstrip": false,
858
+ "single_word": false,
859
+ "special": true
860
+ },
861
+ "156998": {
862
+ "content": "<|reserved_token_103|>",
863
+ "lstrip": false,
864
+ "normalized": false,
865
+ "rstrip": false,
866
+ "single_word": false,
867
+ "special": true
868
+ },
869
+ "156999": {
870
+ "content": "<|reserved_token_104|>",
871
+ "lstrip": false,
872
+ "normalized": false,
873
+ "rstrip": false,
874
+ "single_word": false,
875
+ "special": true
876
+ },
877
+ "157000": {
878
+ "content": "<|reserved_token_105|>",
879
+ "lstrip": false,
880
+ "normalized": false,
881
+ "rstrip": false,
882
+ "single_word": false,
883
+ "special": true
884
+ },
885
+ "157001": {
886
+ "content": "<|reserved_token_106|>",
887
+ "lstrip": false,
888
+ "normalized": false,
889
+ "rstrip": false,
890
+ "single_word": false,
891
+ "special": true
892
+ },
893
+ "157002": {
894
+ "content": "<|reserved_token_107|>",
895
+ "lstrip": false,
896
+ "normalized": false,
897
+ "rstrip": false,
898
+ "single_word": false,
899
+ "special": true
900
+ },
901
+ "157003": {
902
+ "content": "<|reserved_token_108|>",
903
+ "lstrip": false,
904
+ "normalized": false,
905
+ "rstrip": false,
906
+ "single_word": false,
907
+ "special": true
908
+ },
909
+ "157004": {
910
+ "content": "<|reserved_token_109|>",
911
+ "lstrip": false,
912
+ "normalized": false,
913
+ "rstrip": false,
914
+ "single_word": false,
915
+ "special": true
916
+ },
917
+ "157005": {
918
+ "content": "<|reserved_token_110|>",
919
+ "lstrip": false,
920
+ "normalized": false,
921
+ "rstrip": false,
922
+ "single_word": false,
923
+ "special": true
924
+ },
925
+ "157006": {
926
+ "content": "<|reserved_token_111|>",
927
+ "lstrip": false,
928
+ "normalized": false,
929
+ "rstrip": false,
930
+ "single_word": false,
931
+ "special": true
932
+ },
933
+ "157007": {
934
+ "content": "<|reserved_token_112|>",
935
+ "lstrip": false,
936
+ "normalized": false,
937
+ "rstrip": false,
938
+ "single_word": false,
939
+ "special": true
940
+ },
941
+ "157008": {
942
+ "content": "<|reserved_token_113|>",
943
+ "lstrip": false,
944
+ "normalized": false,
945
+ "rstrip": false,
946
+ "single_word": false,
947
+ "special": true
948
+ },
949
+ "157009": {
950
+ "content": "<|reserved_token_114|>",
951
+ "lstrip": false,
952
+ "normalized": false,
953
+ "rstrip": false,
954
+ "single_word": false,
955
+ "special": true
956
+ },
957
+ "157010": {
958
+ "content": "<|reserved_token_115|>",
959
+ "lstrip": false,
960
+ "normalized": false,
961
+ "rstrip": false,
962
+ "single_word": false,
963
+ "special": true
964
+ },
965
+ "157011": {
966
+ "content": "<|reserved_token_116|>",
967
+ "lstrip": false,
968
+ "normalized": false,
969
+ "rstrip": false,
970
+ "single_word": false,
971
+ "special": true
972
+ },
973
+ "157012": {
974
+ "content": "<|reserved_token_117|>",
975
+ "lstrip": false,
976
+ "normalized": false,
977
+ "rstrip": false,
978
+ "single_word": false,
979
+ "special": true
980
+ },
981
+ "157013": {
982
+ "content": "<|reserved_token_118|>",
983
+ "lstrip": false,
984
+ "normalized": false,
985
+ "rstrip": false,
986
+ "single_word": false,
987
+ "special": true
988
+ },
989
+ "157014": {
990
+ "content": "<|reserved_token_119|>",
991
+ "lstrip": false,
992
+ "normalized": false,
993
+ "rstrip": false,
994
+ "single_word": false,
995
+ "special": true
996
+ },
997
+ "157015": {
998
+ "content": "<|reserved_token_120|>",
999
+ "lstrip": false,
1000
+ "normalized": false,
1001
+ "rstrip": false,
1002
+ "single_word": false,
1003
+ "special": true
1004
+ },
1005
+ "157016": {
1006
+ "content": "<|reserved_token_121|>",
1007
+ "lstrip": false,
1008
+ "normalized": false,
1009
+ "rstrip": false,
1010
+ "single_word": false,
1011
+ "special": true
1012
+ },
1013
+ "157017": {
1014
+ "content": "<|reserved_token_122|>",
1015
+ "lstrip": false,
1016
+ "normalized": false,
1017
+ "rstrip": false,
1018
+ "single_word": false,
1019
+ "special": true
1020
+ },
1021
+ "157018": {
1022
+ "content": "<|reserved_token_123|>",
1023
+ "lstrip": false,
1024
+ "normalized": false,
1025
+ "rstrip": false,
1026
+ "single_word": false,
1027
+ "special": true
1028
+ },
1029
+ "157019": {
1030
+ "content": "<|reserved_token_124|>",
1031
+ "lstrip": false,
1032
+ "normalized": false,
1033
+ "rstrip": false,
1034
+ "single_word": false,
1035
+ "special": true
1036
+ },
1037
+ "157020": {
1038
+ "content": "<|reserved_token_125|>",
1039
+ "lstrip": false,
1040
+ "normalized": false,
1041
+ "rstrip": false,
1042
+ "single_word": false,
1043
+ "special": true
1044
+ },
1045
+ "157021": {
1046
+ "content": "<|reserved_token_126|>",
1047
+ "lstrip": false,
1048
+ "normalized": false,
1049
+ "rstrip": false,
1050
+ "single_word": false,
1051
+ "special": true
1052
+ },
1053
+ "157022": {
1054
+ "content": "<|reserved_token_127|>",
1055
+ "lstrip": false,
1056
+ "normalized": false,
1057
+ "rstrip": false,
1058
+ "single_word": false,
1059
+ "special": true
1060
+ },
1061
+ "157023": {
1062
+ "content": "<|reserved_token_128|>",
1063
+ "lstrip": false,
1064
+ "normalized": false,
1065
+ "rstrip": false,
1066
+ "single_word": false,
1067
+ "special": true
1068
+ },
1069
+ "157024": {
1070
+ "content": "<|reserved_token_129|>",
1071
+ "lstrip": false,
1072
+ "normalized": false,
1073
+ "rstrip": false,
1074
+ "single_word": false,
1075
+ "special": true
1076
+ },
1077
+ "157025": {
1078
+ "content": "<|reserved_token_130|>",
1079
+ "lstrip": false,
1080
+ "normalized": false,
1081
+ "rstrip": false,
1082
+ "single_word": false,
1083
+ "special": true
1084
+ },
1085
+ "157026": {
1086
+ "content": "<|reserved_token_131|>",
1087
+ "lstrip": false,
1088
+ "normalized": false,
1089
+ "rstrip": false,
1090
+ "single_word": false,
1091
+ "special": true
1092
+ },
1093
+ "157027": {
1094
+ "content": "<|reserved_token_132|>",
1095
+ "lstrip": false,
1096
+ "normalized": false,
1097
+ "rstrip": false,
1098
+ "single_word": false,
1099
+ "special": true
1100
+ },
1101
+ "157028": {
1102
+ "content": "<|reserved_token_133|>",
1103
+ "lstrip": false,
1104
+ "normalized": false,
1105
+ "rstrip": false,
1106
+ "single_word": false,
1107
+ "special": true
1108
+ },
1109
+ "157029": {
1110
+ "content": "<|reserved_token_134|>",
1111
+ "lstrip": false,
1112
+ "normalized": false,
1113
+ "rstrip": false,
1114
+ "single_word": false,
1115
+ "special": true
1116
+ },
1117
+ "157030": {
1118
+ "content": "<|reserved_token_135|>",
1119
+ "lstrip": false,
1120
+ "normalized": false,
1121
+ "rstrip": false,
1122
+ "single_word": false,
1123
+ "special": true
1124
+ },
1125
+ "157031": {
1126
+ "content": "<|reserved_token_136|>",
1127
+ "lstrip": false,
1128
+ "normalized": false,
1129
+ "rstrip": false,
1130
+ "single_word": false,
1131
+ "special": true
1132
+ },
1133
+ "157032": {
1134
+ "content": "<|reserved_token_137|>",
1135
+ "lstrip": false,
1136
+ "normalized": false,
1137
+ "rstrip": false,
1138
+ "single_word": false,
1139
+ "special": true
1140
+ },
1141
+ "157033": {
1142
+ "content": "<|reserved_token_138|>",
1143
+ "lstrip": false,
1144
+ "normalized": false,
1145
+ "rstrip": false,
1146
+ "single_word": false,
1147
+ "special": true
1148
+ },
1149
+ "157034": {
1150
+ "content": "<|reserved_token_139|>",
1151
+ "lstrip": false,
1152
+ "normalized": false,
1153
+ "rstrip": false,
1154
+ "single_word": false,
1155
+ "special": true
1156
+ },
1157
+ "157035": {
1158
+ "content": "<|reserved_token_140|>",
1159
+ "lstrip": false,
1160
+ "normalized": false,
1161
+ "rstrip": false,
1162
+ "single_word": false,
1163
+ "special": true
1164
+ },
1165
+ "157036": {
1166
+ "content": "<|reserved_token_141|>",
1167
+ "lstrip": false,
1168
+ "normalized": false,
1169
+ "rstrip": false,
1170
+ "single_word": false,
1171
+ "special": true
1172
+ },
1173
+ "157037": {
1174
+ "content": "<|reserved_token_142|>",
1175
+ "lstrip": false,
1176
+ "normalized": false,
1177
+ "rstrip": false,
1178
+ "single_word": false,
1179
+ "special": true
1180
+ },
1181
+ "157038": {
1182
+ "content": "<|reserved_token_143|>",
1183
+ "lstrip": false,
1184
+ "normalized": false,
1185
+ "rstrip": false,
1186
+ "single_word": false,
1187
+ "special": true
1188
+ },
1189
+ "157039": {
1190
+ "content": "<|reserved_token_144|>",
1191
+ "lstrip": false,
1192
+ "normalized": false,
1193
+ "rstrip": false,
1194
+ "single_word": false,
1195
+ "special": true
1196
+ },
1197
+ "157040": {
1198
+ "content": "<|reserved_token_145|>",
1199
+ "lstrip": false,
1200
+ "normalized": false,
1201
+ "rstrip": false,
1202
+ "single_word": false,
1203
+ "special": true
1204
+ },
1205
+ "157041": {
1206
+ "content": "<|reserved_token_146|>",
1207
+ "lstrip": false,
1208
+ "normalized": false,
1209
+ "rstrip": false,
1210
+ "single_word": false,
1211
+ "special": true
1212
+ },
1213
+ "157042": {
1214
+ "content": "<|reserved_token_147|>",
1215
+ "lstrip": false,
1216
+ "normalized": false,
1217
+ "rstrip": false,
1218
+ "single_word": false,
1219
+ "special": true
1220
+ },
1221
+ "157043": {
1222
+ "content": "<|reserved_token_148|>",
1223
+ "lstrip": false,
1224
+ "normalized": false,
1225
+ "rstrip": false,
1226
+ "single_word": false,
1227
+ "special": true
1228
+ },
1229
+ "157044": {
1230
+ "content": "<|reserved_token_149|>",
1231
+ "lstrip": false,
1232
+ "normalized": false,
1233
+ "rstrip": false,
1234
+ "single_word": false,
1235
+ "special": true
1236
+ },
1237
+ "157045": {
1238
+ "content": "<|reserved_token_150|>",
1239
+ "lstrip": false,
1240
+ "normalized": false,
1241
+ "rstrip": false,
1242
+ "single_word": false,
1243
+ "special": true
1244
+ },
1245
+ "157046": {
1246
+ "content": "<|reserved_token_151|>",
1247
+ "lstrip": false,
1248
+ "normalized": false,
1249
+ "rstrip": false,
1250
+ "single_word": false,
1251
+ "special": true
1252
+ },
1253
+ "157047": {
1254
+ "content": "<|reserved_token_152|>",
1255
+ "lstrip": false,
1256
+ "normalized": false,
1257
+ "rstrip": false,
1258
+ "single_word": false,
1259
+ "special": true
1260
+ },
1261
+ "157048": {
1262
+ "content": "<|reserved_token_153|>",
1263
+ "lstrip": false,
1264
+ "normalized": false,
1265
+ "rstrip": false,
1266
+ "single_word": false,
1267
+ "special": true
1268
+ },
1269
+ "157049": {
1270
+ "content": "<|reserved_token_154|>",
1271
+ "lstrip": false,
1272
+ "normalized": false,
1273
+ "rstrip": false,
1274
+ "single_word": false,
1275
+ "special": true
1276
+ },
1277
+ "157050": {
1278
+ "content": "<|reserved_token_155|>",
1279
+ "lstrip": false,
1280
+ "normalized": false,
1281
+ "rstrip": false,
1282
+ "single_word": false,
1283
+ "special": true
1284
+ },
1285
+ "157051": {
1286
+ "content": "<|reserved_token_156|>",
1287
+ "lstrip": false,
1288
+ "normalized": false,
1289
+ "rstrip": false,
1290
+ "single_word": false,
1291
+ "special": true
1292
+ },
1293
+ "157052": {
1294
+ "content": "<|reserved_token_157|>",
1295
+ "lstrip": false,
1296
+ "normalized": false,
1297
+ "rstrip": false,
1298
+ "single_word": false,
1299
+ "special": true
1300
+ },
1301
+ "157053": {
1302
+ "content": "<|reserved_token_158|>",
1303
+ "lstrip": false,
1304
+ "normalized": false,
1305
+ "rstrip": false,
1306
+ "single_word": false,
1307
+ "special": true
1308
+ },
1309
+ "157054": {
1310
+ "content": "<|reserved_token_159|>",
1311
+ "lstrip": false,
1312
+ "normalized": false,
1313
+ "rstrip": false,
1314
+ "single_word": false,
1315
+ "special": true
1316
+ },
1317
+ "157055": {
1318
+ "content": "<|reserved_token_160|>",
1319
+ "lstrip": false,
1320
+ "normalized": false,
1321
+ "rstrip": false,
1322
+ "single_word": false,
1323
+ "special": true
1324
+ },
1325
+ "157056": {
1326
+ "content": "<|reserved_token_161|>",
1327
+ "lstrip": false,
1328
+ "normalized": false,
1329
+ "rstrip": false,
1330
+ "single_word": false,
1331
+ "special": true
1332
+ },
1333
+ "157057": {
1334
+ "content": "<|reserved_token_162|>",
1335
+ "lstrip": false,
1336
+ "normalized": false,
1337
+ "rstrip": false,
1338
+ "single_word": false,
1339
+ "special": true
1340
+ },
1341
+ "157058": {
1342
+ "content": "<|reserved_token_163|>",
1343
+ "lstrip": false,
1344
+ "normalized": false,
1345
+ "rstrip": false,
1346
+ "single_word": false,
1347
+ "special": true
1348
+ },
1349
+ "157059": {
1350
+ "content": "<|reserved_token_164|>",
1351
+ "lstrip": false,
1352
+ "normalized": false,
1353
+ "rstrip": false,
1354
+ "single_word": false,
1355
+ "special": true
1356
+ },
1357
+ "157060": {
1358
+ "content": "<|reserved_token_165|>",
1359
+ "lstrip": false,
1360
+ "normalized": false,
1361
+ "rstrip": false,
1362
+ "single_word": false,
1363
+ "special": true
1364
+ },
1365
+ "157061": {
1366
+ "content": "<|reserved_token_166|>",
1367
+ "lstrip": false,
1368
+ "normalized": false,
1369
+ "rstrip": false,
1370
+ "single_word": false,
1371
+ "special": true
1372
+ },
1373
+ "157062": {
1374
+ "content": "<|reserved_token_167|>",
1375
+ "lstrip": false,
1376
+ "normalized": false,
1377
+ "rstrip": false,
1378
+ "single_word": false,
1379
+ "special": true
1380
+ },
1381
+ "157063": {
1382
+ "content": "<|reserved_token_168|>",
1383
+ "lstrip": false,
1384
+ "normalized": false,
1385
+ "rstrip": false,
1386
+ "single_word": false,
1387
+ "special": true
1388
+ },
1389
+ "157064": {
1390
+ "content": "<|reserved_token_169|>",
1391
+ "lstrip": false,
1392
+ "normalized": false,
1393
+ "rstrip": false,
1394
+ "single_word": false,
1395
+ "special": true
1396
+ },
1397
+ "157065": {
1398
+ "content": "<|reserved_token_170|>",
1399
+ "lstrip": false,
1400
+ "normalized": false,
1401
+ "rstrip": false,
1402
+ "single_word": false,
1403
+ "special": true
1404
+ },
1405
+ "157066": {
1406
+ "content": "<|reserved_token_171|>",
1407
+ "lstrip": false,
1408
+ "normalized": false,
1409
+ "rstrip": false,
1410
+ "single_word": false,
1411
+ "special": true
1412
+ },
1413
+ "157067": {
1414
+ "content": "<|reserved_token_172|>",
1415
+ "lstrip": false,
1416
+ "normalized": false,
1417
+ "rstrip": false,
1418
+ "single_word": false,
1419
+ "special": true
1420
+ },
1421
+ "157068": {
1422
+ "content": "<|reserved_token_173|>",
1423
+ "lstrip": false,
1424
+ "normalized": false,
1425
+ "rstrip": false,
1426
+ "single_word": false,
1427
+ "special": true
1428
+ },
1429
+ "157069": {
1430
+ "content": "<|reserved_token_174|>",
1431
+ "lstrip": false,
1432
+ "normalized": false,
1433
+ "rstrip": false,
1434
+ "single_word": false,
1435
+ "special": true
1436
+ },
1437
+ "157070": {
1438
+ "content": "<|reserved_token_175|>",
1439
+ "lstrip": false,
1440
+ "normalized": false,
1441
+ "rstrip": false,
1442
+ "single_word": false,
1443
+ "special": true
1444
+ },
1445
+ "157071": {
1446
+ "content": "<|reserved_token_176|>",
1447
+ "lstrip": false,
1448
+ "normalized": false,
1449
+ "rstrip": false,
1450
+ "single_word": false,
1451
+ "special": true
1452
+ },
1453
+ "157072": {
1454
+ "content": "<|reserved_token_177|>",
1455
+ "lstrip": false,
1456
+ "normalized": false,
1457
+ "rstrip": false,
1458
+ "single_word": false,
1459
+ "special": true
1460
+ },
1461
+ "157073": {
1462
+ "content": "<|reserved_token_178|>",
1463
+ "lstrip": false,
1464
+ "normalized": false,
1465
+ "rstrip": false,
1466
+ "single_word": false,
1467
+ "special": true
1468
+ },
1469
+ "157074": {
1470
+ "content": "<|reserved_token_179|>",
1471
+ "lstrip": false,
1472
+ "normalized": false,
1473
+ "rstrip": false,
1474
+ "single_word": false,
1475
+ "special": true
1476
+ },
1477
+ "157075": {
1478
+ "content": "<|reserved_token_180|>",
1479
+ "lstrip": false,
1480
+ "normalized": false,
1481
+ "rstrip": false,
1482
+ "single_word": false,
1483
+ "special": true
1484
+ },
1485
+ "157076": {
1486
+ "content": "<|reserved_token_181|>",
1487
+ "lstrip": false,
1488
+ "normalized": false,
1489
+ "rstrip": false,
1490
+ "single_word": false,
1491
+ "special": true
1492
+ },
1493
+ "157077": {
1494
+ "content": "<|reserved_token_182|>",
1495
+ "lstrip": false,
1496
+ "normalized": false,
1497
+ "rstrip": false,
1498
+ "single_word": false,
1499
+ "special": true
1500
+ },
1501
+ "157078": {
1502
+ "content": "<|reserved_token_183|>",
1503
+ "lstrip": false,
1504
+ "normalized": false,
1505
+ "rstrip": false,
1506
+ "single_word": false,
1507
+ "special": true
1508
+ },
1509
+ "157079": {
1510
+ "content": "<|reserved_token_184|>",
1511
+ "lstrip": false,
1512
+ "normalized": false,
1513
+ "rstrip": false,
1514
+ "single_word": false,
1515
+ "special": true
1516
+ },
1517
+ "157080": {
1518
+ "content": "<|reserved_token_185|>",
1519
+ "lstrip": false,
1520
+ "normalized": false,
1521
+ "rstrip": false,
1522
+ "single_word": false,
1523
+ "special": true
1524
+ },
1525
+ "157081": {
1526
+ "content": "<|reserved_token_186|>",
1527
+ "lstrip": false,
1528
+ "normalized": false,
1529
+ "rstrip": false,
1530
+ "single_word": false,
1531
+ "special": true
1532
+ },
1533
+ "157082": {
1534
+ "content": "<|reserved_token_187|>",
1535
+ "lstrip": false,
1536
+ "normalized": false,
1537
+ "rstrip": false,
1538
+ "single_word": false,
1539
+ "special": true
1540
+ },
1541
+ "157083": {
1542
+ "content": "<|reserved_token_188|>",
1543
+ "lstrip": false,
1544
+ "normalized": false,
1545
+ "rstrip": false,
1546
+ "single_word": false,
1547
+ "special": true
1548
+ },
1549
+ "157084": {
1550
+ "content": "<|reserved_token_189|>",
1551
+ "lstrip": false,
1552
+ "normalized": false,
1553
+ "rstrip": false,
1554
+ "single_word": false,
1555
+ "special": true
1556
+ },
1557
+ "157085": {
1558
+ "content": "<|reserved_token_190|>",
1559
+ "lstrip": false,
1560
+ "normalized": false,
1561
+ "rstrip": false,
1562
+ "single_word": false,
1563
+ "special": true
1564
+ },
1565
+ "157086": {
1566
+ "content": "<|reserved_token_191|>",
1567
+ "lstrip": false,
1568
+ "normalized": false,
1569
+ "rstrip": false,
1570
+ "single_word": false,
1571
+ "special": true
1572
+ },
1573
+ "157087": {
1574
+ "content": "<|reserved_token_192|>",
1575
+ "lstrip": false,
1576
+ "normalized": false,
1577
+ "rstrip": false,
1578
+ "single_word": false,
1579
+ "special": true
1580
+ },
1581
+ "157088": {
1582
+ "content": "<|reserved_token_193|>",
1583
+ "lstrip": false,
1584
+ "normalized": false,
1585
+ "rstrip": false,
1586
+ "single_word": false,
1587
+ "special": true
1588
+ },
1589
+ "157089": {
1590
+ "content": "<|reserved_token_194|>",
1591
+ "lstrip": false,
1592
+ "normalized": false,
1593
+ "rstrip": false,
1594
+ "single_word": false,
1595
+ "special": true
1596
+ },
1597
+ "157090": {
1598
+ "content": "<|reserved_token_195|>",
1599
+ "lstrip": false,
1600
+ "normalized": false,
1601
+ "rstrip": false,
1602
+ "single_word": false,
1603
+ "special": true
1604
+ },
1605
+ "157091": {
1606
+ "content": "<|reserved_token_196|>",
1607
+ "lstrip": false,
1608
+ "normalized": false,
1609
+ "rstrip": false,
1610
+ "single_word": false,
1611
+ "special": true
1612
+ },
1613
+ "157092": {
1614
+ "content": "<|reserved_token_197|>",
1615
+ "lstrip": false,
1616
+ "normalized": false,
1617
+ "rstrip": false,
1618
+ "single_word": false,
1619
+ "special": true
1620
+ },
1621
+ "157093": {
1622
+ "content": "<|reserved_token_198|>",
1623
+ "lstrip": false,
1624
+ "normalized": false,
1625
+ "rstrip": false,
1626
+ "single_word": false,
1627
+ "special": true
1628
+ },
1629
+ "157094": {
1630
+ "content": "<|reserved_token_199|>",
1631
+ "lstrip": false,
1632
+ "normalized": false,
1633
+ "rstrip": false,
1634
+ "single_word": false,
1635
+ "special": true
1636
+ },
1637
+ "157095": {
1638
+ "content": "<|reserved_token_200|>",
1639
+ "lstrip": false,
1640
+ "normalized": false,
1641
+ "rstrip": false,
1642
+ "single_word": false,
1643
+ "special": true
1644
+ },
1645
+ "157096": {
1646
+ "content": "<|reserved_token_201|>",
1647
+ "lstrip": false,
1648
+ "normalized": false,
1649
+ "rstrip": false,
1650
+ "single_word": false,
1651
+ "special": true
1652
+ },
1653
+ "157097": {
1654
+ "content": "<|reserved_token_202|>",
1655
+ "lstrip": false,
1656
+ "normalized": false,
1657
+ "rstrip": false,
1658
+ "single_word": false,
1659
+ "special": true
1660
+ },
1661
+ "157098": {
1662
+ "content": "<|reserved_token_203|>",
1663
+ "lstrip": false,
1664
+ "normalized": false,
1665
+ "rstrip": false,
1666
+ "single_word": false,
1667
+ "special": true
1668
+ },
1669
+ "157099": {
1670
+ "content": "<|reserved_token_204|>",
1671
+ "lstrip": false,
1672
+ "normalized": false,
1673
+ "rstrip": false,
1674
+ "single_word": false,
1675
+ "special": true
1676
+ },
1677
+ "157100": {
1678
+ "content": "<|reserved_token_205|>",
1679
+ "lstrip": false,
1680
+ "normalized": false,
1681
+ "rstrip": false,
1682
+ "single_word": false,
1683
+ "special": true
1684
+ },
1685
+ "157101": {
1686
+ "content": "<|reserved_token_206|>",
1687
+ "lstrip": false,
1688
+ "normalized": false,
1689
+ "rstrip": false,
1690
+ "single_word": false,
1691
+ "special": true
1692
+ },
1693
+ "157102": {
1694
+ "content": "<|reserved_token_207|>",
1695
+ "lstrip": false,
1696
+ "normalized": false,
1697
+ "rstrip": false,
1698
+ "single_word": false,
1699
+ "special": true
1700
+ },
1701
+ "157103": {
1702
+ "content": "<|reserved_token_208|>",
1703
+ "lstrip": false,
1704
+ "normalized": false,
1705
+ "rstrip": false,
1706
+ "single_word": false,
1707
+ "special": true
1708
+ },
1709
+ "157104": {
1710
+ "content": "<|reserved_token_209|>",
1711
+ "lstrip": false,
1712
+ "normalized": false,
1713
+ "rstrip": false,
1714
+ "single_word": false,
1715
+ "special": true
1716
+ },
1717
+ "157105": {
1718
+ "content": "<|reserved_token_210|>",
1719
+ "lstrip": false,
1720
+ "normalized": false,
1721
+ "rstrip": false,
1722
+ "single_word": false,
1723
+ "special": true
1724
+ },
1725
+ "157106": {
1726
+ "content": "<|reserved_token_211|>",
1727
+ "lstrip": false,
1728
+ "normalized": false,
1729
+ "rstrip": false,
1730
+ "single_word": false,
1731
+ "special": true
1732
+ },
1733
+ "157107": {
1734
+ "content": "<|reserved_token_212|>",
1735
+ "lstrip": false,
1736
+ "normalized": false,
1737
+ "rstrip": false,
1738
+ "single_word": false,
1739
+ "special": true
1740
+ },
1741
+ "157108": {
1742
+ "content": "<|reserved_token_213|>",
1743
+ "lstrip": false,
1744
+ "normalized": false,
1745
+ "rstrip": false,
1746
+ "single_word": false,
1747
+ "special": true
1748
+ },
1749
+ "157109": {
1750
+ "content": "<|reserved_token_214|>",
1751
+ "lstrip": false,
1752
+ "normalized": false,
1753
+ "rstrip": false,
1754
+ "single_word": false,
1755
+ "special": true
1756
+ },
1757
+ "157110": {
1758
+ "content": "<|reserved_token_215|>",
1759
+ "lstrip": false,
1760
+ "normalized": false,
1761
+ "rstrip": false,
1762
+ "single_word": false,
1763
+ "special": true
1764
+ },
1765
+ "157111": {
1766
+ "content": "<|reserved_token_216|>",
1767
+ "lstrip": false,
1768
+ "normalized": false,
1769
+ "rstrip": false,
1770
+ "single_word": false,
1771
+ "special": true
1772
+ },
1773
+ "157112": {
1774
+ "content": "<|reserved_token_217|>",
1775
+ "lstrip": false,
1776
+ "normalized": false,
1777
+ "rstrip": false,
1778
+ "single_word": false,
1779
+ "special": true
1780
+ },
1781
+ "157113": {
1782
+ "content": "<|reserved_token_218|>",
1783
+ "lstrip": false,
1784
+ "normalized": false,
1785
+ "rstrip": false,
1786
+ "single_word": false,
1787
+ "special": true
1788
+ },
1789
+ "157114": {
1790
+ "content": "<|reserved_token_219|>",
1791
+ "lstrip": false,
1792
+ "normalized": false,
1793
+ "rstrip": false,
1794
+ "single_word": false,
1795
+ "special": true
1796
+ },
1797
+ "157115": {
1798
+ "content": "<|reserved_token_220|>",
1799
+ "lstrip": false,
1800
+ "normalized": false,
1801
+ "rstrip": false,
1802
+ "single_word": false,
1803
+ "special": true
1804
+ },
1805
+ "157116": {
1806
+ "content": "<|reserved_token_221|>",
1807
+ "lstrip": false,
1808
+ "normalized": false,
1809
+ "rstrip": false,
1810
+ "single_word": false,
1811
+ "special": true
1812
+ },
1813
+ "157117": {
1814
+ "content": "<|reserved_token_222|>",
1815
+ "lstrip": false,
1816
+ "normalized": false,
1817
+ "rstrip": false,
1818
+ "single_word": false,
1819
+ "special": true
1820
+ },
1821
+ "157118": {
1822
+ "content": "<|reserved_token_223|>",
1823
+ "lstrip": false,
1824
+ "normalized": false,
1825
+ "rstrip": false,
1826
+ "single_word": false,
1827
+ "special": true
1828
+ },
1829
+ "157119": {
1830
+ "content": "<|reserved_token_224|>",
1831
+ "lstrip": false,
1832
+ "normalized": false,
1833
+ "rstrip": false,
1834
+ "single_word": false,
1835
+ "special": true
1836
+ },
1837
+ "157120": {
1838
+ "content": "<|reserved_token_225|>",
1839
+ "lstrip": false,
1840
+ "normalized": false,
1841
+ "rstrip": false,
1842
+ "single_word": false,
1843
+ "special": true
1844
+ },
1845
+ "157121": {
1846
+ "content": "<|reserved_token_226|>",
1847
+ "lstrip": false,
1848
+ "normalized": false,
1849
+ "rstrip": false,
1850
+ "single_word": false,
1851
+ "special": true
1852
+ },
1853
+ "157122": {
1854
+ "content": "<|reserved_token_227|>",
1855
+ "lstrip": false,
1856
+ "normalized": false,
1857
+ "rstrip": false,
1858
+ "single_word": false,
1859
+ "special": true
1860
+ },
1861
+ "157123": {
1862
+ "content": "<|reserved_token_228|>",
1863
+ "lstrip": false,
1864
+ "normalized": false,
1865
+ "rstrip": false,
1866
+ "single_word": false,
1867
+ "special": true
1868
+ },
1869
+ "157124": {
1870
+ "content": "<|reserved_token_229|>",
1871
+ "lstrip": false,
1872
+ "normalized": false,
1873
+ "rstrip": false,
1874
+ "single_word": false,
1875
+ "special": true
1876
+ },
1877
+ "157125": {
1878
+ "content": "<|reserved_token_230|>",
1879
+ "lstrip": false,
1880
+ "normalized": false,
1881
+ "rstrip": false,
1882
+ "single_word": false,
1883
+ "special": true
1884
+ },
1885
+ "157126": {
1886
+ "content": "<|reserved_token_231|>",
1887
+ "lstrip": false,
1888
+ "normalized": false,
1889
+ "rstrip": false,
1890
+ "single_word": false,
1891
+ "special": true
1892
+ },
1893
+ "157127": {
1894
+ "content": "<|reserved_token_232|>",
1895
+ "lstrip": false,
1896
+ "normalized": false,
1897
+ "rstrip": false,
1898
+ "single_word": false,
1899
+ "special": true
1900
+ },
1901
+ "157128": {
1902
+ "content": "<|reserved_token_233|>",
1903
+ "lstrip": false,
1904
+ "normalized": false,
1905
+ "rstrip": false,
1906
+ "single_word": false,
1907
+ "special": true
1908
+ },
1909
+ "157129": {
1910
+ "content": "<|reserved_token_234|>",
1911
+ "lstrip": false,
1912
+ "normalized": false,
1913
+ "rstrip": false,
1914
+ "single_word": false,
1915
+ "special": true
1916
+ },
1917
+ "157130": {
1918
+ "content": "<|reserved_token_235|>",
1919
+ "lstrip": false,
1920
+ "normalized": false,
1921
+ "rstrip": false,
1922
+ "single_word": false,
1923
+ "special": true
1924
+ },
1925
+ "157131": {
1926
+ "content": "<|reserved_token_236|>",
1927
+ "lstrip": false,
1928
+ "normalized": false,
1929
+ "rstrip": false,
1930
+ "single_word": false,
1931
+ "special": true
1932
+ },
1933
+ "157132": {
1934
+ "content": "<|reserved_token_237|>",
1935
+ "lstrip": false,
1936
+ "normalized": false,
1937
+ "rstrip": false,
1938
+ "single_word": false,
1939
+ "special": true
1940
+ },
1941
+ "157133": {
1942
+ "content": "<|reserved_token_238|>",
1943
+ "lstrip": false,
1944
+ "normalized": false,
1945
+ "rstrip": false,
1946
+ "single_word": false,
1947
+ "special": true
1948
+ },
1949
+ "157134": {
1950
+ "content": "<|reserved_token_239|>",
1951
+ "lstrip": false,
1952
+ "normalized": false,
1953
+ "rstrip": false,
1954
+ "single_word": false,
1955
+ "special": true
1956
+ },
1957
+ "157135": {
1958
+ "content": "<|reserved_token_240|>",
1959
+ "lstrip": false,
1960
+ "normalized": false,
1961
+ "rstrip": false,
1962
+ "single_word": false,
1963
+ "special": true
1964
+ },
1965
+ "157136": {
1966
+ "content": "<|reserved_token_241|>",
1967
+ "lstrip": false,
1968
+ "normalized": false,
1969
+ "rstrip": false,
1970
+ "single_word": false,
1971
+ "special": true
1972
+ },
1973
+ "157137": {
1974
+ "content": "<|reserved_token_242|>",
1975
+ "lstrip": false,
1976
+ "normalized": false,
1977
+ "rstrip": false,
1978
+ "single_word": false,
1979
+ "special": true
1980
+ },
1981
+ "157138": {
1982
+ "content": "<|reserved_token_243|>",
1983
+ "lstrip": false,
1984
+ "normalized": false,
1985
+ "rstrip": false,
1986
+ "single_word": false,
1987
+ "special": true
1988
+ },
1989
+ "157139": {
1990
+ "content": "<|reserved_token_244|>",
1991
+ "lstrip": false,
1992
+ "normalized": false,
1993
+ "rstrip": false,
1994
+ "single_word": false,
1995
+ "special": true
1996
+ },
1997
+ "157140": {
1998
+ "content": "<|reserved_token_245|>",
1999
+ "lstrip": false,
2000
+ "normalized": false,
2001
+ "rstrip": false,
2002
+ "single_word": false,
2003
+ "special": true
2004
+ },
2005
+ "157141": {
2006
+ "content": "<|reserved_token_246|>",
2007
+ "lstrip": false,
2008
+ "normalized": false,
2009
+ "rstrip": false,
2010
+ "single_word": false,
2011
+ "special": true
2012
+ },
2013
+ "157142": {
2014
+ "content": "<|reserved_token_247|>",
2015
+ "lstrip": false,
2016
+ "normalized": false,
2017
+ "rstrip": false,
2018
+ "single_word": false,
2019
+ "special": true
2020
+ },
2021
+ "157143": {
2022
+ "content": "<|reserved_token_248|>",
2023
+ "lstrip": false,
2024
+ "normalized": false,
2025
+ "rstrip": false,
2026
+ "single_word": false,
2027
+ "special": true
2028
+ },
2029
+ "157144": {
2030
+ "content": "<|reserved_token_249|>",
2031
+ "lstrip": false,
2032
+ "normalized": false,
2033
+ "rstrip": false,
2034
+ "single_word": false,
2035
+ "special": true
2036
+ },
2037
+ "157145": {
2038
+ "content": "<|reserved_token_250|>",
2039
+ "lstrip": false,
2040
+ "normalized": false,
2041
+ "rstrip": false,
2042
+ "single_word": false,
2043
+ "special": true
2044
+ },
2045
+ "157146": {
2046
+ "content": "<|reserved_token_251|>",
2047
+ "lstrip": false,
2048
+ "normalized": false,
2049
+ "rstrip": false,
2050
+ "single_word": false,
2051
+ "special": true
2052
+ },
2053
+ "157147": {
2054
+ "content": "<|reserved_token_252|>",
2055
+ "lstrip": false,
2056
+ "normalized": false,
2057
+ "rstrip": false,
2058
+ "single_word": false,
2059
+ "special": true
2060
+ },
2061
+ "157148": {
2062
+ "content": "<|reserved_token_253|>",
2063
+ "lstrip": false,
2064
+ "normalized": false,
2065
+ "rstrip": false,
2066
+ "single_word": false,
2067
+ "special": true
2068
+ },
2069
+ "157149": {
2070
+ "content": "<|reserved_token_254|>",
2071
+ "lstrip": false,
2072
+ "normalized": false,
2073
+ "rstrip": false,
2074
+ "single_word": false,
2075
+ "special": true
2076
+ },
2077
+ "157150": {
2078
+ "content": "<|reserved_token_255|>",
2079
+ "lstrip": false,
2080
+ "normalized": false,
2081
+ "rstrip": false,
2082
+ "single_word": false,
2083
+ "special": true
2084
+ },
2085
+ "157151": {
2086
+ "content": "<role>",
2087
+ "lstrip": false,
2088
+ "normalized": false,
2089
+ "rstrip": false,
2090
+ "single_word": false,
2091
+ "special": true
2092
+ },
2093
+ "157152": {
2094
+ "content": "</role>",
2095
+ "lstrip": false,
2096
+ "normalized": false,
2097
+ "rstrip": false,
2098
+ "single_word": false,
2099
+ "special": true
2100
+ },
2101
+ "157153": {
2102
+ "content": "<|arithmetic_end|>",
2103
+ "lstrip": false,
2104
+ "normalized": false,
2105
+ "rstrip": false,
2106
+ "single_word": false,
2107
+ "special": true
2108
+ },
2109
+ "157154": {
2110
+ "content": "<|number_end|>",
2111
+ "lstrip": false,
2112
+ "normalized": false,
2113
+ "rstrip": false,
2114
+ "single_word": false,
2115
+ "special": true
2116
+ },
2117
+ "157155": {
2118
+ "content": "<|arithmetic_start|>",
2119
+ "lstrip": false,
2120
+ "normalized": false,
2121
+ "rstrip": false,
2122
+ "single_word": false,
2123
+ "special": true
2124
+ },
2125
+ "157156": {
2126
+ "content": "<|number_start|>",
2127
+ "lstrip": false,
2128
+ "normalized": false,
2129
+ "rstrip": false,
2130
+ "single_word": false,
2131
+ "special": true
2132
+ },
2133
+ "157157": {
2134
+ "content": "<imagePatch>",
2135
+ "lstrip": false,
2136
+ "normalized": false,
2137
+ "rstrip": false,
2138
+ "single_word": false,
2139
+ "special": true
2140
+ },
2141
+ "157158": {
2142
+ "content": "<image>",
2143
+ "lstrip": false,
2144
+ "normalized": false,
2145
+ "rstrip": false,
2146
+ "single_word": false,
2147
+ "special": true
2148
+ },
2149
+ "157159": {
2150
+ "content": "</image>",
2151
+ "lstrip": false,
2152
+ "normalized": false,
2153
+ "rstrip": false,
2154
+ "single_word": false,
2155
+ "special": true
2156
+ },
2157
+ "157160": {
2158
+ "content": "<video>",
2159
+ "lstrip": false,
2160
+ "normalized": false,
2161
+ "rstrip": false,
2162
+ "single_word": false,
2163
+ "special": true
2164
+ },
2165
+ "157161": {
2166
+ "content": "</video>",
2167
+ "lstrip": false,
2168
+ "normalized": false,
2169
+ "rstrip": false,
2170
+ "single_word": false,
2171
+ "special": true
2172
+ },
2173
+ "157162": {
2174
+ "content": "<gen_imagePatch>",
2175
+ "lstrip": false,
2176
+ "normalized": false,
2177
+ "rstrip": false,
2178
+ "single_word": false,
2179
+ "special": true
2180
+ },
2181
+ "157163": {
2182
+ "content": "<gen_image>",
2183
+ "lstrip": false,
2184
+ "normalized": false,
2185
+ "rstrip": false,
2186
+ "single_word": false,
2187
+ "special": true
2188
+ },
2189
+ "157164": {
2190
+ "content": "</gen_image>",
2191
+ "lstrip": false,
2192
+ "normalized": false,
2193
+ "rstrip": false,
2194
+ "single_word": false,
2195
+ "special": true
2196
+ },
2197
+ "157165": {
2198
+ "content": "<imageHere>",
2199
+ "lstrip": false,
2200
+ "normalized": false,
2201
+ "rstrip": false,
2202
+ "single_word": false,
2203
+ "special": true
2204
+ },
2205
+ "157166": {
2206
+ "content": "<end_of_chunk>",
2207
+ "lstrip": false,
2208
+ "normalized": false,
2209
+ "rstrip": false,
2210
+ "single_word": false,
2211
+ "special": true
2212
+ },
2213
+ "157167": {
2214
+ "content": "<end_of_audio>",
2215
+ "lstrip": false,
2216
+ "normalized": false,
2217
+ "rstrip": false,
2218
+ "single_word": false,
2219
+ "special": true
2220
+ },
2221
+ "157168": {
2222
+ "content": "<audioPatch>",
2223
+ "lstrip": false,
2224
+ "normalized": false,
2225
+ "rstrip": false,
2226
+ "single_word": false,
2227
+ "special": true
2228
+ },
2229
+ "157169": {
2230
+ "content": "<audio>",
2231
+ "lstrip": false,
2232
+ "normalized": false,
2233
+ "rstrip": false,
2234
+ "single_word": false,
2235
+ "special": true
2236
+ },
2237
+ "157170": {
2238
+ "content": "</audio>",
2239
+ "lstrip": false,
2240
+ "normalized": false,
2241
+ "rstrip": false,
2242
+ "single_word": false,
2243
+ "special": true
2244
+ },
2245
+ "157171": {
2246
+ "content": "<gen_audioPatch>",
2247
+ "lstrip": false,
2248
+ "normalized": false,
2249
+ "rstrip": false,
2250
+ "single_word": false,
2251
+ "special": true
2252
+ },
2253
+ "157172": {
2254
+ "content": "<gen_audio>",
2255
+ "lstrip": false,
2256
+ "normalized": false,
2257
+ "rstrip": false,
2258
+ "single_word": false,
2259
+ "special": true
2260
+ },
2261
+ "157173": {
2262
+ "content": "</gen_audio>",
2263
+ "lstrip": false,
2264
+ "normalized": false,
2265
+ "rstrip": false,
2266
+ "single_word": false,
2267
+ "special": true
2268
+ },
2269
+ "157174": {
2270
+ "content": "<audioHere>",
2271
+ "lstrip": false,
2272
+ "normalized": false,
2273
+ "rstrip": false,
2274
+ "single_word": false,
2275
+ "special": true
2276
+ },
2277
+ "157175": {
2278
+ "content": "<framePatch>",
2279
+ "lstrip": false,
2280
+ "normalized": false,
2281
+ "rstrip": false,
2282
+ "single_word": false,
2283
+ "special": true
2284
+ },
2285
+ "157176": {
2286
+ "content": "<text>",
2287
+ "lstrip": false,
2288
+ "normalized": false,
2289
+ "rstrip": false,
2290
+ "single_word": false,
2291
+ "special": true
2292
+ },
2293
+ "157177": {
2294
+ "content": "<asr>",
2295
+ "lstrip": false,
2296
+ "normalized": false,
2297
+ "rstrip": false,
2298
+ "single_word": false,
2299
+ "special": true
2300
+ },
2301
+ "157178": {
2302
+ "content": "<tts>",
2303
+ "lstrip": false,
2304
+ "normalized": false,
2305
+ "rstrip": false,
2306
+ "single_word": false,
2307
+ "special": true
2308
+ }
2309
+ },
2310
+ "additional_special_tokens": [
2311
+ "<|arithmetic_end|>",
2312
+ "<role>",
2313
+ "<|number_end|>",
2314
+ "<|arithmetic_start|>",
2315
+ "</role>",
2316
+ "<|number_start|>",
2317
+ "<|role_end|>",
2318
+ "<tool_call>",
2319
+ "</tool_call>",
2320
+ "<tool_response>",
2321
+ "</tool_response>"
2322
+ ],
2323
+ "bos_token": "<|startoftext|>",
2324
+ "chat_template": "{% set thinking_option = 'off' %}\n{{- '<role>SYSTEM</role>' }}\n{%- if messages[0].role == 'system' %}\n {{- messages[0].content + '\\n\\n' }}\n{%- endif %}\n{%- if tools %}\n {{- \"# Tools\\n\\nYou may call one or more functions to assist with the user query.\\n\\nYou are provided with function signatures within <tools></tools> XML tags:\\n<tools>\" }}\n {%- for tool in tools %}\n {{- \"\\n\" }}\n {{- tool | tojson }}\n {%- endfor %}\n {{- \"\\n</tools>\\n\\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\\n<tool_call>\\n{\\\"name\\\": <function-name>, \\\"arguments\\\": <args-json-object>}\\n</tool_call>\\n\" }}\n{%- endif %}\n{{- 'detailed thinking ' + thinking_option + '<|role_end|>' }}\n{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}\n{%- for message in messages[::-1] %}\n {%- set index = (messages|length - 1) - loop.index0 %}\n {%- if ns.multi_step_tool and message.role == \"user\" and message.content is string and not(message.content.startswith('<tool_response>') and message.content.endswith('</tool_response>')) %}\n {%- set ns.multi_step_tool = false %}\n {%- set ns.last_query_index = index %}\n {%- endif %}\n{%- endfor %}\n{%- for message in messages %}\n {%- if message.content is string %}\n {%- set content = message.content %}\n {%- else %}\n {%- set content = '' %}\n {%- endif %}\n {%- if message.role == \"user\" %}\n {{- '<role>HUMAN</role>' + message.content + '<|role_end|>' }}\n {%- elif message.role == \"system\" and not loop.first %}\n {{- '<role>SYSTEM</role>' + message.content + '<|role_end|>' }}\n {%- elif message.role == \"assistant\" %}\n {%- set reasoning_content = '' %}\n {%- if message.reasoning_content is string %}\n {%- set reasoning_content = message.reasoning_content %}\n {%- else %}\n {%- if '</think>' in content %}\n {%- set reasoning_content = content.split('</think>')[0].rstrip('\\n').split('<think>')[-1].lstrip('\\n') %}\n {%- set content = content.split('</think>')[-1].lstrip('\\n') %}\n {%- endif %}\n {%- endif %}\n {%- if loop.index0 > ns.last_query_index %}\n {%- if reasoning_content %}\n {{- '<role>ASSISTANT</role>' + '\\n<think>\\n' + reasoning_content.strip('\\n') + '\\n</think>\\n\\n' + content.lstrip('\\n') }}\n {%- else %}\n {{- '<role>ASSISTANT</role>' + content }}\n {%- endif %}\n {%- else %}\n {{- '<role>ASSISTANT</role>' + content }}\n {%- endif %}\n {%- if message.tool_calls %}\n {%- for tool_call in message.tool_calls %}\n {%- if (loop.first and content) or (not loop.first) %}\n {{- '\\n' }}\n {%- endif %}\n {%- if tool_call.function %}\n {%- set tool_call = tool_call.function %}\n {%- endif %}\n {{- '<tool_call>\\n{\"name\": \"' }}\n {{- tool_call.name }}\n {{- '\", \"arguments\": ' }}\n {%- if tool_call.arguments is string %}\n {{- tool_call.arguments }}\n {%- else %}\n {{- tool_call.arguments | tojson }}\n {%- endif %}\n {{- '}\\n</tool_call>' }}\n {%- endfor %}\n {%- endif %}\n {{- '<|role_end|>' }}\n {%- elif message.role == \"tool\" %}\n {%- if loop.first or (messages[loop.index0 - 1].role != \"tool\") %}\n {{- '<role>OBSERVATION</role>' }}\n {%- endif %}\n {{- '\\n<tool_response>\\n' }}\n {{- content }}\n {{- '\\n</tool_response>' }}\n {%- if loop.last or (messages[loop.index0 + 1].role != \"tool\") %}\n {{- '<|role_end|>' }}\n {%- endif %}\n {%- endif %}\n{%- endfor %}\n{%- if add_generation_prompt %}\n {{- '<role>ASSISTANT</role>' }}\n{%- endif %}",
2325
+ "clean_up_tokenization_spaces": false,
2326
+ "cls_token": "[CLS]",
2327
+ "eos_token": "<|role_end|>",
2328
+ "extra_special_tokens": {},
2329
+ "fast_tokenizer": true,
2330
+ "gmask_token": "[gMASK]",
2331
+ "merges_file": null,
2332
+ "model_max_length": 1000000000000000019884624838656,
2333
+ "pad_token": "<|role_end|>"
2334
+ }
mlp/config.json ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "use_identity_mlp": true,
3
+ "use_vlm_directvlm_condition": true,
4
+ "use_learnable_token_condition": true,
5
+ "diffusion_c_input_dim": 2560,
6
+ "text_encoder_norm": false,
7
+ "connector_norm": false,
8
+ "img_gen_scales": [
9
+ 16
10
+ ],
11
+ "selected_hidden_states_layers": [
12
+ 5,
13
+ 12,
14
+ 20
15
+ ],
16
+ "diffusion_inner_dim": 3840
17
+ }
samples/cabin_upstream.json ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "canvas_settings": {
3
+ "aspect_ratio": "1:1, 2048 × 2048 px",
4
+ "ambient_lighting": "Four season-specific lighting conditions unified by one fixed eye-level sun direction, consistent horizon at 48% canvas height, matching atmospheric depth, and natural cinematic exposure.",
5
+ "image_style": "High-resolution cinematic photorealistic environmental concept sheet; four equal edge-to-edge vertical panels with crisp shared boundaries and no gutters, frames, extra title, or decorative inter-panel objects. Fixed 35 mm lens, eye-level three-quarter front-right camera angle, identical perspective and cabin registration in every panel. Each panel uses restrained bold uppercase condensed sans-serif labeling at the bottom."
6
+ },
7
+ "layers": [
8
+ {
9
+ "description": "Leftmost spring panel. The same small one-story rectangular log cabin appears at a fixed local position and scale: horizontal dark-brown logs, pale stone foundation, centered vertical-plank front door, two square four-pane front windows, one matching right-side window, and a charcoal standing-seam gabled roof with identical ridge height and overhang. Fresh pale-green grass, budding branches, white and pink blossoms, rain-darkened soil, shallow puddles, wet roof highlights, fine recent-rain droplets, and soft mist fill the landscape beneath an overcast clearing sky. Render the exact label \"SPRING\" centered near the bottom in bold condensed uppercase white sans-serif with a subtle dark shadow.",
10
+ "coordinates": "cx: 0.125, cy: 0.500, w: 0.250, h: 1",
11
+ "hierarchy_and_relation": "Complete first panel; aligned to the shared horizon and fixed cabin registration, flush with the canvas left edge and directly abutting the second panel without a gutter.",
12
+ "color_specs": ["#A9C7D2", "#D8E4DE", "#7FAF69", "#B7D58B", "#F2D7DE", "#F4F1E8", "#5A3B28", "#2C3032", "#7B7468", "#FFFFFF"]
13
+ },
14
+ {
15
+ "description": "Second summer panel. Repeat the cabin with exactly the same architecture, camera angle, local placement, dimensions, centered plank door, front and side windows, stone foundation, roof pitch, ridge height, and overhang. Surround it with lush meadow grass, dense layered deciduous foliage, full tree canopies, small sunlit shrubs, and dry natural ground; use a bright clear sky, strong warm sunlight from the same direction, crisp leaf highlights, controlled shadows, and mild atmospheric haze. Render the exact label \"SUMMER\" centered near the bottom in bold condensed uppercase white sans-serif with a subtle dark shadow.",
16
+ "coordinates": "cx: 0.375, cy: 0.500, w: 0.250, h: 1",
17
+ "hierarchy_and_relation": "Complete second panel; aligned to the shared horizon and fixed cabin registration, directly abutting the first and third panels without gutters.",
18
+ "color_specs": ["#4FA9D8", "#BDE8F2", "#26733E", "#4F9A45", "#79B84A", "#D8C66A", "#5A3B28", "#2C3032", "#7B7468", "#FFFFFF"]
19
+ },
20
+ {
21
+ "description": "Third autumn panel. Repeat the cabin with exactly the same architecture, camera angle, local placement, dimensions, centered plank door, front and side windows, stone foundation, roof pitch, ridge height, and overhang. Cover the landscape with amber grass, rust-orange and golden deciduous foliage, scattered fallen leaves, and a few leaves drifting naturally in the air. Use warm late-afternoon sunlight from the same direction, long soft shadows, copper rim light on logs and roof, a pale warm sky, and subtle atmospheric depth. Render the exact label \"AUTUMN\" centered near the bottom in bold condensed uppercase white sans-serif with a subtle dark shadow.",
22
+ "coordinates": "cx: 0.625, cy: 0.500, w: 0.250, h: 1",
23
+ "hierarchy_and_relation": "Complete third panel; aligned to the shared horizon and fixed cabin registration, directly abutting the second and fourth panels without gutters.",
24
+ "color_specs": ["#E6B06A", "#F0D6A2", "#D66A2C", "#A84324", "#E2A62B", "#80613B", "#5A3B28", "#2C3032", "#7B7468", "#FFFFFF"]
25
+ },
26
+ {
27
+ "description": "Rightmost winter panel. Repeat the cabin with exactly the same architecture, camera angle, local placement, dimensions, centered plank door, front and side windows, stone foundation, roof pitch, ridge height, and overhang. Blanket the ground and roof with natural snow accumulation while keeping doors and windows readable; add bare deciduous branches, snow-laden low shrubs, faint footprints away from the cabin, and sparse fine snowfall. Use cool blue ambient light, a pale clouded sky, soft low-contrast shadows from the same direction, icy edge highlights, and slight winter haze. Render the exact label \"WINTER\" centered near the bottom in bold condensed uppercase white sans-serif with a subtle dark shadow.",
28
+ "coordinates": "cx: 0.875, cy: 0.500, w: 0.250, h: 1",
29
+ "hierarchy_and_relation": "Complete fourth panel; aligned to the shared horizon and fixed cabin registration, directly abutting the third panel without a gutter and flush with the canvas right edge.",
30
+ "color_specs": ["#AFC8DE", "#DCE8F1", "#F4F7F8", "#C6D8E5", "#8198AA", "#6D5949", "#5A3B28", "#2C3032", "#7B7468", "#FFFFFF"]
31
+ }
32
+ ]
33
+ }
samples/e2e_daily_grind.caption.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ A landing page hero section for a coffee subscription service called Daily Grind: the headline Fresh beans every Monday, a short tagline, a Start your subscription button, and a photo of latte art on the right. Warm earthy palette.
samples/e2e_daily_grind.json ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "canvas_settings": {
3
+ "aspect_ratio": "1:1",
4
+ "ambient_lighting": "Warm, softly diffused natural light from the upper left, producing gentle cream highlights, muted brown shadows, and a cozy editorial café atmosphere.",
5
+ "image_style": "Polished square landing-page hero with a warm minimalist editorial layout; the left side is clean graphic space and the right side is a softly blurred lifestyle photograph."
6
+ },
7
+ "layers": [
8
+ {
9
+ "description": "Full square hero background divided visually into a warm off-white paper-like left field and a softly defocused café-photo field on the right. The left field has subtle beige grain and a faint horizontal tonal transition near the lower third; the right field shows blurred amber-brown wood, pale reflections, and indistinct dark plant silhouettes.",
10
+ "coordinates": "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000",
11
+ "hierarchy_and_relation": "Base layer beneath all typography, controls, icons, and foreground drink elements; the left and right fields meet near the vertical center without a hard border.",
12
+ "color_specs": [
13
+ "#F7F1E8",
14
+ "#E9DDCB",
15
+ "#C8955D",
16
+ "#6A442D",
17
+ "#2F2118"
18
+ ]
19
+ },
20
+ {
21
+ "description": "Upper-left brand lockup: a small dark-brown outlined coffee-bean icon tilted diagonally, followed by the widely tracked uppercase wordmark 'DAILY GRIND'. The icon and wordmark share one horizontal row, with the wordmark aligned to the icon's right and vertically centered against it. The lettering is compact, bold, sans-serif, and aligned horizontally with generous surrounding whitespace.",
22
+ "coordinates": "cx: 0.151, cy: 0.082, w: 0.205, h: 0.045",
23
+ "hierarchy_and_relation": "Sits directly over the off-white background at the top-left and remains visually separate from the headline below.",
24
+ "color_specs": [
25
+ "#3B2418",
26
+ "#F7F1E8"
27
+ ]
28
+ },
29
+ {
30
+ "description": "Large two-line hero headline in a heavy rounded sans-serif. The first line reads 'Fresh beans' and the second line reads 'every Monday'. The two rendered-text occurrences form a left-aligned vertical stack, with the second row directly beneath the first and both sharing the same left edge. Both rows are left aligned, dark espresso brown, tightly stacked, and occupy most of the left half.",
31
+ "coordinates": "cx: 0.276, cy: 0.350, w: 0.458, h: 0.210",
32
+ "hierarchy_and_relation": "Primary typographic focal point on the left, positioned below the brand and above the supporting copy; it remains unobstructed against the pale background.",
33
+ "color_specs": [
34
+ "#352117",
35
+ "#F7F1E8"
36
+ ]
37
+ },
38
+ {
39
+ "description": "Two-line supporting statement in a smaller medium-weight sans-serif: first line 'Handpicked roasts delivered weekly.', second line 'Start your week with better coffee.'. The two rendered-text occurrences are stacked as consecutive left-aligned rows beneath the headline, with the second row directly below the first. The text is left aligned, muted warm brown, and evenly spaced.",
40
+ "coordinates": "cx: 0.224, cy: 0.505, w: 0.348, h: 0.068",
41
+ "hierarchy_and_relation": "Placed beneath the headline and above the call-to-action, with strong contrast against the off-white field.",
42
+ "color_specs": [
43
+ "#6F503D",
44
+ "#F7F1E8"
45
+ ]
46
+ },
47
+ {
48
+ "description": "Rounded rectangular call-to-action button with a solid burnt-orange fill. Centered inside is the white bold sans-serif label 'Start your subscription', followed by a thin white right-pointing arrow icon. The label and arrow share one horizontal row, with the arrow positioned to the label's right and both vertically centered within the button. The button has softly rounded corners and a subtle warm shadow.",
49
+ "coordinates": "cx: 0.202, cy: 0.612, w: 0.299, h: 0.064",
50
+ "hierarchy_and_relation": "Floats over the left background below the supporting statement; its saturated fill creates the strongest interactive accent on the page.",
51
+ "color_specs": [
52
+ "#C9682F",
53
+ "#FFFFFF",
54
+ "#A95225"
55
+ ]
56
+ },
57
+ {
58
+ "description": "Three compact benefit items arranged in one horizontal row beneath the button. Each item pairs a thin brown outline icon with two stacked text lines: a coffee-bean icon above 'Ethically' and 'sourced'; a leaf icon above 'Small-batch' and 'roasted'; and a calendar icon above 'Delivered' and 'every Monday'. Within each item, the two rendered-text occurrences form a left-aligned two-row column beside its icon; the three icon-and-text groups proceed left to right on a shared horizontal row. Headings are darker and slightly bolder than the smaller muted sublines.",
59
+ "coordinates": "cx: 0.237, cy: 0.720, w: 0.375, h: 0.060",
60
+ "hierarchy_and_relation": "Anchored along the lower-left area, aligned beneath the call-to-action and separated from the photograph by open negative space.",
61
+ "color_specs": [
62
+ "#5A3827",
63
+ "#806451",
64
+ "#F7F1E8"
65
+ ]
66
+ },
67
+ {
68
+ "description": "Large shallow ceramic latte cup viewed from a slightly elevated front angle, positioned on the right. The cup has a thick cream rim, a rounded body with warm beige speckling, and a broad curved handle extending to the right. It rests on a dark brown wooden saucer or tabletop with visible grain and soft highlights.",
69
+ "coordinates": "cx: 0.735, cy: 0.666, w: 0.491, h: 0.430",
70
+ "hierarchy_and_relation": "Foreground photographic subject overlapping the blurred café background; the cup occludes the wood surface and supports the latte-art surface above it.",
71
+ "color_specs": [
72
+ "#F3E7D4",
73
+ "#D8B992",
74
+ "#B98255",
75
+ "#5A321F",
76
+ "#2B1A12"
77
+ ]
78
+ },
79
+ {
80
+ "description": "Creamy latte surface inside the cup, filling the upper bowl with a pale beige foam. Symmetrical rosette latte art is formed by crisp white milk foam and darker caramel-brown coffee channels, with a pointed leaf-like center and soft blurred foam edges.",
81
+ "coordinates": "cx: 0.724, cy: 0.535, w: 0.315, h: 0.225",
82
+ "hierarchy_and_relation": "Contained within and partially occluded by the cup rim; the bright foam contrasts against the darker coffee and draws attention toward the upper-right.",
83
+ "color_specs": [
84
+ "#F6E8D2",
85
+ "#D9B27D",
86
+ "#A9683D",
87
+ "#7B4328"
88
+ ]
89
+ },
90
+ {
91
+ "description": "Foreground coffee beans scattered across the lower-right wooden surface, including several sharply focused oval beans near the bottom edge and softer beans behind them. The beans are roasted chestnut brown with glossy highlights and natural central creases.",
92
+ "coordinates": "cx: 0.858, cy: 0.842, w: 0.282, h: 0.210",
93
+ "hierarchy_and_relation": "Topmost photographic foreground detail, overlapping the wooden surface and partially approaching the cup without covering its main body.",
94
+ "color_specs": [
95
+ "#4A2818",
96
+ "#6E3D23",
97
+ "#8B502D",
98
+ "#C18A5B"
99
+ ]
100
+ }
101
+ ]
102
+ }
samples/info_water.json ADDED
@@ -0,0 +1,141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "canvas_settings": {
3
+ "aspect_ratio": "1:1",
4
+ "ambient_lighting": "Bright, even white studio illumination with soft pastel gradients and no cast shadows.",
5
+ "image_style": "Clean flat-vector educational infographic with rounded geometry, thin white outlines, and a centered four-stage cycle."
6
+ },
7
+ "layers": [
8
+ {
9
+ "description": "Solid white square background with generous margins around a centered water-cycle infographic.",
10
+ "coordinates": "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000",
11
+ "hierarchy_and_relation": "Base layer beneath every illustration, arrow, panel, and text element.",
12
+ "color_specs": [
13
+ "#FFFFFF"
14
+ ]
15
+ },
16
+ {
17
+ "description": "Large circular ocean-evaporation scene: a blue sun with short radial yellow rays rises over layered turquoise and deep-blue water. White-capped waves fill the lower half; a pale blue dashed vertical path rises from the water toward the top of the circle.",
18
+ "coordinates": "cx: 0.247, cy: 0.258, w: 0.333, h: 0.333",
19
+ "hierarchy_and_relation": "Occupies the upper-left cycle position; its top edge is partially occluded by the adjacent upward arrow, and it sits behind the stage label below.",
20
+ "color_specs": [
21
+ "#29A9E2",
22
+ "#7ED6F4",
23
+ "#087FC1",
24
+ "#075E91",
25
+ "#FFFFFF",
26
+ "#FFD547"
27
+ ]
28
+ },
29
+ {
30
+ "description": "Large circular cloud-formation scene: a pale blue sky contains several white and light-gray puffy clouds, with a darker gray rain cloud near the lower center. A pale blue dashed vertical path descends from the upper cloud area.",
31
+ "coordinates": "cx: 0.750, cy: 0.258, w: 0.333, h: 0.333",
32
+ "hierarchy_and_relation": "Occupies the upper-right cycle position; its left edge is partially occluded by the incoming arrow from the evaporation stage and its lower edge is partially occluded by the downward arrow.",
33
+ "color_specs": [
34
+ "#B9E5F5",
35
+ "#FFFFFF",
36
+ "#D9D9D9",
37
+ "#858585",
38
+ "#29A9E2"
39
+ ]
40
+ },
41
+ {
42
+ "description": "Large circular rainfall scene: dark blue-gray storm clouds and falling blue raindrops fill the upper portion above a vivid turquoise-blue water surface with small white wave marks.",
43
+ "coordinates": "cx: 0.750, cy: 0.742, w: 0.333, h: 0.333",
44
+ "hierarchy_and_relation": "Occupies the lower-right cycle position; receives the downward arrow from the cloud stage and sends a curved arrow toward the collection stage.",
45
+ "color_specs": [
46
+ "#555F68",
47
+ "#747D85",
48
+ "#29A9E2",
49
+ "#087FC1",
50
+ "#FFFFFF"
51
+ ]
52
+ },
53
+ {
54
+ "description": "Large circular land-collection scene: green hills and grassy banks surround a blue lake or river. Dark evergreen trees stand on the left and right banks, pale gray mountains recede in the distance, and a small white cloud floats at upper left.",
55
+ "coordinates": "cx: 0.247, cy: 0.742, w: 0.333, h: 0.333",
56
+ "hierarchy_and_relation": "Occupies the lower-left cycle position; receives the curved arrow from rainfall and sends a long curved arrow upward toward evaporation.",
57
+ "color_specs": [
58
+ "#238F38",
59
+ "#56A94B",
60
+ "#087FC1",
61
+ "#29A9E2",
62
+ "#A9C7D8",
63
+ "#FFFFFF",
64
+ "#4F7E35"
65
+ ]
66
+ },
67
+ {
68
+ "description": "Thick curved directional arrows with rounded shafts and triangular heads: one rises from the upper-left circle toward the upper-right circle, one descends from the upper-right circle toward the lower-right circle, one curves from the lower-right circle toward the lower-left circle, and one curves from the lower-left circle toward the upper-left circle. Each arrow is medium blue with a white outline.",
69
+ "coordinates": "cx: 0.500, cy: 0.500, w: 0.744, h: 0.744",
70
+ "hierarchy_and_relation": "Overlays the gaps between the four circular scenes, visually linking them into a clockwise cycle while remaining behind the stage labels.",
71
+ "color_specs": [
72
+ "#29A9E2",
73
+ "#FFFFFF"
74
+ ]
75
+ },
76
+ {
77
+ "description": "Bold dark-blue uppercase stage heading 'EVAPORATION', centered above the upper-left circular scene. The heading forms the first row of a centered two-row label group, stacked directly above its explanatory copy.",
78
+ "coordinates": "cx: 0.247, cy: 0.070, w: 0.245, h: 0.031",
79
+ "hierarchy_and_relation": "Topmost label for the evaporation stage, positioned above and not touching the circular water illustration.",
80
+ "color_specs": [
81
+ "#075E91"
82
+ ]
83
+ },
84
+ {
85
+ "description": "Two centered lines of small dark-gray explanatory copy beneath the upper-left heading: first line 'The sun heats up water', second line 'and it turns into vapor.'. The two occurrences form a compact stacked pair on separate rows, share a common center axis, and sit directly beneath the heading.",
86
+ "coordinates": "cx: 0.247, cy: 0.111, w: 0.190, h: 0.043",
87
+ "hierarchy_and_relation": "Secondary text block directly below the evaporation heading and above the ocean scene.",
88
+ "color_specs": [
89
+ "#333333"
90
+ ]
91
+ },
92
+ {
93
+ "description": "Bold dark-blue uppercase stage heading 'CONDENSATION', centered above the upper-right circular scene. The heading forms the first row of a centered two-row label group, stacked directly above its explanatory copy.",
94
+ "coordinates": "cx: 0.750, cy: 0.070, w: 0.245, h: 0.031",
95
+ "hierarchy_and_relation": "Topmost label for the condensation stage, positioned above the cloud illustration.",
96
+ "color_specs": [
97
+ "#075E91"
98
+ ]
99
+ },
100
+ {
101
+ "description": "Two centered lines of small dark-gray explanatory copy beneath the upper-right heading: first line 'Water vapor cools', second line 'and forms clouds.'. The two occurrences form a compact stacked pair on separate rows, share a common center axis, and sit directly beneath the heading.",
102
+ "coordinates": "cx: 0.750, cy: 0.111, w: 0.170, h: 0.043",
103
+ "hierarchy_and_relation": "Secondary text block directly below the condensation heading and above the cloud scene.",
104
+ "color_specs": [
105
+ "#333333"
106
+ ]
107
+ },
108
+ {
109
+ "description": "Bold dark-blue uppercase stage heading 'PRECIPITATION', centered beside the lower-right circular scene. The heading forms the first row of a centered two-row label group, stacked directly above its explanatory copy.",
110
+ "coordinates": "cx: 0.750, cy: 0.918, w: 0.245, h: 0.031",
111
+ "hierarchy_and_relation": "Topmost label below the rainfall illustration, aligned to its horizontal center.",
112
+ "color_specs": [
113
+ "#075E91"
114
+ ]
115
+ },
116
+ {
117
+ "description": "Two centered lines of small dark-gray explanatory copy beneath the lower-right heading: first line 'Water falls as rain,', second line 'snow, or hail.'. The two occurrences form a compact stacked pair on separate rows, share a common center axis, and sit directly beneath the heading.",
118
+ "coordinates": "cx: 0.750, cy: 0.959, w: 0.170, h: 0.043",
119
+ "hierarchy_and_relation": "Secondary text block directly below the precipitation heading and above the bottom canvas margin.",
120
+ "color_specs": [
121
+ "#333333"
122
+ ]
123
+ },
124
+ {
125
+ "description": "Bold dark-blue uppercase stage heading 'COLLECTION', centered above the lower-left circular scene. The heading forms the first row of a centered two-row label group, stacked directly above its explanatory copy.",
126
+ "coordinates": "cx: 0.247, cy: 0.918, w: 0.200, h: 0.031",
127
+ "hierarchy_and_relation": "Topmost label below the land-and-water illustration, aligned to its horizontal center.",
128
+ "color_specs": [
129
+ "#075E91"
130
+ ]
131
+ },
132
+ {
133
+ "description": "Two centered lines of small dark-gray explanatory copy beneath the lower-left heading: first line 'Water collects in rivers, lakes,', second line 'and oceans.'. The two occurrences form a compact stacked pair on separate rows, share a common center axis, and sit directly beneath the heading.",
134
+ "coordinates": "cx: 0.247, cy: 0.959, w: 0.220, h: 0.043",
135
+ "hierarchy_and_relation": "Secondary text block directly below the collection heading and above the bottom canvas margin.",
136
+ "color_specs": [
137
+ "#333333"
138
+ ]
139
+ }
140
+ ]
141
+ }
samples/poster_jazz.json ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "canvas_settings": {
3
+ "aspect_ratio": "1:1",
4
+ "ambient_lighting": "Low-key nocturnal lighting with a soft cyan-blue glow rising from the lower center, fading into near-black navy at the top; the saxophone silhouette is rim-lit by this cool backlight.",
5
+ "image_style": "Minimalist square concert-poster design with smooth digital gradients, crisp centered typography, and a single high-contrast photographic-style instrument silhouette."
6
+ },
7
+ "layers": [
8
+ {
9
+ "description": "Full-square background formed by a smooth deep navy-to-black vertical gradient, subtly brighter blue around the lower center and nearly black along the upper edge; no visible texture beyond the soft tonal falloff.",
10
+ "coordinates": "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000",
11
+ "hierarchy_and_relation": "Base layer beneath all typography, ornament, and instrument imagery; its lower-center glow silhouettes the saxophone.",
12
+ "color_specs": [
13
+ "#020A18",
14
+ "#06152A",
15
+ "#0A2848",
16
+ "#123E63"
17
+ ]
18
+ },
19
+ {
20
+ "description": "Large centered event title 'Blue Hour' in an elegant high-contrast serif typeface, pale blue-white, with very large letterforms and generous spacing; the title occupies the upper-middle area on a single centered baseline and is horizontally centered over the saxophone and lower information block.",
21
+ "coordinates": "cx: 0.500, cy: 0.250, w: 0.760, h: 0.145",
22
+ "hierarchy_and_relation": "Topmost typography over the dark gradient, centered above the saxophone and visually dominant over all other text.",
23
+ "color_specs": [
24
+ "#D8E8F7",
25
+ "#BFD8EF"
26
+ ]
27
+ },
28
+ {
29
+ "description": "Small uppercase subtitle 'A JAZZ NIGHT' centered below the main title, set in widely tracked sans-serif letters; a thin horizontal rule appears on each side, aligned to the subtitle baseline and extending outward symmetrically toward the left and right margins.",
30
+ "coordinates": "cx: 0.500, cy: 0.345, w: 0.520, h: 0.030",
31
+ "hierarchy_and_relation": "Placed directly beneath the title and above the saxophone, with the flanking rules framing the subtitle without touching the letters.",
32
+ "color_specs": [
33
+ "#8FB4D2",
34
+ "#5F88AA"
35
+ ]
36
+ },
37
+ {
38
+ "description": "Vertical saxophone silhouette in dark navy-black, shown in side profile facing left with the bell flaring at the lower left; curved neck, mouthpiece, body keys, and bell are visible, with a narrow cool-blue highlight tracing the right edge and lower rim.",
39
+ "coordinates": "cx: 0.493, cy: 0.574, w: 0.285, h: 0.445",
40
+ "hierarchy_and_relation": "Central photographic-style object over the gradient; it sits below the title block and behind the lower event-information typography, with its lower portion partially occluded by the date row.",
41
+ "color_specs": [
42
+ "#020B18",
43
+ "#06152A",
44
+ "#0B2742",
45
+ "#2E6F9D"
46
+ ]
47
+ },
48
+ {
49
+ "description": "Thin circular blue outline behind the saxophone, centered around the instrument body; the ring is delicate, mostly visible along the left and lower arc, and partially hidden by the saxophone silhouette.",
50
+ "coordinates": "cx: 0.499, cy: 0.574, w: 0.350, h: 0.350",
51
+ "hierarchy_and_relation": "Decorative midground element between the gradient background and saxophone; the saxophone occludes portions of the ring while the ring remains visible around the left and bottom edges.",
52
+ "color_specs": [
53
+ "#174A73",
54
+ "#0D3153"
55
+ ]
56
+ },
57
+ {
58
+ "description": "Centered date line 'Friday, October 3' in a clean light sans-serif typeface, pale blue-gray, positioned below the saxophone bell and body.",
59
+ "coordinates": "cx: 0.500, cy: 0.786, w: 0.430, h: 0.040",
60
+ "hierarchy_and_relation": "Foreground text below the instrument, centered over the lower gradient and above the time-and-venue row.",
61
+ "color_specs": [
62
+ "#C7D9EA",
63
+ "#9FB9D0"
64
+ ]
65
+ },
66
+ {
67
+ "description": "Centered lower information row with '8 PM' at left and 'The Lantern Room' at right, both in small widely tracked uppercase sans-serif lettering; a thin horizontal rule separates this row from the date above and aligns with the left and right text endpoints.",
68
+ "coordinates": "cx: 0.500, cy: 0.835, w: 0.430, h: 0.055",
69
+ "hierarchy_and_relation": "Foreground information layer beneath the date line; the two text occurrences share a baseline and are separated by open space, while the rule spans behind them horizontally.",
70
+ "color_specs": [
71
+ "#A8C2D9",
72
+ "#6F93B1",
73
+ "#315776"
74
+ ]
75
+ },
76
+ {
77
+ "description": "Small bottom decorative mark centered near the lower edge: a thin vertical stem with a tiny blue dot above it and a short horizontal base line.",
78
+ "coordinates": "cx: 0.500, cy: 0.923, w: 0.060, h: 0.055",
79
+ "hierarchy_and_relation": "Topmost minimal ornament below the event information, isolated against the dark background and aligned to the poster center axis.",
80
+ "color_specs": [
81
+ "#1E5A86",
82
+ "#0B2D4A"
83
+ ]
84
+ }
85
+ ]
86
+ }
samples/ui_banking.json ADDED
@@ -0,0 +1,94 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "canvas_settings": {
3
+ "aspect_ratio": "1:1",
4
+ "ambient_lighting": "soft cool daylight UI lighting with pale blue-white background glow and subtle card shadows",
5
+ "image_style": "clean modern fintech mobile app UI, rounded white cards, soft drop shadows, sans-serif typography, flat vector icons"
6
+ },
7
+ "layers": [
8
+ {
9
+ "description": "Full-screen mobile app background with a very pale blue-to-white gradient, slightly brighter at the top and lower edges, smooth glassy fintech aesthetic with no visible texture beyond soft gradient shading.",
10
+ "coordinates": "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000",
11
+ "hierarchy_and_relation": "Z-index 0 background layer; all cards, text, icons, and navigation elements sit above it with soft shadows separating them from the pale surface.",
12
+ "color_specs": [
13
+ "#F3F8FF",
14
+ "#EAF3FF",
15
+ "#FFFFFF"
16
+ ]
17
+ },
18
+ {
19
+ "description": "Top status bar with dark time text '9:41' at the upper left, and compact black cellular signal bars, Wi-Fi arcs, and a rounded battery outline with dark fill at the upper right. The time and the right-side icon cluster sit on the same horizontal status-bar row, aligned near the top edge with wide empty space between them.",
20
+ "coordinates": "cx: 0.500, cy: 0.030, w: 0.850, h: 0.035",
21
+ "hierarchy_and_relation": "Sits above the background and above the app header, visually aligned to the top safe area without overlapping the balance card.",
22
+ "color_specs": [
23
+ "#111827",
24
+ "#000000",
25
+ "#FFFFFF",
26
+ "#D1D5DB"
27
+ ]
28
+ },
29
+ {
30
+ "description": "Header area with bold dark greeting text 'Good morning, Alex' near the upper left, smaller muted subtitle 'Here's your financial overview' directly underneath, and a circular white notification button on the upper right containing a thin outlined bell icon with a small blue unread dot. The greeting and subtitle form a left-aligned two-line stack, while the notification button is horizontally aligned with the greeting block and separated to the far right.",
31
+ "coordinates": "cx: 0.500, cy: 0.095, w: 0.850, h: 0.075",
32
+ "hierarchy_and_relation": "Placed below the status bar and above the balance card; the notification button floats on the same horizontal band and casts a faint shadow over the background.",
33
+ "color_specs": [
34
+ "#111827",
35
+ "#6B7280",
36
+ "#FFFFFF",
37
+ "#2F80ED",
38
+ "#E5E7EB"
39
+ ]
40
+ },
41
+ {
42
+ "description": "Large rounded balance card with deep blue gradient fill, subtle curved wave decoration in the lower right, and a white outlined eye icon button near the top right. Left side contains small label 'Total Balance', oversized white balance amount '$8,450.75', and smaller white currency label 'USD'. Right lower area contains a white rounded square icon tile with a blue wallet/card symbol. The label, amount, and currency label are stacked in a left-aligned column inside the card, while the eye button aligns near the card's top-right corner and the wallet icon tile anchors the lower-right area.",
43
+ "coordinates": "cx: 0.500, cy: 0.250, w: 0.850, h: 0.225",
44
+ "hierarchy_and_relation": "Main hero card sits above the background and below the header; its shadow and saturated blue fill make it the dominant visual block, with the wallet icon tile layered on top of the card gradient.",
45
+ "color_specs": [
46
+ "#0F62FE",
47
+ "#1B4FD8",
48
+ "#FFFFFF",
49
+ "#EAF2FF",
50
+ "#1E5BDB"
51
+ ]
52
+ },
53
+ {
54
+ "description": "White rounded recent transactions panel with title 'Recent Transactions' at the top left and blue link 'View all' at the top right. First transaction row shows a green circular food icon with a white shopping bag outline, merchant name 'Whole Foods Market', date line 'May 20, 2024', right-aligned negative amount '-$85.42', and category 'Groceries'. Second row shows a blue circular ride icon with a white car outline, merchant name 'Uber Ride', date line 'May 20, 2024', right-aligned negative amount '-$18.75', and category 'Transport'. Third row shows a purple circular salary icon with a white wallet outline, merchant name 'Salary Deposit', date line 'May 19, 2024', right-aligned positive amount '+$2,500.00', and category 'Income'. Thin light gray dividers separate rows. The panel title and link share the top header row, and each transaction row uses a consistent left-to-right layout of icon, two-line merchant/date text block, and right-aligned amount/category stack.",
55
+ "coordinates": "cx: 0.500, cy: 0.520, w: 0.850, h: 0.285",
56
+ "hierarchy_and_relation": "Panel sits below the balance card with a soft shadow; each icon, text block, amount, and category is contained within the white card, with dividers layered behind the row content.",
57
+ "color_specs": [
58
+ "#FFFFFF",
59
+ "#111827",
60
+ "#6B7280",
61
+ "#2F80ED",
62
+ "#22C55E",
63
+ "#8B5CF6",
64
+ "#E5E7EB"
65
+ ]
66
+ },
67
+ {
68
+ "description": "White rounded quick-actions panel with heading 'Quick Actions' at the top left. Three evenly spaced action columns appear below: left column has a blue circular icon with a white paper-plane/send symbol and label 'Send' with sublabel 'Transfer money'; middle column has a green circular icon with a white payment-card symbol and label 'Pay' with sublabel 'Pay bills & more'; right column has a purple circular icon with a white plus symbol and label 'Top Up' with sublabel 'Add funds'. The heading sits above a single row of three equal action columns, each column stacking its circular icon above the bold label and smaller sublabel, all centered within its column.",
69
+ "coordinates": "cx: 0.500, cy: 0.765, w: 0.850, h: 0.170",
70
+ "hierarchy_and_relation": "Positioned below the transactions panel and above the bottom navigation; the three action groups are children of the white rounded card and are separated by spacing rather than visible dividers.",
71
+ "color_specs": [
72
+ "#FFFFFF",
73
+ "#111827",
74
+ "#6B7280",
75
+ "#2F80ED",
76
+ "#22C55E",
77
+ "#8B5CF6",
78
+ "#F3F4F6"
79
+ ]
80
+ },
81
+ {
82
+ "description": "Bottom navigation bar as a wide white rounded rectangle with four evenly spaced tabs. Left tab is active with a blue home icon and blue label 'Home'. Second tab has a gray vertical card/statement icon and label 'Cards'. Third tab has a gray pie-chart icon and label 'Insights'. Fourth tab has a gray person outline icon and label 'Profile'. The four tabs form a single evenly spaced horizontal row, with each icon centered above its label and all labels aligned along the same baseline.",
83
+ "coordinates": "cx: 0.500, cy: 0.915, w: 0.850, h: 0.095",
84
+ "hierarchy_and_relation": "Topmost persistent navigation layer sits above the background and below no other content; it overlaps the lower margin with a soft shadow and keeps all tab labels aligned along the same baseline.",
85
+ "color_specs": [
86
+ "#FFFFFF",
87
+ "#2F80ED",
88
+ "#6B7280",
89
+ "#E5E7EB",
90
+ "#111827"
91
+ ]
92
+ }
93
+ ]
94
+ }
scheduler/scheduler_config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "_class_name": "FlowMatchEulerDiscreteScheduler",
3
+ "_diffusers_version": "0.37.0.dev0",
4
+ "num_train_timesteps": 1000,
5
+ "use_dynamic_shifting": false,
6
+ "shift": 6.0
7
+ }
transformer/config.json ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_class_name": "DiffusionTransformer",
3
+ "_diffusers_version": "0.36.0",
4
+ "_name_or_path": "",
5
+ "all_f_patch_size": [
6
+ 1
7
+ ],
8
+ "all_patch_size": [
9
+ 2
10
+ ],
11
+ "axes_dims": [
12
+ 32,
13
+ 48,
14
+ 48
15
+ ],
16
+ "axes_lens": [
17
+ 20480,
18
+ 512,
19
+ 512
20
+ ],
21
+ "cap_feat_dim": 2560,
22
+ "dim": 3840,
23
+ "in_channels": 16,
24
+ "n_heads": 30,
25
+ "n_kv_heads": 30,
26
+ "n_layers": 30,
27
+ "n_refiner_layers": 2,
28
+ "norm_eps": 1e-05,
29
+ "qk_norm": true,
30
+ "rope_theta": 256.0,
31
+ "siglip_feat_dim": null,
32
+ "t_scale": 1000.0,
33
+ "alignment_padding_mode": "zero_masked",
34
+ "multi_frame_output": false
35
+ }
transformer/diffusion_pytorch_model.safetensors.index.json ADDED
@@ -0,0 +1,526 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "metadata": {
3
+ "total_size": 12309802112
4
+ },
5
+ "weight_map": {
6
+ "all_final_layer.2-1.adaLN_modulation.1.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
7
+ "all_final_layer.2-1.adaLN_modulation.1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
8
+ "all_final_layer.2-1.linear.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
9
+ "all_final_layer.2-1.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
10
+ "all_x_embedder.2-1.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
11
+ "all_x_embedder.2-1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
12
+ "cap_embedder.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
13
+ "cap_embedder.1.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
14
+ "cap_embedder.1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
15
+ "context_refiner.0.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
16
+ "context_refiner.0.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
17
+ "context_refiner.0.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
18
+ "context_refiner.0.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
19
+ "context_refiner.0.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
20
+ "context_refiner.0.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
21
+ "context_refiner.0.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
22
+ "context_refiner.0.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
23
+ "context_refiner.0.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
24
+ "context_refiner.0.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
25
+ "context_refiner.0.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
26
+ "context_refiner.0.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
27
+ "context_refiner.0.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
28
+ "context_refiner.1.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
29
+ "context_refiner.1.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
30
+ "context_refiner.1.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
31
+ "context_refiner.1.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
32
+ "context_refiner.1.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
33
+ "context_refiner.1.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
34
+ "context_refiner.1.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
35
+ "context_refiner.1.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
36
+ "context_refiner.1.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
37
+ "context_refiner.1.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
38
+ "context_refiner.1.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
39
+ "context_refiner.1.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
40
+ "context_refiner.1.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
41
+ "layers.0.adaLN_modulation.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
42
+ "layers.0.adaLN_modulation.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
43
+ "layers.0.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
44
+ "layers.0.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
45
+ "layers.0.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
46
+ "layers.0.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
47
+ "layers.0.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
48
+ "layers.0.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
49
+ "layers.0.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
50
+ "layers.0.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
51
+ "layers.0.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
52
+ "layers.0.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
53
+ "layers.0.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
54
+ "layers.0.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
55
+ "layers.0.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
56
+ "layers.1.adaLN_modulation.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
57
+ "layers.1.adaLN_modulation.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
58
+ "layers.1.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
59
+ "layers.1.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
60
+ "layers.1.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
61
+ "layers.1.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
62
+ "layers.1.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
63
+ "layers.1.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
64
+ "layers.1.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
65
+ "layers.1.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
66
+ "layers.1.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
67
+ "layers.1.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
68
+ "layers.1.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
69
+ "layers.1.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
70
+ "layers.1.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
71
+ "layers.10.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
72
+ "layers.10.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
73
+ "layers.10.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
74
+ "layers.10.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
75
+ "layers.10.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
76
+ "layers.10.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
77
+ "layers.10.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
78
+ "layers.10.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
79
+ "layers.10.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
80
+ "layers.10.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
81
+ "layers.10.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
82
+ "layers.10.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
83
+ "layers.10.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
84
+ "layers.10.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
85
+ "layers.10.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
86
+ "layers.11.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
87
+ "layers.11.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
88
+ "layers.11.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
89
+ "layers.11.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
90
+ "layers.11.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
91
+ "layers.11.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
92
+ "layers.11.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
93
+ "layers.11.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
94
+ "layers.11.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
95
+ "layers.11.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
96
+ "layers.11.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
97
+ "layers.11.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
98
+ "layers.11.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
99
+ "layers.11.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
100
+ "layers.11.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
101
+ "layers.12.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
102
+ "layers.12.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
103
+ "layers.12.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
104
+ "layers.12.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
105
+ "layers.12.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
106
+ "layers.12.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
107
+ "layers.12.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
108
+ "layers.12.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
109
+ "layers.12.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
110
+ "layers.12.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
111
+ "layers.12.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
112
+ "layers.12.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
113
+ "layers.12.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
114
+ "layers.12.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
115
+ "layers.12.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
116
+ "layers.13.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
117
+ "layers.13.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
118
+ "layers.13.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
119
+ "layers.13.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
120
+ "layers.13.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
121
+ "layers.13.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
122
+ "layers.13.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
123
+ "layers.13.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
124
+ "layers.13.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
125
+ "layers.13.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
126
+ "layers.13.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
127
+ "layers.13.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
128
+ "layers.13.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
129
+ "layers.13.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
130
+ "layers.13.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
131
+ "layers.14.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
132
+ "layers.14.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
133
+ "layers.14.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
134
+ "layers.14.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
135
+ "layers.14.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
136
+ "layers.14.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
137
+ "layers.14.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
138
+ "layers.14.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
139
+ "layers.14.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
140
+ "layers.14.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
141
+ "layers.14.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
142
+ "layers.14.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
143
+ "layers.14.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
144
+ "layers.14.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
145
+ "layers.14.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
146
+ "layers.15.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
147
+ "layers.15.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
148
+ "layers.15.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
149
+ "layers.15.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
150
+ "layers.15.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
151
+ "layers.15.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
152
+ "layers.15.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
153
+ "layers.15.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
154
+ "layers.15.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
155
+ "layers.15.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
156
+ "layers.15.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
157
+ "layers.15.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
158
+ "layers.15.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
159
+ "layers.15.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
160
+ "layers.15.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
161
+ "layers.16.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
162
+ "layers.16.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
163
+ "layers.16.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
164
+ "layers.16.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
165
+ "layers.16.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
166
+ "layers.16.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
167
+ "layers.16.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
168
+ "layers.16.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
169
+ "layers.16.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
170
+ "layers.16.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
171
+ "layers.16.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
172
+ "layers.16.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
173
+ "layers.16.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
174
+ "layers.16.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
175
+ "layers.16.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
176
+ "layers.17.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
177
+ "layers.17.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
178
+ "layers.17.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
179
+ "layers.17.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
180
+ "layers.17.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
181
+ "layers.17.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
182
+ "layers.17.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
183
+ "layers.17.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
184
+ "layers.17.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
185
+ "layers.17.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
186
+ "layers.17.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
187
+ "layers.17.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
188
+ "layers.17.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
189
+ "layers.17.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
190
+ "layers.17.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
191
+ "layers.18.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
192
+ "layers.18.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
193
+ "layers.18.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
194
+ "layers.18.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
195
+ "layers.18.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
196
+ "layers.18.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
197
+ "layers.18.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
198
+ "layers.18.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
199
+ "layers.18.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
200
+ "layers.18.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
201
+ "layers.18.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
202
+ "layers.18.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
203
+ "layers.18.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
204
+ "layers.18.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
205
+ "layers.18.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
206
+ "layers.19.adaLN_modulation.0.bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
207
+ "layers.19.adaLN_modulation.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
208
+ "layers.19.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
209
+ "layers.19.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
210
+ "layers.19.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
211
+ "layers.19.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
212
+ "layers.19.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
213
+ "layers.19.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
214
+ "layers.19.attention_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
215
+ "layers.19.attention_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
216
+ "layers.19.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
217
+ "layers.19.feed_forward.w2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
218
+ "layers.19.feed_forward.w3.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
219
+ "layers.19.ffn_norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
220
+ "layers.19.ffn_norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
221
+ "layers.2.adaLN_modulation.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
222
+ "layers.2.adaLN_modulation.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
223
+ "layers.2.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
224
+ "layers.2.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
225
+ "layers.2.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
226
+ "layers.2.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
227
+ "layers.2.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
228
+ "layers.2.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
229
+ "layers.2.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
230
+ "layers.2.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
231
+ "layers.2.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
232
+ "layers.2.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
233
+ "layers.2.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
234
+ "layers.2.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
235
+ "layers.2.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
236
+ "layers.20.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
237
+ "layers.20.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
238
+ "layers.20.attention.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
239
+ "layers.20.attention.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
240
+ "layers.20.attention.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
241
+ "layers.20.attention.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
242
+ "layers.20.attention.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
243
+ "layers.20.attention.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
244
+ "layers.20.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
245
+ "layers.20.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
246
+ "layers.20.feed_forward.w1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
247
+ "layers.20.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
248
+ "layers.20.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
249
+ "layers.20.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
250
+ "layers.20.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
251
+ "layers.21.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
252
+ "layers.21.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
253
+ "layers.21.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
254
+ "layers.21.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
255
+ "layers.21.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
256
+ "layers.21.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
257
+ "layers.21.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
258
+ "layers.21.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
259
+ "layers.21.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
260
+ "layers.21.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
261
+ "layers.21.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
262
+ "layers.21.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
263
+ "layers.21.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
264
+ "layers.21.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
265
+ "layers.21.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
266
+ "layers.22.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
267
+ "layers.22.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
268
+ "layers.22.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
269
+ "layers.22.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
270
+ "layers.22.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
271
+ "layers.22.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
272
+ "layers.22.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
273
+ "layers.22.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
274
+ "layers.22.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
275
+ "layers.22.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
276
+ "layers.22.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
277
+ "layers.22.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
278
+ "layers.22.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
279
+ "layers.22.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
280
+ "layers.22.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
281
+ "layers.23.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
282
+ "layers.23.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
283
+ "layers.23.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
284
+ "layers.23.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
285
+ "layers.23.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
286
+ "layers.23.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
287
+ "layers.23.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
288
+ "layers.23.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
289
+ "layers.23.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
290
+ "layers.23.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
291
+ "layers.23.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
292
+ "layers.23.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
293
+ "layers.23.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
294
+ "layers.23.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
295
+ "layers.23.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
296
+ "layers.24.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
297
+ "layers.24.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
298
+ "layers.24.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
299
+ "layers.24.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
300
+ "layers.24.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
301
+ "layers.24.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
302
+ "layers.24.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
303
+ "layers.24.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
304
+ "layers.24.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
305
+ "layers.24.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
306
+ "layers.24.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
307
+ "layers.24.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
308
+ "layers.24.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
309
+ "layers.24.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
310
+ "layers.24.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
311
+ "layers.25.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
312
+ "layers.25.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
313
+ "layers.25.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
314
+ "layers.25.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
315
+ "layers.25.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
316
+ "layers.25.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
317
+ "layers.25.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
318
+ "layers.25.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
319
+ "layers.25.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
320
+ "layers.25.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
321
+ "layers.25.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
322
+ "layers.25.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
323
+ "layers.25.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
324
+ "layers.25.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
325
+ "layers.25.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
326
+ "layers.26.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
327
+ "layers.26.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
328
+ "layers.26.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
329
+ "layers.26.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
330
+ "layers.26.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
331
+ "layers.26.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
332
+ "layers.26.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
333
+ "layers.26.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
334
+ "layers.26.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
335
+ "layers.26.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
336
+ "layers.26.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
337
+ "layers.26.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
338
+ "layers.26.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
339
+ "layers.26.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
340
+ "layers.26.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
341
+ "layers.27.adaLN_modulation.0.bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
342
+ "layers.27.adaLN_modulation.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
343
+ "layers.27.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
344
+ "layers.27.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
345
+ "layers.27.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
346
+ "layers.27.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
347
+ "layers.27.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
348
+ "layers.27.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
349
+ "layers.27.attention_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
350
+ "layers.27.attention_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
351
+ "layers.27.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
352
+ "layers.27.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
353
+ "layers.27.feed_forward.w3.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
354
+ "layers.27.ffn_norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
355
+ "layers.27.ffn_norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
356
+ "layers.28.adaLN_modulation.0.bias": "diffusion_pytorch_model-00005-of-00005.safetensors",
357
+ "layers.28.adaLN_modulation.0.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
358
+ "layers.28.attention.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
359
+ "layers.28.attention.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
360
+ "layers.28.attention.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
361
+ "layers.28.attention.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
362
+ "layers.28.attention.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
363
+ "layers.28.attention.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
364
+ "layers.28.attention_norm1.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
365
+ "layers.28.attention_norm2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
366
+ "layers.28.feed_forward.w1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
367
+ "layers.28.feed_forward.w2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
368
+ "layers.28.feed_forward.w3.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
369
+ "layers.28.ffn_norm1.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
370
+ "layers.28.ffn_norm2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
371
+ "layers.29.adaLN_modulation.0.bias": "diffusion_pytorch_model-00005-of-00005.safetensors",
372
+ "layers.29.adaLN_modulation.0.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
373
+ "layers.29.attention.norm_k.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
374
+ "layers.29.attention.norm_q.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
375
+ "layers.29.attention.to_k.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
376
+ "layers.29.attention.to_out.0.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
377
+ "layers.29.attention.to_q.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
378
+ "layers.29.attention.to_v.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
379
+ "layers.29.attention_norm1.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
380
+ "layers.29.attention_norm2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
381
+ "layers.29.feed_forward.w1.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
382
+ "layers.29.feed_forward.w2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
383
+ "layers.29.feed_forward.w3.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
384
+ "layers.29.ffn_norm1.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
385
+ "layers.29.ffn_norm2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
386
+ "layers.3.adaLN_modulation.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
387
+ "layers.3.adaLN_modulation.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
388
+ "layers.3.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
389
+ "layers.3.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
390
+ "layers.3.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
391
+ "layers.3.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
392
+ "layers.3.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
393
+ "layers.3.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
394
+ "layers.3.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
395
+ "layers.3.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
396
+ "layers.3.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
397
+ "layers.3.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
398
+ "layers.3.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
399
+ "layers.3.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
400
+ "layers.3.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
401
+ "layers.4.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
402
+ "layers.4.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
403
+ "layers.4.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
404
+ "layers.4.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
405
+ "layers.4.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
406
+ "layers.4.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
407
+ "layers.4.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
408
+ "layers.4.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
409
+ "layers.4.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
410
+ "layers.4.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
411
+ "layers.4.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
412
+ "layers.4.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
413
+ "layers.4.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
414
+ "layers.4.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
415
+ "layers.4.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
416
+ "layers.5.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
417
+ "layers.5.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
418
+ "layers.5.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
419
+ "layers.5.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
420
+ "layers.5.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
421
+ "layers.5.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
422
+ "layers.5.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
423
+ "layers.5.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
424
+ "layers.5.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
425
+ "layers.5.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
426
+ "layers.5.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
427
+ "layers.5.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
428
+ "layers.5.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
429
+ "layers.5.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
430
+ "layers.5.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
431
+ "layers.6.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
432
+ "layers.6.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
433
+ "layers.6.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
434
+ "layers.6.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
435
+ "layers.6.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
436
+ "layers.6.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
437
+ "layers.6.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
438
+ "layers.6.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
439
+ "layers.6.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
440
+ "layers.6.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
441
+ "layers.6.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
442
+ "layers.6.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
443
+ "layers.6.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
444
+ "layers.6.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
445
+ "layers.6.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
446
+ "layers.7.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
447
+ "layers.7.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
448
+ "layers.7.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
449
+ "layers.7.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
450
+ "layers.7.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
451
+ "layers.7.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
452
+ "layers.7.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
453
+ "layers.7.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
454
+ "layers.7.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
455
+ "layers.7.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
456
+ "layers.7.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
457
+ "layers.7.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
458
+ "layers.7.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
459
+ "layers.7.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
460
+ "layers.7.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
461
+ "layers.8.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
462
+ "layers.8.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
463
+ "layers.8.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
464
+ "layers.8.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
465
+ "layers.8.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
466
+ "layers.8.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
467
+ "layers.8.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
468
+ "layers.8.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
469
+ "layers.8.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
470
+ "layers.8.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
471
+ "layers.8.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
472
+ "layers.8.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
473
+ "layers.8.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
474
+ "layers.8.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
475
+ "layers.8.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
476
+ "layers.9.adaLN_modulation.0.bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
477
+ "layers.9.adaLN_modulation.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
478
+ "layers.9.attention.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
479
+ "layers.9.attention.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
480
+ "layers.9.attention.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
481
+ "layers.9.attention.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
482
+ "layers.9.attention.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
483
+ "layers.9.attention.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
484
+ "layers.9.attention_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
485
+ "layers.9.attention_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
486
+ "layers.9.feed_forward.w1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
487
+ "layers.9.feed_forward.w2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
488
+ "layers.9.feed_forward.w3.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
489
+ "layers.9.ffn_norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
490
+ "layers.9.ffn_norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
491
+ "noise_refiner.0.adaLN_modulation.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
492
+ "noise_refiner.0.adaLN_modulation.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
493
+ "noise_refiner.0.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
494
+ "noise_refiner.0.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
495
+ "noise_refiner.0.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
496
+ "noise_refiner.0.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
497
+ "noise_refiner.0.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
498
+ "noise_refiner.0.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
499
+ "noise_refiner.0.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
500
+ "noise_refiner.0.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
501
+ "noise_refiner.0.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
502
+ "noise_refiner.0.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
503
+ "noise_refiner.0.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
504
+ "noise_refiner.0.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
505
+ "noise_refiner.0.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
506
+ "noise_refiner.1.adaLN_modulation.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
507
+ "noise_refiner.1.adaLN_modulation.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
508
+ "noise_refiner.1.attention.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
509
+ "noise_refiner.1.attention.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
510
+ "noise_refiner.1.attention.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
511
+ "noise_refiner.1.attention.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
512
+ "noise_refiner.1.attention.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
513
+ "noise_refiner.1.attention.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
514
+ "noise_refiner.1.attention_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
515
+ "noise_refiner.1.attention_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
516
+ "noise_refiner.1.feed_forward.w1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
517
+ "noise_refiner.1.feed_forward.w2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
518
+ "noise_refiner.1.feed_forward.w3.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
519
+ "noise_refiner.1.ffn_norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
520
+ "noise_refiner.1.ffn_norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
521
+ "t_embedder.mlp.0.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
522
+ "t_embedder.mlp.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
523
+ "t_embedder.mlp.2.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
524
+ "t_embedder.mlp.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors"
525
+ }
526
+ }
vae/config.json ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_class_name": "AutoencoderKLQwenImage",
3
+ "_diffusers_version": "0.36.0.dev0",
4
+ "attn_scales": [],
5
+ "base_dim": 96,
6
+ "dim_mult": [
7
+ 1,
8
+ 2,
9
+ 4,
10
+ 4
11
+ ],
12
+ "dropout": 0.0,
13
+ "input_channels": 4,
14
+ "latents_mean": null,
15
+ "latents_std": null,
16
+ "scaling_factor": 8.0064,
17
+ "shift_factor": 0.0,
18
+ "num_res_blocks": 2,
19
+ "temperal_downsample": [
20
+ false,
21
+ true,
22
+ true
23
+ ],
24
+ "z_dim": 16
25
+ }