ymyy307 commited on
Commit
0ed6b0e
·
verified ·
1 Parent(s): 5a5e271

Upload folder using huggingface_hub (part 2)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. diffsynth.egg-info/PKG-INFO +0 -0
  2. diffsynth.egg-info/SOURCES.txt +271 -0
  3. diffsynth.egg-info/dependency_links.txt +1 -0
  4. diffsynth.egg-info/requires.txt +64 -0
  5. diffsynth.egg-info/top_level.txt +1 -0
  6. diffsynth/utils/state_dict_converters/sdxl_text_encoder_2.py +345 -0
  7. diffsynth/utils/state_dict_converters/sdxl_vae.py +265 -0
  8. diffsynth/utils/state_dict_converters/stable_diffusion_text_encoder.py +7 -0
  9. diffsynth/utils/state_dict_converters/stable_diffusion_vae.py +18 -0
  10. diffsynth/utils/state_dict_converters/stable_diffusion_xl_text_encoder.py +13 -0
  11. diffsynth/utils/state_dict_converters/step1x_connector.py +7 -0
  12. diffsynth/utils/state_dict_converters/wan_video_animate_adapter.py +6 -0
  13. diffsynth/utils/state_dict_converters/wan_video_dit.py +83 -0
  14. diffsynth/utils/state_dict_converters/wan_video_image_encoder.py +8 -0
  15. diffsynth/utils/state_dict_converters/wan_video_mot.py +78 -0
  16. diffsynth/utils/state_dict_converters/wan_video_vace.py +3 -0
  17. diffsynth/utils/state_dict_converters/wan_video_vae.py +7 -0
  18. diffsynth/utils/state_dict_converters/wans2v_audio_encoder.py +12 -0
  19. diffsynth/utils/state_dict_converters/z_image_dit.py +3 -0
  20. diffsynth/utils/state_dict_converters/z_image_text_encoder.py +6 -0
  21. diffsynth/utils/tile/__init__.py +1 -0
  22. diffsynth/utils/tile/tile_worker.py +55 -0
  23. diffsynth/utils/xfuser/__init__.py +1 -0
  24. diffsynth/utils/xfuser/xdit_context_parallel.py +219 -0
  25. diffsynth/version.py +5 -0
  26. docs/en/.readthedocs.yaml +28 -0
  27. docs/en/API_Reference/core/attention.md +80 -0
  28. docs/en/API_Reference/core/data.md +151 -0
  29. docs/en/API_Reference/core/gradient.md +69 -0
  30. docs/en/API_Reference/core/loader.md +141 -0
  31. docs/en/API_Reference/core/quant.md +295 -0
  32. docs/en/API_Reference/core/vram.md +66 -0
  33. docs/en/Developer_Guide/Building_a_Pipeline.md +254 -0
  34. docs/en/Developer_Guide/Enabling_VRAM_management.md +455 -0
  35. docs/en/Developer_Guide/Integrating_Quantization_Backend.md +473 -0
  36. docs/en/Developer_Guide/Integrating_Your_Model.md +186 -0
  37. docs/en/Developer_Guide/Training_Diffusion_Models.md +66 -0
  38. docs/en/Diffusion_Templates/Introducing_Diffusion_Templates.md +76 -0
  39. docs/en/Diffusion_Templates/Template_Model_Inference.md +333 -0
  40. docs/en/Diffusion_Templates/Template_Model_Training.md +344 -0
  41. docs/en/Diffusion_Templates/Understanding_Diffusion_Templates.md +62 -0
  42. docs/en/Makefile +20 -0
  43. docs/en/Model_Details/ACE-Step.md +166 -0
  44. docs/en/Model_Details/Anima.md +140 -0
  45. docs/en/Model_Details/Boogu-Image.md +148 -0
  46. docs/en/Model_Details/ERNIE-Image.md +135 -0
  47. docs/en/Model_Details/FLUX.md +185 -0
  48. docs/en/Model_Details/FLUX2.md +155 -0
  49. docs/en/Model_Details/HiDream-O1-Image.md +143 -0
  50. docs/en/Model_Details/Ideogram-4.md +151 -0
diffsynth.egg-info/PKG-INFO ADDED
The diff for this file is too large to render. See raw diff
 
diffsynth.egg-info/SOURCES.txt ADDED
@@ -0,0 +1,271 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ LICENSE
2
+ README.md
3
+ pyproject.toml
4
+ diffsynth/__init__.py
5
+ diffsynth/version.py
6
+ diffsynth.egg-info/PKG-INFO
7
+ diffsynth.egg-info/SOURCES.txt
8
+ diffsynth.egg-info/dependency_links.txt
9
+ diffsynth.egg-info/requires.txt
10
+ diffsynth.egg-info/top_level.txt
11
+ diffsynth/configs/__init__.py
12
+ diffsynth/configs/model_configs.py
13
+ diffsynth/configs/vram_management_module_maps.py
14
+ diffsynth/core/__init__.py
15
+ diffsynth/core/attention/__init__.py
16
+ diffsynth/core/attention/attention.py
17
+ diffsynth/core/data/__init__.py
18
+ diffsynth/core/data/operators.py
19
+ diffsynth/core/data/unified_dataset.py
20
+ diffsynth/core/device/__init__.py
21
+ diffsynth/core/device/npu_compatible_device.py
22
+ diffsynth/core/gradient/__init__.py
23
+ diffsynth/core/gradient/gradient_checkpoint.py
24
+ diffsynth/core/loader/__init__.py
25
+ diffsynth/core/loader/config.py
26
+ diffsynth/core/loader/file.py
27
+ diffsynth/core/loader/model.py
28
+ diffsynth/core/npu_patch/npu_fused_operator.py
29
+ diffsynth/core/offload_training/__init__.py
30
+ diffsynth/core/offload_training/manager.py
31
+ diffsynth/core/offload_training/memory_buffer.py
32
+ diffsynth/core/offload_training/offloader.py
33
+ diffsynth/core/quant/__init__.py
34
+ diffsynth/core/quant/base.py
35
+ diffsynth/core/quant/config.py
36
+ diffsynth/core/quant/backends/__init__.py
37
+ diffsynth/core/quant/backends/bitsandbytes.py
38
+ diffsynth/core/quant/backends/comfy_kitchen.py
39
+ diffsynth/core/quant/backends/torchao.py
40
+ diffsynth/core/vram/__init__.py
41
+ diffsynth/core/vram/disk_map.py
42
+ diffsynth/core/vram/initialization.py
43
+ diffsynth/core/vram/layers.py
44
+ diffsynth/diffusion/__init__.py
45
+ diffsynth/diffusion/base_pipeline.py
46
+ diffsynth/diffusion/ddim_scheduler.py
47
+ diffsynth/diffusion/dmd2.py
48
+ diffsynth/diffusion/flow_match.py
49
+ diffsynth/diffusion/logger.py
50
+ diffsynth/diffusion/loss.py
51
+ diffsynth/diffusion/parsers.py
52
+ diffsynth/diffusion/runner.py
53
+ diffsynth/diffusion/template.py
54
+ diffsynth/diffusion/training_module.py
55
+ diffsynth/metrics/__init__.py
56
+ diffsynth/metrics/aesthetic.py
57
+ diffsynth/metrics/base.py
58
+ diffsynth/metrics/bioclip.py
59
+ diffsynth/metrics/clip.py
60
+ diffsynth/metrics/fid.py
61
+ diffsynth/metrics/hpsv2.py
62
+ diffsynth/metrics/hpsv3.py
63
+ diffsynth/metrics/image_reward.py
64
+ diffsynth/metrics/lpips.py
65
+ diffsynth/metrics/pickscore.py
66
+ diffsynth/metrics/qwen_image_bench.py
67
+ diffsynth/metrics/unified_reward_2.py
68
+ diffsynth/metrics/unified_reward_edit.py
69
+ diffsynth/models/ace_step_conditioner.py
70
+ diffsynth/models/ace_step_dit.py
71
+ diffsynth/models/ace_step_residual_fsq.py
72
+ diffsynth/models/ace_step_text_encoder.py
73
+ diffsynth/models/ace_step_tokenizer.py
74
+ diffsynth/models/ace_step_vae.py
75
+ diffsynth/models/aesthetic.py
76
+ diffsynth/models/anima_dit.py
77
+ diffsynth/models/bioclip.py
78
+ diffsynth/models/boogu_image_dit.py
79
+ diffsynth/models/clip.py
80
+ diffsynth/models/demucs.py
81
+ diffsynth/models/dinov3_image_encoder.py
82
+ diffsynth/models/ernie_image_dit.py
83
+ diffsynth/models/ernie_image_text_encoder.py
84
+ diffsynth/models/fid.py
85
+ diffsynth/models/flux2_dit.py
86
+ diffsynth/models/flux2_text_encoder.py
87
+ diffsynth/models/flux2_vae.py
88
+ diffsynth/models/flux_controlnet.py
89
+ diffsynth/models/flux_dit.py
90
+ diffsynth/models/flux_infiniteyou.py
91
+ diffsynth/models/flux_ipadapter.py
92
+ diffsynth/models/flux_lora_encoder.py
93
+ diffsynth/models/flux_lora_patcher.py
94
+ diffsynth/models/flux_redux.py
95
+ diffsynth/models/flux_text_encoder_clip.py
96
+ diffsynth/models/flux_text_encoder_t5.py
97
+ diffsynth/models/flux_vae.py
98
+ diffsynth/models/flux_value_control.py
99
+ diffsynth/models/general_modules.py
100
+ diffsynth/models/hidream_common.py
101
+ diffsynth/models/hidream_o1_image_dit.py
102
+ diffsynth/models/hpsv2.py
103
+ diffsynth/models/hpsv3.py
104
+ diffsynth/models/ideogram4_dit.py
105
+ diffsynth/models/ideogram4_text_encoder.py
106
+ diffsynth/models/ideogram4_vae.py
107
+ diffsynth/models/image_reward.py
108
+ diffsynth/models/joyai_image_dit.py
109
+ diffsynth/models/joyai_image_text_encoder.py
110
+ diffsynth/models/krea2_dit.py
111
+ diffsynth/models/krea2_text_encoder.py
112
+ diffsynth/models/lingbot_video_dit.py
113
+ diffsynth/models/longcat_video_dit.py
114
+ diffsynth/models/lpips.py
115
+ diffsynth/models/ltx2_audio_vae.py
116
+ diffsynth/models/ltx2_common.py
117
+ diffsynth/models/ltx2_dit.py
118
+ diffsynth/models/ltx2_text_encoder.py
119
+ diffsynth/models/ltx2_upsampler.py
120
+ diffsynth/models/ltx2_video_vae.py
121
+ diffsynth/models/minimax_h3_audio_vae.py
122
+ diffsynth/models/minimax_h3_dit.py
123
+ diffsynth/models/minimax_h3_dit_comfy.py
124
+ diffsynth/models/minimax_h3_text_encoder.py
125
+ diffsynth/models/minimax_h3_video_vae.py
126
+ diffsynth/models/minimax_music3_condition_encoder.py
127
+ diffsynth/models/minimax_music3_dit.py
128
+ diffsynth/models/minimax_music3_rvq_depth_decoder.py
129
+ diffsynth/models/minimax_music3_text_encoder.py
130
+ diffsynth/models/minimax_music3_vocoder.py
131
+ diffsynth/models/model_loader.py
132
+ diffsynth/models/mova_audio_dit.py
133
+ diffsynth/models/mova_audio_vae.py
134
+ diffsynth/models/mova_dual_tower_bridge.py
135
+ diffsynth/models/nexus_gen.py
136
+ diffsynth/models/nexus_gen_ar_model.py
137
+ diffsynth/models/nexus_gen_projector.py
138
+ diffsynth/models/pickscore.py
139
+ diffsynth/models/qwen_image_bench.py
140
+ diffsynth/models/qwen_image_controlnet.py
141
+ diffsynth/models/qwen_image_dit.py
142
+ diffsynth/models/qwen_image_image2lora.py
143
+ diffsynth/models/qwen_image_text_encoder.py
144
+ diffsynth/models/qwen_image_vae.py
145
+ diffsynth/models/qwen_video_edit_dit.py
146
+ diffsynth/models/sd_text_encoder.py
147
+ diffsynth/models/siglip2_image_encoder.py
148
+ diffsynth/models/stable_diffusion_text_encoder.py
149
+ diffsynth/models/stable_diffusion_unet.py
150
+ diffsynth/models/stable_diffusion_vae.py
151
+ diffsynth/models/stable_diffusion_xl_text_encoder.py
152
+ diffsynth/models/stable_diffusion_xl_unet.py
153
+ diffsynth/models/step1x_connector.py
154
+ diffsynth/models/step1x_text_encoder.py
155
+ diffsynth/models/unified_reward_2.py
156
+ diffsynth/models/unified_reward_edit.py
157
+ diffsynth/models/wan_animate_2_dit.py
158
+ diffsynth/models/wan_video_animate_adapter.py
159
+ diffsynth/models/wan_video_camera_controller.py
160
+ diffsynth/models/wan_video_dit.py
161
+ diffsynth/models/wan_video_dit_s2v.py
162
+ diffsynth/models/wan_video_image_encoder.py
163
+ diffsynth/models/wan_video_mot.py
164
+ diffsynth/models/wan_video_motion_controller.py
165
+ diffsynth/models/wan_video_text_encoder.py
166
+ diffsynth/models/wan_video_vace.py
167
+ diffsynth/models/wan_video_vae.py
168
+ diffsynth/models/wantodance.py
169
+ diffsynth/models/wav2vec.py
170
+ diffsynth/models/z_image_controlnet.py
171
+ diffsynth/models/z_image_dit.py
172
+ diffsynth/models/z_image_image2lora.py
173
+ diffsynth/models/z_image_text_encoder.py
174
+ diffsynth/pipelines/ace_step.py
175
+ diffsynth/pipelines/anima_image.py
176
+ diffsynth/pipelines/boogu_image.py
177
+ diffsynth/pipelines/ernie_image.py
178
+ diffsynth/pipelines/flux2_image.py
179
+ diffsynth/pipelines/flux_image.py
180
+ diffsynth/pipelines/hidream_o1_image.py
181
+ diffsynth/pipelines/ideogram4.py
182
+ diffsynth/pipelines/joyai_image.py
183
+ diffsynth/pipelines/krea2.py
184
+ diffsynth/pipelines/lingbot_video.py
185
+ diffsynth/pipelines/ltx2_audio_video.py
186
+ diffsynth/pipelines/minimax_h3_audio_video.py
187
+ diffsynth/pipelines/minimax_music3.py
188
+ diffsynth/pipelines/mova_audio_video.py
189
+ diffsynth/pipelines/qwen_image.py
190
+ diffsynth/pipelines/qwen_video_edit.py
191
+ diffsynth/pipelines/stable_diffusion.py
192
+ diffsynth/pipelines/stable_diffusion_xl.py
193
+ diffsynth/pipelines/wan_video.py
194
+ diffsynth/pipelines/z_image.py
195
+ diffsynth/utils/controlnet/__init__.py
196
+ diffsynth/utils/controlnet/annotator.py
197
+ diffsynth/utils/controlnet/controlnet_input.py
198
+ diffsynth/utils/data/__init__.py
199
+ diffsynth/utils/data/audio.py
200
+ diffsynth/utils/data/audio_video.py
201
+ diffsynth/utils/data/media_io_ltx2.py
202
+ diffsynth/utils/data/minimax_h3.py
203
+ diffsynth/utils/demucs/__init__.py
204
+ diffsynth/utils/dequantizer/__init__.py
205
+ diffsynth/utils/lora/__init__.py
206
+ diffsynth/utils/lora/flux.py
207
+ diffsynth/utils/lora/flux_timestep.py
208
+ diffsynth/utils/lora/general.py
209
+ diffsynth/utils/lora/krea2.py
210
+ diffsynth/utils/lora/merge.py
211
+ diffsynth/utils/lora/minimax_h3.py
212
+ diffsynth/utils/lora/reset_rank.py
213
+ diffsynth/utils/lora/sdxl.py
214
+ diffsynth/utils/quant/serialization.py
215
+ diffsynth/utils/ses/__init__.py
216
+ diffsynth/utils/ses/ses.py
217
+ diffsynth/utils/state_dict_converters/__init__.py
218
+ diffsynth/utils/state_dict_converters/ace_step_conditioner.py
219
+ diffsynth/utils/state_dict_converters/ace_step_dit.py
220
+ diffsynth/utils/state_dict_converters/ace_step_text_encoder.py
221
+ diffsynth/utils/state_dict_converters/ace_step_tokenizer.py
222
+ diffsynth/utils/state_dict_converters/anima_dit.py
223
+ diffsynth/utils/state_dict_converters/dino_v3.py
224
+ diffsynth/utils/state_dict_converters/ernie_image_text_encoder.py
225
+ diffsynth/utils/state_dict_converters/flux2_text_encoder.py
226
+ diffsynth/utils/state_dict_converters/flux_controlnet.py
227
+ diffsynth/utils/state_dict_converters/flux_dit.py
228
+ diffsynth/utils/state_dict_converters/flux_infiniteyou.py
229
+ diffsynth/utils/state_dict_converters/flux_ipadapter.py
230
+ diffsynth/utils/state_dict_converters/flux_text_encoder_clip.py
231
+ diffsynth/utils/state_dict_converters/flux_text_encoder_t5.py
232
+ diffsynth/utils/state_dict_converters/flux_vae.py
233
+ diffsynth/utils/state_dict_converters/ideogram4_text_encoder.py
234
+ diffsynth/utils/state_dict_converters/image_metrics.py
235
+ diffsynth/utils/state_dict_converters/joyai_image_text_encoder.py
236
+ diffsynth/utils/state_dict_converters/krea2_dit.py
237
+ diffsynth/utils/state_dict_converters/krea2_text_encoder.py
238
+ diffsynth/utils/state_dict_converters/lingbot_video_dit.py
239
+ diffsynth/utils/state_dict_converters/ltx2_audio_vae.py
240
+ diffsynth/utils/state_dict_converters/ltx2_dit.py
241
+ diffsynth/utils/state_dict_converters/ltx2_text_encoder.py
242
+ diffsynth/utils/state_dict_converters/ltx2_video_vae.py
243
+ diffsynth/utils/state_dict_converters/minimax_h3_audio_vae.py
244
+ diffsynth/utils/state_dict_converters/minimax_h3_text_encoder.py
245
+ diffsynth/utils/state_dict_converters/minimax_h3_video_vae.py
246
+ diffsynth/utils/state_dict_converters/minimax_music3_text_encoder.py
247
+ diffsynth/utils/state_dict_converters/nexus_gen.py
248
+ diffsynth/utils/state_dict_converters/nexus_gen_projector.py
249
+ diffsynth/utils/state_dict_converters/qwen_image_text_encoder.py
250
+ diffsynth/utils/state_dict_converters/qwen_video_edit.py
251
+ diffsynth/utils/state_dict_converters/sdxl.py
252
+ diffsynth/utils/state_dict_converters/sdxl_text_encoder.py
253
+ diffsynth/utils/state_dict_converters/sdxl_text_encoder_2.py
254
+ diffsynth/utils/state_dict_converters/sdxl_vae.py
255
+ diffsynth/utils/state_dict_converters/stable_diffusion_text_encoder.py
256
+ diffsynth/utils/state_dict_converters/stable_diffusion_vae.py
257
+ diffsynth/utils/state_dict_converters/stable_diffusion_xl_text_encoder.py
258
+ diffsynth/utils/state_dict_converters/step1x_connector.py
259
+ diffsynth/utils/state_dict_converters/wan_video_animate_adapter.py
260
+ diffsynth/utils/state_dict_converters/wan_video_dit.py
261
+ diffsynth/utils/state_dict_converters/wan_video_image_encoder.py
262
+ diffsynth/utils/state_dict_converters/wan_video_mot.py
263
+ diffsynth/utils/state_dict_converters/wan_video_vace.py
264
+ diffsynth/utils/state_dict_converters/wan_video_vae.py
265
+ diffsynth/utils/state_dict_converters/wans2v_audio_encoder.py
266
+ diffsynth/utils/state_dict_converters/z_image_dit.py
267
+ diffsynth/utils/state_dict_converters/z_image_text_encoder.py
268
+ diffsynth/utils/tile/__init__.py
269
+ diffsynth/utils/tile/tile_worker.py
270
+ diffsynth/utils/xfuser/__init__.py
271
+ diffsynth/utils/xfuser/xdit_context_parallel.py
diffsynth.egg-info/dependency_links.txt ADDED
@@ -0,0 +1 @@
 
 
1
+
diffsynth.egg-info/requires.txt ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ torch==2.6.0
2
+ torchvision
3
+ transformers
4
+ imageio[ffmpeg]
5
+ safetensors
6
+ einops
7
+ modelscope
8
+ ftfy
9
+ pandas
10
+ accelerate
11
+ peft
12
+
13
+ [all]
14
+ av
15
+ torchaudio
16
+ torchcodec
17
+ librosa
18
+ bitsandbytes
19
+ comfy-kitchen
20
+ torchao>=0.16
21
+ deepspeed
22
+ tensorboard
23
+ swanlab
24
+ wandb
25
+
26
+ [audio]
27
+ av
28
+ torchaudio
29
+ torchcodec
30
+ librosa
31
+
32
+ [infiniteyou]
33
+ insightface
34
+ facexlib
35
+
36
+ [logger]
37
+ tensorboard
38
+ swanlab
39
+ wandb
40
+
41
+ [nexusgen]
42
+ qwen_vl_utils
43
+ transformers==4.49.0
44
+
45
+ [npu]
46
+ torch==2.7.1+cpu
47
+ torch-npu==2.7.1
48
+ torchvision==0.22.1+cpu
49
+
50
+ [npu_aarch64]
51
+ torch==2.7.1
52
+ torch-npu==2.7.1
53
+ torchvision==0.22.1
54
+
55
+ [quant]
56
+ bitsandbytes
57
+ comfy-kitchen
58
+ torchao>=0.16
59
+
60
+ [ses]
61
+ pywt
62
+
63
+ [training]
64
+ deepspeed
diffsynth.egg-info/top_level.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ diffsynth
diffsynth/utils/state_dict_converters/sdxl_text_encoder_2.py ADDED
@@ -0,0 +1,345 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ rename_dict = {
2
+ "model.text_model.embeddings.position_embedding.weight": "conditioner.embedders.1.model.positional_embedding",
3
+ "model.text_model.embeddings.token_embedding.weight": "conditioner.embedders.1.model.token_embedding.weight",
4
+ "model.text_model.encoder.layers.0.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.0.ln_1.bias",
5
+ "model.text_model.encoder.layers.0.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.0.ln_1.weight",
6
+ "model.text_model.encoder.layers.0.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.0.ln_2.bias",
7
+ "model.text_model.encoder.layers.0.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.0.ln_2.weight",
8
+ "model.text_model.encoder.layers.0.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_fc.bias",
9
+ "model.text_model.encoder.layers.0.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_fc.weight",
10
+ "model.text_model.encoder.layers.0.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_proj.bias",
11
+ "model.text_model.encoder.layers.0.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_proj.weight",
12
+ "model.text_model.encoder.layers.0.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.0.attn.out_proj.bias",
13
+ "model.text_model.encoder.layers.0.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.0.attn.out_proj.weight",
14
+ "model.text_model.encoder.layers.1.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.1.ln_1.bias",
15
+ "model.text_model.encoder.layers.1.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.1.ln_1.weight",
16
+ "model.text_model.encoder.layers.1.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.1.ln_2.bias",
17
+ "model.text_model.encoder.layers.1.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.1.ln_2.weight",
18
+ "model.text_model.encoder.layers.1.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_fc.bias",
19
+ "model.text_model.encoder.layers.1.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_fc.weight",
20
+ "model.text_model.encoder.layers.1.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_proj.bias",
21
+ "model.text_model.encoder.layers.1.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_proj.weight",
22
+ "model.text_model.encoder.layers.1.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.1.attn.out_proj.bias",
23
+ "model.text_model.encoder.layers.1.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.1.attn.out_proj.weight",
24
+ "model.text_model.encoder.layers.10.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.10.ln_1.bias",
25
+ "model.text_model.encoder.layers.10.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.10.ln_1.weight",
26
+ "model.text_model.encoder.layers.10.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.10.ln_2.bias",
27
+ "model.text_model.encoder.layers.10.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.10.ln_2.weight",
28
+ "model.text_model.encoder.layers.10.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_fc.bias",
29
+ "model.text_model.encoder.layers.10.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_fc.weight",
30
+ "model.text_model.encoder.layers.10.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_proj.bias",
31
+ "model.text_model.encoder.layers.10.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_proj.weight",
32
+ "model.text_model.encoder.layers.10.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.10.attn.out_proj.bias",
33
+ "model.text_model.encoder.layers.10.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.10.attn.out_proj.weight",
34
+ "model.text_model.encoder.layers.11.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.11.ln_1.bias",
35
+ "model.text_model.encoder.layers.11.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.11.ln_1.weight",
36
+ "model.text_model.encoder.layers.11.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.11.ln_2.bias",
37
+ "model.text_model.encoder.layers.11.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.11.ln_2.weight",
38
+ "model.text_model.encoder.layers.11.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_fc.bias",
39
+ "model.text_model.encoder.layers.11.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_fc.weight",
40
+ "model.text_model.encoder.layers.11.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_proj.bias",
41
+ "model.text_model.encoder.layers.11.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_proj.weight",
42
+ "model.text_model.encoder.layers.11.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.11.attn.out_proj.bias",
43
+ "model.text_model.encoder.layers.11.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.11.attn.out_proj.weight",
44
+ "model.text_model.encoder.layers.12.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.12.ln_1.bias",
45
+ "model.text_model.encoder.layers.12.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.12.ln_1.weight",
46
+ "model.text_model.encoder.layers.12.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.12.ln_2.bias",
47
+ "model.text_model.encoder.layers.12.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.12.ln_2.weight",
48
+ "model.text_model.encoder.layers.12.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_fc.bias",
49
+ "model.text_model.encoder.layers.12.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_fc.weight",
50
+ "model.text_model.encoder.layers.12.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_proj.bias",
51
+ "model.text_model.encoder.layers.12.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_proj.weight",
52
+ "model.text_model.encoder.layers.12.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.12.attn.out_proj.bias",
53
+ "model.text_model.encoder.layers.12.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.12.attn.out_proj.weight",
54
+ "model.text_model.encoder.layers.13.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.13.ln_1.bias",
55
+ "model.text_model.encoder.layers.13.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.13.ln_1.weight",
56
+ "model.text_model.encoder.layers.13.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.13.ln_2.bias",
57
+ "model.text_model.encoder.layers.13.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.13.ln_2.weight",
58
+ "model.text_model.encoder.layers.13.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_fc.bias",
59
+ "model.text_model.encoder.layers.13.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_fc.weight",
60
+ "model.text_model.encoder.layers.13.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_proj.bias",
61
+ "model.text_model.encoder.layers.13.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_proj.weight",
62
+ "model.text_model.encoder.layers.13.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.13.attn.out_proj.bias",
63
+ "model.text_model.encoder.layers.13.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.13.attn.out_proj.weight",
64
+ "model.text_model.encoder.layers.14.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.14.ln_1.bias",
65
+ "model.text_model.encoder.layers.14.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.14.ln_1.weight",
66
+ "model.text_model.encoder.layers.14.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.14.ln_2.bias",
67
+ "model.text_model.encoder.layers.14.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.14.ln_2.weight",
68
+ "model.text_model.encoder.layers.14.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_fc.bias",
69
+ "model.text_model.encoder.layers.14.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_fc.weight",
70
+ "model.text_model.encoder.layers.14.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_proj.bias",
71
+ "model.text_model.encoder.layers.14.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_proj.weight",
72
+ "model.text_model.encoder.layers.14.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.14.attn.out_proj.bias",
73
+ "model.text_model.encoder.layers.14.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.14.attn.out_proj.weight",
74
+ "model.text_model.encoder.layers.15.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.15.ln_1.bias",
75
+ "model.text_model.encoder.layers.15.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.15.ln_1.weight",
76
+ "model.text_model.encoder.layers.15.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.15.ln_2.bias",
77
+ "model.text_model.encoder.layers.15.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.15.ln_2.weight",
78
+ "model.text_model.encoder.layers.15.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_fc.bias",
79
+ "model.text_model.encoder.layers.15.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_fc.weight",
80
+ "model.text_model.encoder.layers.15.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_proj.bias",
81
+ "model.text_model.encoder.layers.15.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_proj.weight",
82
+ "model.text_model.encoder.layers.15.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.15.attn.out_proj.bias",
83
+ "model.text_model.encoder.layers.15.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.15.attn.out_proj.weight",
84
+ "model.text_model.encoder.layers.16.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.16.ln_1.bias",
85
+ "model.text_model.encoder.layers.16.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.16.ln_1.weight",
86
+ "model.text_model.encoder.layers.16.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.16.ln_2.bias",
87
+ "model.text_model.encoder.layers.16.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.16.ln_2.weight",
88
+ "model.text_model.encoder.layers.16.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_fc.bias",
89
+ "model.text_model.encoder.layers.16.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_fc.weight",
90
+ "model.text_model.encoder.layers.16.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_proj.bias",
91
+ "model.text_model.encoder.layers.16.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_proj.weight",
92
+ "model.text_model.encoder.layers.16.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.16.attn.out_proj.bias",
93
+ "model.text_model.encoder.layers.16.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.16.attn.out_proj.weight",
94
+ "model.text_model.encoder.layers.17.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.17.ln_1.bias",
95
+ "model.text_model.encoder.layers.17.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.17.ln_1.weight",
96
+ "model.text_model.encoder.layers.17.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.17.ln_2.bias",
97
+ "model.text_model.encoder.layers.17.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.17.ln_2.weight",
98
+ "model.text_model.encoder.layers.17.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_fc.bias",
99
+ "model.text_model.encoder.layers.17.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_fc.weight",
100
+ "model.text_model.encoder.layers.17.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_proj.bias",
101
+ "model.text_model.encoder.layers.17.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_proj.weight",
102
+ "model.text_model.encoder.layers.17.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.17.attn.out_proj.bias",
103
+ "model.text_model.encoder.layers.17.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.17.attn.out_proj.weight",
104
+ "model.text_model.encoder.layers.18.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.18.ln_1.bias",
105
+ "model.text_model.encoder.layers.18.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.18.ln_1.weight",
106
+ "model.text_model.encoder.layers.18.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.18.ln_2.bias",
107
+ "model.text_model.encoder.layers.18.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.18.ln_2.weight",
108
+ "model.text_model.encoder.layers.18.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_fc.bias",
109
+ "model.text_model.encoder.layers.18.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_fc.weight",
110
+ "model.text_model.encoder.layers.18.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_proj.bias",
111
+ "model.text_model.encoder.layers.18.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_proj.weight",
112
+ "model.text_model.encoder.layers.18.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.18.attn.out_proj.bias",
113
+ "model.text_model.encoder.layers.18.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.18.attn.out_proj.weight",
114
+ "model.text_model.encoder.layers.19.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.19.ln_1.bias",
115
+ "model.text_model.encoder.layers.19.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.19.ln_1.weight",
116
+ "model.text_model.encoder.layers.19.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.19.ln_2.bias",
117
+ "model.text_model.encoder.layers.19.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.19.ln_2.weight",
118
+ "model.text_model.encoder.layers.19.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_fc.bias",
119
+ "model.text_model.encoder.layers.19.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_fc.weight",
120
+ "model.text_model.encoder.layers.19.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_proj.bias",
121
+ "model.text_model.encoder.layers.19.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_proj.weight",
122
+ "model.text_model.encoder.layers.19.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.19.attn.out_proj.bias",
123
+ "model.text_model.encoder.layers.19.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.19.attn.out_proj.weight",
124
+ "model.text_model.encoder.layers.2.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.2.ln_1.bias",
125
+ "model.text_model.encoder.layers.2.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.2.ln_1.weight",
126
+ "model.text_model.encoder.layers.2.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.2.ln_2.bias",
127
+ "model.text_model.encoder.layers.2.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.2.ln_2.weight",
128
+ "model.text_model.encoder.layers.2.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_fc.bias",
129
+ "model.text_model.encoder.layers.2.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_fc.weight",
130
+ "model.text_model.encoder.layers.2.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_proj.bias",
131
+ "model.text_model.encoder.layers.2.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_proj.weight",
132
+ "model.text_model.encoder.layers.2.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.2.attn.out_proj.bias",
133
+ "model.text_model.encoder.layers.2.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.2.attn.out_proj.weight",
134
+ "model.text_model.encoder.layers.20.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.20.ln_1.bias",
135
+ "model.text_model.encoder.layers.20.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.20.ln_1.weight",
136
+ "model.text_model.encoder.layers.20.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.20.ln_2.bias",
137
+ "model.text_model.encoder.layers.20.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.20.ln_2.weight",
138
+ "model.text_model.encoder.layers.20.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_fc.bias",
139
+ "model.text_model.encoder.layers.20.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_fc.weight",
140
+ "model.text_model.encoder.layers.20.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_proj.bias",
141
+ "model.text_model.encoder.layers.20.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_proj.weight",
142
+ "model.text_model.encoder.layers.20.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.20.attn.out_proj.bias",
143
+ "model.text_model.encoder.layers.20.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.20.attn.out_proj.weight",
144
+ "model.text_model.encoder.layers.21.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.21.ln_1.bias",
145
+ "model.text_model.encoder.layers.21.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.21.ln_1.weight",
146
+ "model.text_model.encoder.layers.21.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.21.ln_2.bias",
147
+ "model.text_model.encoder.layers.21.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.21.ln_2.weight",
148
+ "model.text_model.encoder.layers.21.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_fc.bias",
149
+ "model.text_model.encoder.layers.21.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_fc.weight",
150
+ "model.text_model.encoder.layers.21.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_proj.bias",
151
+ "model.text_model.encoder.layers.21.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_proj.weight",
152
+ "model.text_model.encoder.layers.21.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.21.attn.out_proj.bias",
153
+ "model.text_model.encoder.layers.21.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.21.attn.out_proj.weight",
154
+ "model.text_model.encoder.layers.22.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.22.ln_1.bias",
155
+ "model.text_model.encoder.layers.22.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.22.ln_1.weight",
156
+ "model.text_model.encoder.layers.22.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.22.ln_2.bias",
157
+ "model.text_model.encoder.layers.22.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.22.ln_2.weight",
158
+ "model.text_model.encoder.layers.22.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_fc.bias",
159
+ "model.text_model.encoder.layers.22.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_fc.weight",
160
+ "model.text_model.encoder.layers.22.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_proj.bias",
161
+ "model.text_model.encoder.layers.22.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_proj.weight",
162
+ "model.text_model.encoder.layers.22.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.22.attn.out_proj.bias",
163
+ "model.text_model.encoder.layers.22.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.22.attn.out_proj.weight",
164
+ "model.text_model.encoder.layers.23.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.23.ln_1.bias",
165
+ "model.text_model.encoder.layers.23.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.23.ln_1.weight",
166
+ "model.text_model.encoder.layers.23.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.23.ln_2.bias",
167
+ "model.text_model.encoder.layers.23.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.23.ln_2.weight",
168
+ "model.text_model.encoder.layers.23.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_fc.bias",
169
+ "model.text_model.encoder.layers.23.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_fc.weight",
170
+ "model.text_model.encoder.layers.23.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_proj.bias",
171
+ "model.text_model.encoder.layers.23.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_proj.weight",
172
+ "model.text_model.encoder.layers.23.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.23.attn.out_proj.bias",
173
+ "model.text_model.encoder.layers.23.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.23.attn.out_proj.weight",
174
+ "model.text_model.encoder.layers.24.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.24.ln_1.bias",
175
+ "model.text_model.encoder.layers.24.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.24.ln_1.weight",
176
+ "model.text_model.encoder.layers.24.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.24.ln_2.bias",
177
+ "model.text_model.encoder.layers.24.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.24.ln_2.weight",
178
+ "model.text_model.encoder.layers.24.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_fc.bias",
179
+ "model.text_model.encoder.layers.24.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_fc.weight",
180
+ "model.text_model.encoder.layers.24.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_proj.bias",
181
+ "model.text_model.encoder.layers.24.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_proj.weight",
182
+ "model.text_model.encoder.layers.24.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.24.attn.out_proj.bias",
183
+ "model.text_model.encoder.layers.24.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.24.attn.out_proj.weight",
184
+ "model.text_model.encoder.layers.25.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.25.ln_1.bias",
185
+ "model.text_model.encoder.layers.25.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.25.ln_1.weight",
186
+ "model.text_model.encoder.layers.25.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.25.ln_2.bias",
187
+ "model.text_model.encoder.layers.25.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.25.ln_2.weight",
188
+ "model.text_model.encoder.layers.25.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_fc.bias",
189
+ "model.text_model.encoder.layers.25.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_fc.weight",
190
+ "model.text_model.encoder.layers.25.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_proj.bias",
191
+ "model.text_model.encoder.layers.25.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_proj.weight",
192
+ "model.text_model.encoder.layers.25.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.25.attn.out_proj.bias",
193
+ "model.text_model.encoder.layers.25.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.25.attn.out_proj.weight",
194
+ "model.text_model.encoder.layers.26.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.26.ln_1.bias",
195
+ "model.text_model.encoder.layers.26.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.26.ln_1.weight",
196
+ "model.text_model.encoder.layers.26.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.26.ln_2.bias",
197
+ "model.text_model.encoder.layers.26.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.26.ln_2.weight",
198
+ "model.text_model.encoder.layers.26.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_fc.bias",
199
+ "model.text_model.encoder.layers.26.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_fc.weight",
200
+ "model.text_model.encoder.layers.26.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_proj.bias",
201
+ "model.text_model.encoder.layers.26.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_proj.weight",
202
+ "model.text_model.encoder.layers.26.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.26.attn.out_proj.bias",
203
+ "model.text_model.encoder.layers.26.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.26.attn.out_proj.weight",
204
+ "model.text_model.encoder.layers.27.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.27.ln_1.bias",
205
+ "model.text_model.encoder.layers.27.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.27.ln_1.weight",
206
+ "model.text_model.encoder.layers.27.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.27.ln_2.bias",
207
+ "model.text_model.encoder.layers.27.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.27.ln_2.weight",
208
+ "model.text_model.encoder.layers.27.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_fc.bias",
209
+ "model.text_model.encoder.layers.27.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_fc.weight",
210
+ "model.text_model.encoder.layers.27.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_proj.bias",
211
+ "model.text_model.encoder.layers.27.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_proj.weight",
212
+ "model.text_model.encoder.layers.27.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.27.attn.out_proj.bias",
213
+ "model.text_model.encoder.layers.27.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.27.attn.out_proj.weight",
214
+ "model.text_model.encoder.layers.28.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.28.ln_1.bias",
215
+ "model.text_model.encoder.layers.28.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.28.ln_1.weight",
216
+ "model.text_model.encoder.layers.28.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.28.ln_2.bias",
217
+ "model.text_model.encoder.layers.28.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.28.ln_2.weight",
218
+ "model.text_model.encoder.layers.28.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_fc.bias",
219
+ "model.text_model.encoder.layers.28.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_fc.weight",
220
+ "model.text_model.encoder.layers.28.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_proj.bias",
221
+ "model.text_model.encoder.layers.28.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_proj.weight",
222
+ "model.text_model.encoder.layers.28.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.28.attn.out_proj.bias",
223
+ "model.text_model.encoder.layers.28.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.28.attn.out_proj.weight",
224
+ "model.text_model.encoder.layers.29.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.29.ln_1.bias",
225
+ "model.text_model.encoder.layers.29.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.29.ln_1.weight",
226
+ "model.text_model.encoder.layers.29.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.29.ln_2.bias",
227
+ "model.text_model.encoder.layers.29.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.29.ln_2.weight",
228
+ "model.text_model.encoder.layers.29.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_fc.bias",
229
+ "model.text_model.encoder.layers.29.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_fc.weight",
230
+ "model.text_model.encoder.layers.29.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_proj.bias",
231
+ "model.text_model.encoder.layers.29.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_proj.weight",
232
+ "model.text_model.encoder.layers.29.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.29.attn.out_proj.bias",
233
+ "model.text_model.encoder.layers.29.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.29.attn.out_proj.weight",
234
+ "model.text_model.encoder.layers.3.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.3.ln_1.bias",
235
+ "model.text_model.encoder.layers.3.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.3.ln_1.weight",
236
+ "model.text_model.encoder.layers.3.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.3.ln_2.bias",
237
+ "model.text_model.encoder.layers.3.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.3.ln_2.weight",
238
+ "model.text_model.encoder.layers.3.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_fc.bias",
239
+ "model.text_model.encoder.layers.3.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_fc.weight",
240
+ "model.text_model.encoder.layers.3.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_proj.bias",
241
+ "model.text_model.encoder.layers.3.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_proj.weight",
242
+ "model.text_model.encoder.layers.3.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.3.attn.out_proj.bias",
243
+ "model.text_model.encoder.layers.3.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.3.attn.out_proj.weight",
244
+ "model.text_model.encoder.layers.30.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.30.ln_1.bias",
245
+ "model.text_model.encoder.layers.30.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.30.ln_1.weight",
246
+ "model.text_model.encoder.layers.30.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.30.ln_2.bias",
247
+ "model.text_model.encoder.layers.30.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.30.ln_2.weight",
248
+ "model.text_model.encoder.layers.30.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_fc.bias",
249
+ "model.text_model.encoder.layers.30.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_fc.weight",
250
+ "model.text_model.encoder.layers.30.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_proj.bias",
251
+ "model.text_model.encoder.layers.30.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_proj.weight",
252
+ "model.text_model.encoder.layers.30.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.30.attn.out_proj.bias",
253
+ "model.text_model.encoder.layers.30.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.30.attn.out_proj.weight",
254
+ "model.text_model.encoder.layers.31.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.31.ln_1.bias",
255
+ "model.text_model.encoder.layers.31.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.31.ln_1.weight",
256
+ "model.text_model.encoder.layers.31.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.31.ln_2.bias",
257
+ "model.text_model.encoder.layers.31.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.31.ln_2.weight",
258
+ "model.text_model.encoder.layers.31.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_fc.bias",
259
+ "model.text_model.encoder.layers.31.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_fc.weight",
260
+ "model.text_model.encoder.layers.31.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_proj.bias",
261
+ "model.text_model.encoder.layers.31.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_proj.weight",
262
+ "model.text_model.encoder.layers.31.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.31.attn.out_proj.bias",
263
+ "model.text_model.encoder.layers.31.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.31.attn.out_proj.weight",
264
+ "model.text_model.encoder.layers.4.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.4.ln_1.bias",
265
+ "model.text_model.encoder.layers.4.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.4.ln_1.weight",
266
+ "model.text_model.encoder.layers.4.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.4.ln_2.bias",
267
+ "model.text_model.encoder.layers.4.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.4.ln_2.weight",
268
+ "model.text_model.encoder.layers.4.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_fc.bias",
269
+ "model.text_model.encoder.layers.4.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_fc.weight",
270
+ "model.text_model.encoder.layers.4.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_proj.bias",
271
+ "model.text_model.encoder.layers.4.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_proj.weight",
272
+ "model.text_model.encoder.layers.4.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.4.attn.out_proj.bias",
273
+ "model.text_model.encoder.layers.4.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.4.attn.out_proj.weight",
274
+ "model.text_model.encoder.layers.5.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.5.ln_1.bias",
275
+ "model.text_model.encoder.layers.5.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.5.ln_1.weight",
276
+ "model.text_model.encoder.layers.5.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.5.ln_2.bias",
277
+ "model.text_model.encoder.layers.5.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.5.ln_2.weight",
278
+ "model.text_model.encoder.layers.5.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_fc.bias",
279
+ "model.text_model.encoder.layers.5.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_fc.weight",
280
+ "model.text_model.encoder.layers.5.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_proj.bias",
281
+ "model.text_model.encoder.layers.5.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_proj.weight",
282
+ "model.text_model.encoder.layers.5.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.5.attn.out_proj.bias",
283
+ "model.text_model.encoder.layers.5.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.5.attn.out_proj.weight",
284
+ "model.text_model.encoder.layers.6.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.6.ln_1.bias",
285
+ "model.text_model.encoder.layers.6.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.6.ln_1.weight",
286
+ "model.text_model.encoder.layers.6.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.6.ln_2.bias",
287
+ "model.text_model.encoder.layers.6.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.6.ln_2.weight",
288
+ "model.text_model.encoder.layers.6.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_fc.bias",
289
+ "model.text_model.encoder.layers.6.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_fc.weight",
290
+ "model.text_model.encoder.layers.6.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_proj.bias",
291
+ "model.text_model.encoder.layers.6.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_proj.weight",
292
+ "model.text_model.encoder.layers.6.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.6.attn.out_proj.bias",
293
+ "model.text_model.encoder.layers.6.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.6.attn.out_proj.weight",
294
+ "model.text_model.encoder.layers.7.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.7.ln_1.bias",
295
+ "model.text_model.encoder.layers.7.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.7.ln_1.weight",
296
+ "model.text_model.encoder.layers.7.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.7.ln_2.bias",
297
+ "model.text_model.encoder.layers.7.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.7.ln_2.weight",
298
+ "model.text_model.encoder.layers.7.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_fc.bias",
299
+ "model.text_model.encoder.layers.7.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_fc.weight",
300
+ "model.text_model.encoder.layers.7.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_proj.bias",
301
+ "model.text_model.encoder.layers.7.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_proj.weight",
302
+ "model.text_model.encoder.layers.7.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.7.attn.out_proj.bias",
303
+ "model.text_model.encoder.layers.7.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.7.attn.out_proj.weight",
304
+ "model.text_model.encoder.layers.8.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.8.ln_1.bias",
305
+ "model.text_model.encoder.layers.8.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.8.ln_1.weight",
306
+ "model.text_model.encoder.layers.8.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.8.ln_2.bias",
307
+ "model.text_model.encoder.layers.8.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.8.ln_2.weight",
308
+ "model.text_model.encoder.layers.8.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_fc.bias",
309
+ "model.text_model.encoder.layers.8.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_fc.weight",
310
+ "model.text_model.encoder.layers.8.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_proj.bias",
311
+ "model.text_model.encoder.layers.8.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_proj.weight",
312
+ "model.text_model.encoder.layers.8.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.8.attn.out_proj.bias",
313
+ "model.text_model.encoder.layers.8.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.8.attn.out_proj.weight",
314
+ "model.text_model.encoder.layers.9.layer_norm1.bias": "conditioner.embedders.1.model.transformer.resblocks.9.ln_1.bias",
315
+ "model.text_model.encoder.layers.9.layer_norm1.weight": "conditioner.embedders.1.model.transformer.resblocks.9.ln_1.weight",
316
+ "model.text_model.encoder.layers.9.layer_norm2.bias": "conditioner.embedders.1.model.transformer.resblocks.9.ln_2.bias",
317
+ "model.text_model.encoder.layers.9.layer_norm2.weight": "conditioner.embedders.1.model.transformer.resblocks.9.ln_2.weight",
318
+ "model.text_model.encoder.layers.9.mlp.fc1.bias": "conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_fc.bias",
319
+ "model.text_model.encoder.layers.9.mlp.fc1.weight": "conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_fc.weight",
320
+ "model.text_model.encoder.layers.9.mlp.fc2.bias": "conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_proj.bias",
321
+ "model.text_model.encoder.layers.9.mlp.fc2.weight": "conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_proj.weight",
322
+ "model.text_model.encoder.layers.9.self_attn.out_proj.bias": "conditioner.embedders.1.model.transformer.resblocks.9.attn.out_proj.bias",
323
+ "model.text_model.encoder.layers.9.self_attn.out_proj.weight": "conditioner.embedders.1.model.transformer.resblocks.9.attn.out_proj.weight",
324
+ "model.text_model.final_layer_norm.bias": "conditioner.embedders.1.model.ln_final.bias",
325
+ "model.text_model.final_layer_norm.weight": "conditioner.embedders.1.model.ln_final.weight",
326
+ }
327
+
328
+ def SDXLTextEncoder2StateDictConverter_Original2Diffusers(state_dict):
329
+ state_dict_ = {name: state_dict[rename_dict[name]] for name in rename_dict if rename_dict[name] in state_dict}
330
+ for i in range(32):
331
+ name = f"conditioner.embedders.1.model.transformer.resblocks.{i}.attn.in_proj_weight"
332
+ if name not in state_dict:
333
+ continue
334
+ state_dict_[f"model.text_model.encoder.layers.{i}.self_attn.q_proj.weight"] = state_dict[name][:1280]
335
+ state_dict_[f"model.text_model.encoder.layers.{i}.self_attn.k_proj.weight"] = state_dict[name][1280:1280*2]
336
+ state_dict_[f"model.text_model.encoder.layers.{i}.self_attn.v_proj.weight"] = state_dict[name][1280*2:]
337
+ for i in range(32):
338
+ name = f"conditioner.embedders.1.model.transformer.resblocks.{i}.attn.in_proj_bias"
339
+ if name not in state_dict:
340
+ continue
341
+ state_dict_[f"model.text_model.encoder.layers.{i}.self_attn.q_proj.bias"] = state_dict[name][:1280]
342
+ state_dict_[f"model.text_model.encoder.layers.{i}.self_attn.k_proj.bias"] = state_dict[name][1280:1280*2]
343
+ state_dict_[f"model.text_model.encoder.layers.{i}.self_attn.v_proj.bias"] = state_dict[name][1280*2:]
344
+ state_dict_["model.text_projection.weight"] = state_dict["conditioner.embedders.1.model.text_projection"].T
345
+ return state_dict_
diffsynth/utils/state_dict_converters/sdxl_vae.py ADDED
@@ -0,0 +1,265 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ rename_dict = {
2
+ "decoder.conv_in.bias": "first_stage_model.decoder.conv_in.bias",
3
+ "decoder.conv_in.weight": "first_stage_model.decoder.conv_in.weight",
4
+ "decoder.conv_norm_out.bias": "first_stage_model.decoder.norm_out.bias",
5
+ "decoder.conv_norm_out.weight": "first_stage_model.decoder.norm_out.weight",
6
+ "decoder.conv_out.bias": "first_stage_model.decoder.conv_out.bias",
7
+ "decoder.conv_out.weight": "first_stage_model.decoder.conv_out.weight",
8
+ "decoder.mid_block.attentions.0.group_norm.bias": "first_stage_model.decoder.mid.attn_1.norm.bias",
9
+ "decoder.mid_block.attentions.0.group_norm.weight": "first_stage_model.decoder.mid.attn_1.norm.weight",
10
+ "decoder.mid_block.attentions.0.to_k.bias": "first_stage_model.decoder.mid.attn_1.k.bias",
11
+ "decoder.mid_block.attentions.0.to_k.weight": "first_stage_model.decoder.mid.attn_1.k.weight",
12
+ "decoder.mid_block.attentions.0.to_out.0.bias": "first_stage_model.decoder.mid.attn_1.proj_out.bias",
13
+ "decoder.mid_block.attentions.0.to_out.0.weight": "first_stage_model.decoder.mid.attn_1.proj_out.weight",
14
+ "decoder.mid_block.attentions.0.to_q.bias": "first_stage_model.decoder.mid.attn_1.q.bias",
15
+ "decoder.mid_block.attentions.0.to_q.weight": "first_stage_model.decoder.mid.attn_1.q.weight",
16
+ "decoder.mid_block.attentions.0.to_v.bias": "first_stage_model.decoder.mid.attn_1.v.bias",
17
+ "decoder.mid_block.attentions.0.to_v.weight": "first_stage_model.decoder.mid.attn_1.v.weight",
18
+ "decoder.mid_block.resnets.0.conv1.bias": "first_stage_model.decoder.mid.block_1.conv1.bias",
19
+ "decoder.mid_block.resnets.0.conv1.weight": "first_stage_model.decoder.mid.block_1.conv1.weight",
20
+ "decoder.mid_block.resnets.0.conv2.bias": "first_stage_model.decoder.mid.block_1.conv2.bias",
21
+ "decoder.mid_block.resnets.0.conv2.weight": "first_stage_model.decoder.mid.block_1.conv2.weight",
22
+ "decoder.mid_block.resnets.0.norm1.bias": "first_stage_model.decoder.mid.block_1.norm1.bias",
23
+ "decoder.mid_block.resnets.0.norm1.weight": "first_stage_model.decoder.mid.block_1.norm1.weight",
24
+ "decoder.mid_block.resnets.0.norm2.bias": "first_stage_model.decoder.mid.block_1.norm2.bias",
25
+ "decoder.mid_block.resnets.0.norm2.weight": "first_stage_model.decoder.mid.block_1.norm2.weight",
26
+ "decoder.mid_block.resnets.1.conv1.bias": "first_stage_model.decoder.mid.block_2.conv1.bias",
27
+ "decoder.mid_block.resnets.1.conv1.weight": "first_stage_model.decoder.mid.block_2.conv1.weight",
28
+ "decoder.mid_block.resnets.1.conv2.bias": "first_stage_model.decoder.mid.block_2.conv2.bias",
29
+ "decoder.mid_block.resnets.1.conv2.weight": "first_stage_model.decoder.mid.block_2.conv2.weight",
30
+ "decoder.mid_block.resnets.1.norm1.bias": "first_stage_model.decoder.mid.block_2.norm1.bias",
31
+ "decoder.mid_block.resnets.1.norm1.weight": "first_stage_model.decoder.mid.block_2.norm1.weight",
32
+ "decoder.mid_block.resnets.1.norm2.bias": "first_stage_model.decoder.mid.block_2.norm2.bias",
33
+ "decoder.mid_block.resnets.1.norm2.weight": "first_stage_model.decoder.mid.block_2.norm2.weight",
34
+ "decoder.up_blocks.0.resnets.0.conv1.bias": "first_stage_model.decoder.up.3.block.0.conv1.bias",
35
+ "decoder.up_blocks.0.resnets.0.conv1.weight": "first_stage_model.decoder.up.3.block.0.conv1.weight",
36
+ "decoder.up_blocks.0.resnets.0.conv2.bias": "first_stage_model.decoder.up.3.block.0.conv2.bias",
37
+ "decoder.up_blocks.0.resnets.0.conv2.weight": "first_stage_model.decoder.up.3.block.0.conv2.weight",
38
+ "decoder.up_blocks.0.resnets.0.norm1.bias": "first_stage_model.decoder.up.3.block.0.norm1.bias",
39
+ "decoder.up_blocks.0.resnets.0.norm1.weight": "first_stage_model.decoder.up.3.block.0.norm1.weight",
40
+ "decoder.up_blocks.0.resnets.0.norm2.bias": "first_stage_model.decoder.up.3.block.0.norm2.bias",
41
+ "decoder.up_blocks.0.resnets.0.norm2.weight": "first_stage_model.decoder.up.3.block.0.norm2.weight",
42
+ "decoder.up_blocks.0.resnets.1.conv1.bias": "first_stage_model.decoder.up.3.block.1.conv1.bias",
43
+ "decoder.up_blocks.0.resnets.1.conv1.weight": "first_stage_model.decoder.up.3.block.1.conv1.weight",
44
+ "decoder.up_blocks.0.resnets.1.conv2.bias": "first_stage_model.decoder.up.3.block.1.conv2.bias",
45
+ "decoder.up_blocks.0.resnets.1.conv2.weight": "first_stage_model.decoder.up.3.block.1.conv2.weight",
46
+ "decoder.up_blocks.0.resnets.1.norm1.bias": "first_stage_model.decoder.up.3.block.1.norm1.bias",
47
+ "decoder.up_blocks.0.resnets.1.norm1.weight": "first_stage_model.decoder.up.3.block.1.norm1.weight",
48
+ "decoder.up_blocks.0.resnets.1.norm2.bias": "first_stage_model.decoder.up.3.block.1.norm2.bias",
49
+ "decoder.up_blocks.0.resnets.1.norm2.weight": "first_stage_model.decoder.up.3.block.1.norm2.weight",
50
+ "decoder.up_blocks.0.resnets.2.conv1.bias": "first_stage_model.decoder.up.3.block.2.conv1.bias",
51
+ "decoder.up_blocks.0.resnets.2.conv1.weight": "first_stage_model.decoder.up.3.block.2.conv1.weight",
52
+ "decoder.up_blocks.0.resnets.2.conv2.bias": "first_stage_model.decoder.up.3.block.2.conv2.bias",
53
+ "decoder.up_blocks.0.resnets.2.conv2.weight": "first_stage_model.decoder.up.3.block.2.conv2.weight",
54
+ "decoder.up_blocks.0.resnets.2.norm1.bias": "first_stage_model.decoder.up.3.block.2.norm1.bias",
55
+ "decoder.up_blocks.0.resnets.2.norm1.weight": "first_stage_model.decoder.up.3.block.2.norm1.weight",
56
+ "decoder.up_blocks.0.resnets.2.norm2.bias": "first_stage_model.decoder.up.3.block.2.norm2.bias",
57
+ "decoder.up_blocks.0.resnets.2.norm2.weight": "first_stage_model.decoder.up.3.block.2.norm2.weight",
58
+ "decoder.up_blocks.0.upsamplers.0.conv.bias": "first_stage_model.decoder.up.3.upsample.conv.bias",
59
+ "decoder.up_blocks.0.upsamplers.0.conv.weight": "first_stage_model.decoder.up.3.upsample.conv.weight",
60
+ "decoder.up_blocks.1.resnets.0.conv1.bias": "first_stage_model.decoder.up.2.block.0.conv1.bias",
61
+ "decoder.up_blocks.1.resnets.0.conv1.weight": "first_stage_model.decoder.up.2.block.0.conv1.weight",
62
+ "decoder.up_blocks.1.resnets.0.conv2.bias": "first_stage_model.decoder.up.2.block.0.conv2.bias",
63
+ "decoder.up_blocks.1.resnets.0.conv2.weight": "first_stage_model.decoder.up.2.block.0.conv2.weight",
64
+ "decoder.up_blocks.1.resnets.0.norm1.bias": "first_stage_model.decoder.up.2.block.0.norm1.bias",
65
+ "decoder.up_blocks.1.resnets.0.norm1.weight": "first_stage_model.decoder.up.2.block.0.norm1.weight",
66
+ "decoder.up_blocks.1.resnets.0.norm2.bias": "first_stage_model.decoder.up.2.block.0.norm2.bias",
67
+ "decoder.up_blocks.1.resnets.0.norm2.weight": "first_stage_model.decoder.up.2.block.0.norm2.weight",
68
+ "decoder.up_blocks.1.resnets.1.conv1.bias": "first_stage_model.decoder.up.2.block.1.conv1.bias",
69
+ "decoder.up_blocks.1.resnets.1.conv1.weight": "first_stage_model.decoder.up.2.block.1.conv1.weight",
70
+ "decoder.up_blocks.1.resnets.1.conv2.bias": "first_stage_model.decoder.up.2.block.1.conv2.bias",
71
+ "decoder.up_blocks.1.resnets.1.conv2.weight": "first_stage_model.decoder.up.2.block.1.conv2.weight",
72
+ "decoder.up_blocks.1.resnets.1.norm1.bias": "first_stage_model.decoder.up.2.block.1.norm1.bias",
73
+ "decoder.up_blocks.1.resnets.1.norm1.weight": "first_stage_model.decoder.up.2.block.1.norm1.weight",
74
+ "decoder.up_blocks.1.resnets.1.norm2.bias": "first_stage_model.decoder.up.2.block.1.norm2.bias",
75
+ "decoder.up_blocks.1.resnets.1.norm2.weight": "first_stage_model.decoder.up.2.block.1.norm2.weight",
76
+ "decoder.up_blocks.1.resnets.2.conv1.bias": "first_stage_model.decoder.up.2.block.2.conv1.bias",
77
+ "decoder.up_blocks.1.resnets.2.conv1.weight": "first_stage_model.decoder.up.2.block.2.conv1.weight",
78
+ "decoder.up_blocks.1.resnets.2.conv2.bias": "first_stage_model.decoder.up.2.block.2.conv2.bias",
79
+ "decoder.up_blocks.1.resnets.2.conv2.weight": "first_stage_model.decoder.up.2.block.2.conv2.weight",
80
+ "decoder.up_blocks.1.resnets.2.norm1.bias": "first_stage_model.decoder.up.2.block.2.norm1.bias",
81
+ "decoder.up_blocks.1.resnets.2.norm1.weight": "first_stage_model.decoder.up.2.block.2.norm1.weight",
82
+ "decoder.up_blocks.1.resnets.2.norm2.bias": "first_stage_model.decoder.up.2.block.2.norm2.bias",
83
+ "decoder.up_blocks.1.resnets.2.norm2.weight": "first_stage_model.decoder.up.2.block.2.norm2.weight",
84
+ "decoder.up_blocks.1.upsamplers.0.conv.bias": "first_stage_model.decoder.up.2.upsample.conv.bias",
85
+ "decoder.up_blocks.1.upsamplers.0.conv.weight": "first_stage_model.decoder.up.2.upsample.conv.weight",
86
+ "decoder.up_blocks.2.resnets.0.conv1.bias": "first_stage_model.decoder.up.1.block.0.conv1.bias",
87
+ "decoder.up_blocks.2.resnets.0.conv1.weight": "first_stage_model.decoder.up.1.block.0.conv1.weight",
88
+ "decoder.up_blocks.2.resnets.0.conv2.bias": "first_stage_model.decoder.up.1.block.0.conv2.bias",
89
+ "decoder.up_blocks.2.resnets.0.conv2.weight": "first_stage_model.decoder.up.1.block.0.conv2.weight",
90
+ "decoder.up_blocks.2.resnets.0.conv_shortcut.bias": "first_stage_model.decoder.up.1.block.0.nin_shortcut.bias",
91
+ "decoder.up_blocks.2.resnets.0.conv_shortcut.weight": "first_stage_model.decoder.up.1.block.0.nin_shortcut.weight",
92
+ "decoder.up_blocks.2.resnets.0.norm1.bias": "first_stage_model.decoder.up.1.block.0.norm1.bias",
93
+ "decoder.up_blocks.2.resnets.0.norm1.weight": "first_stage_model.decoder.up.1.block.0.norm1.weight",
94
+ "decoder.up_blocks.2.resnets.0.norm2.bias": "first_stage_model.decoder.up.1.block.0.norm2.bias",
95
+ "decoder.up_blocks.2.resnets.0.norm2.weight": "first_stage_model.decoder.up.1.block.0.norm2.weight",
96
+ "decoder.up_blocks.2.resnets.1.conv1.bias": "first_stage_model.decoder.up.1.block.1.conv1.bias",
97
+ "decoder.up_blocks.2.resnets.1.conv1.weight": "first_stage_model.decoder.up.1.block.1.conv1.weight",
98
+ "decoder.up_blocks.2.resnets.1.conv2.bias": "first_stage_model.decoder.up.1.block.1.conv2.bias",
99
+ "decoder.up_blocks.2.resnets.1.conv2.weight": "first_stage_model.decoder.up.1.block.1.conv2.weight",
100
+ "decoder.up_blocks.2.resnets.1.norm1.bias": "first_stage_model.decoder.up.1.block.1.norm1.bias",
101
+ "decoder.up_blocks.2.resnets.1.norm1.weight": "first_stage_model.decoder.up.1.block.1.norm1.weight",
102
+ "decoder.up_blocks.2.resnets.1.norm2.bias": "first_stage_model.decoder.up.1.block.1.norm2.bias",
103
+ "decoder.up_blocks.2.resnets.1.norm2.weight": "first_stage_model.decoder.up.1.block.1.norm2.weight",
104
+ "decoder.up_blocks.2.resnets.2.conv1.bias": "first_stage_model.decoder.up.1.block.2.conv1.bias",
105
+ "decoder.up_blocks.2.resnets.2.conv1.weight": "first_stage_model.decoder.up.1.block.2.conv1.weight",
106
+ "decoder.up_blocks.2.resnets.2.conv2.bias": "first_stage_model.decoder.up.1.block.2.conv2.bias",
107
+ "decoder.up_blocks.2.resnets.2.conv2.weight": "first_stage_model.decoder.up.1.block.2.conv2.weight",
108
+ "decoder.up_blocks.2.resnets.2.norm1.bias": "first_stage_model.decoder.up.1.block.2.norm1.bias",
109
+ "decoder.up_blocks.2.resnets.2.norm1.weight": "first_stage_model.decoder.up.1.block.2.norm1.weight",
110
+ "decoder.up_blocks.2.resnets.2.norm2.bias": "first_stage_model.decoder.up.1.block.2.norm2.bias",
111
+ "decoder.up_blocks.2.resnets.2.norm2.weight": "first_stage_model.decoder.up.1.block.2.norm2.weight",
112
+ "decoder.up_blocks.2.upsamplers.0.conv.bias": "first_stage_model.decoder.up.1.upsample.conv.bias",
113
+ "decoder.up_blocks.2.upsamplers.0.conv.weight": "first_stage_model.decoder.up.1.upsample.conv.weight",
114
+ "decoder.up_blocks.3.resnets.0.conv1.bias": "first_stage_model.decoder.up.0.block.0.conv1.bias",
115
+ "decoder.up_blocks.3.resnets.0.conv1.weight": "first_stage_model.decoder.up.0.block.0.conv1.weight",
116
+ "decoder.up_blocks.3.resnets.0.conv2.bias": "first_stage_model.decoder.up.0.block.0.conv2.bias",
117
+ "decoder.up_blocks.3.resnets.0.conv2.weight": "first_stage_model.decoder.up.0.block.0.conv2.weight",
118
+ "decoder.up_blocks.3.resnets.0.conv_shortcut.bias": "first_stage_model.decoder.up.0.block.0.nin_shortcut.bias",
119
+ "decoder.up_blocks.3.resnets.0.conv_shortcut.weight": "first_stage_model.decoder.up.0.block.0.nin_shortcut.weight",
120
+ "decoder.up_blocks.3.resnets.0.norm1.bias": "first_stage_model.decoder.up.0.block.0.norm1.bias",
121
+ "decoder.up_blocks.3.resnets.0.norm1.weight": "first_stage_model.decoder.up.0.block.0.norm1.weight",
122
+ "decoder.up_blocks.3.resnets.0.norm2.bias": "first_stage_model.decoder.up.0.block.0.norm2.bias",
123
+ "decoder.up_blocks.3.resnets.0.norm2.weight": "first_stage_model.decoder.up.0.block.0.norm2.weight",
124
+ "decoder.up_blocks.3.resnets.1.conv1.bias": "first_stage_model.decoder.up.0.block.1.conv1.bias",
125
+ "decoder.up_blocks.3.resnets.1.conv1.weight": "first_stage_model.decoder.up.0.block.1.conv1.weight",
126
+ "decoder.up_blocks.3.resnets.1.conv2.bias": "first_stage_model.decoder.up.0.block.1.conv2.bias",
127
+ "decoder.up_blocks.3.resnets.1.conv2.weight": "first_stage_model.decoder.up.0.block.1.conv2.weight",
128
+ "decoder.up_blocks.3.resnets.1.norm1.bias": "first_stage_model.decoder.up.0.block.1.norm1.bias",
129
+ "decoder.up_blocks.3.resnets.1.norm1.weight": "first_stage_model.decoder.up.0.block.1.norm1.weight",
130
+ "decoder.up_blocks.3.resnets.1.norm2.bias": "first_stage_model.decoder.up.0.block.1.norm2.bias",
131
+ "decoder.up_blocks.3.resnets.1.norm2.weight": "first_stage_model.decoder.up.0.block.1.norm2.weight",
132
+ "decoder.up_blocks.3.resnets.2.conv1.bias": "first_stage_model.decoder.up.0.block.2.conv1.bias",
133
+ "decoder.up_blocks.3.resnets.2.conv1.weight": "first_stage_model.decoder.up.0.block.2.conv1.weight",
134
+ "decoder.up_blocks.3.resnets.2.conv2.bias": "first_stage_model.decoder.up.0.block.2.conv2.bias",
135
+ "decoder.up_blocks.3.resnets.2.conv2.weight": "first_stage_model.decoder.up.0.block.2.conv2.weight",
136
+ "decoder.up_blocks.3.resnets.2.norm1.bias": "first_stage_model.decoder.up.0.block.2.norm1.bias",
137
+ "decoder.up_blocks.3.resnets.2.norm1.weight": "first_stage_model.decoder.up.0.block.2.norm1.weight",
138
+ "decoder.up_blocks.3.resnets.2.norm2.bias": "first_stage_model.decoder.up.0.block.2.norm2.bias",
139
+ "decoder.up_blocks.3.resnets.2.norm2.weight": "first_stage_model.decoder.up.0.block.2.norm2.weight",
140
+ "encoder.conv_in.bias": "first_stage_model.encoder.conv_in.bias",
141
+ "encoder.conv_in.weight": "first_stage_model.encoder.conv_in.weight",
142
+ "encoder.conv_norm_out.bias": "first_stage_model.encoder.norm_out.bias",
143
+ "encoder.conv_norm_out.weight": "first_stage_model.encoder.norm_out.weight",
144
+ "encoder.conv_out.bias": "first_stage_model.encoder.conv_out.bias",
145
+ "encoder.conv_out.weight": "first_stage_model.encoder.conv_out.weight",
146
+ "encoder.down_blocks.0.downsamplers.0.conv.bias": "first_stage_model.encoder.down.0.downsample.conv.bias",
147
+ "encoder.down_blocks.0.downsamplers.0.conv.weight": "first_stage_model.encoder.down.0.downsample.conv.weight",
148
+ "encoder.down_blocks.0.resnets.0.conv1.bias": "first_stage_model.encoder.down.0.block.0.conv1.bias",
149
+ "encoder.down_blocks.0.resnets.0.conv1.weight": "first_stage_model.encoder.down.0.block.0.conv1.weight",
150
+ "encoder.down_blocks.0.resnets.0.conv2.bias": "first_stage_model.encoder.down.0.block.0.conv2.bias",
151
+ "encoder.down_blocks.0.resnets.0.conv2.weight": "first_stage_model.encoder.down.0.block.0.conv2.weight",
152
+ "encoder.down_blocks.0.resnets.0.norm1.bias": "first_stage_model.encoder.down.0.block.0.norm1.bias",
153
+ "encoder.down_blocks.0.resnets.0.norm1.weight": "first_stage_model.encoder.down.0.block.0.norm1.weight",
154
+ "encoder.down_blocks.0.resnets.0.norm2.bias": "first_stage_model.encoder.down.0.block.0.norm2.bias",
155
+ "encoder.down_blocks.0.resnets.0.norm2.weight": "first_stage_model.encoder.down.0.block.0.norm2.weight",
156
+ "encoder.down_blocks.0.resnets.1.conv1.bias": "first_stage_model.encoder.down.0.block.1.conv1.bias",
157
+ "encoder.down_blocks.0.resnets.1.conv1.weight": "first_stage_model.encoder.down.0.block.1.conv1.weight",
158
+ "encoder.down_blocks.0.resnets.1.conv2.bias": "first_stage_model.encoder.down.0.block.1.conv2.bias",
159
+ "encoder.down_blocks.0.resnets.1.conv2.weight": "first_stage_model.encoder.down.0.block.1.conv2.weight",
160
+ "encoder.down_blocks.0.resnets.1.norm1.bias": "first_stage_model.encoder.down.0.block.1.norm1.bias",
161
+ "encoder.down_blocks.0.resnets.1.norm1.weight": "first_stage_model.encoder.down.0.block.1.norm1.weight",
162
+ "encoder.down_blocks.0.resnets.1.norm2.bias": "first_stage_model.encoder.down.0.block.1.norm2.bias",
163
+ "encoder.down_blocks.0.resnets.1.norm2.weight": "first_stage_model.encoder.down.0.block.1.norm2.weight",
164
+ "encoder.down_blocks.1.downsamplers.0.conv.bias": "first_stage_model.encoder.down.1.downsample.conv.bias",
165
+ "encoder.down_blocks.1.downsamplers.0.conv.weight": "first_stage_model.encoder.down.1.downsample.conv.weight",
166
+ "encoder.down_blocks.1.resnets.0.conv1.bias": "first_stage_model.encoder.down.1.block.0.conv1.bias",
167
+ "encoder.down_blocks.1.resnets.0.conv1.weight": "first_stage_model.encoder.down.1.block.0.conv1.weight",
168
+ "encoder.down_blocks.1.resnets.0.conv2.bias": "first_stage_model.encoder.down.1.block.0.conv2.bias",
169
+ "encoder.down_blocks.1.resnets.0.conv2.weight": "first_stage_model.encoder.down.1.block.0.conv2.weight",
170
+ "encoder.down_blocks.1.resnets.0.conv_shortcut.bias": "first_stage_model.encoder.down.1.block.0.nin_shortcut.bias",
171
+ "encoder.down_blocks.1.resnets.0.conv_shortcut.weight": "first_stage_model.encoder.down.1.block.0.nin_shortcut.weight",
172
+ "encoder.down_blocks.1.resnets.0.norm1.bias": "first_stage_model.encoder.down.1.block.0.norm1.bias",
173
+ "encoder.down_blocks.1.resnets.0.norm1.weight": "first_stage_model.encoder.down.1.block.0.norm1.weight",
174
+ "encoder.down_blocks.1.resnets.0.norm2.bias": "first_stage_model.encoder.down.1.block.0.norm2.bias",
175
+ "encoder.down_blocks.1.resnets.0.norm2.weight": "first_stage_model.encoder.down.1.block.0.norm2.weight",
176
+ "encoder.down_blocks.1.resnets.1.conv1.bias": "first_stage_model.encoder.down.1.block.1.conv1.bias",
177
+ "encoder.down_blocks.1.resnets.1.conv1.weight": "first_stage_model.encoder.down.1.block.1.conv1.weight",
178
+ "encoder.down_blocks.1.resnets.1.conv2.bias": "first_stage_model.encoder.down.1.block.1.conv2.bias",
179
+ "encoder.down_blocks.1.resnets.1.conv2.weight": "first_stage_model.encoder.down.1.block.1.conv2.weight",
180
+ "encoder.down_blocks.1.resnets.1.norm1.bias": "first_stage_model.encoder.down.1.block.1.norm1.bias",
181
+ "encoder.down_blocks.1.resnets.1.norm1.weight": "first_stage_model.encoder.down.1.block.1.norm1.weight",
182
+ "encoder.down_blocks.1.resnets.1.norm2.bias": "first_stage_model.encoder.down.1.block.1.norm2.bias",
183
+ "encoder.down_blocks.1.resnets.1.norm2.weight": "first_stage_model.encoder.down.1.block.1.norm2.weight",
184
+ "encoder.down_blocks.2.downsamplers.0.conv.bias": "first_stage_model.encoder.down.2.downsample.conv.bias",
185
+ "encoder.down_blocks.2.downsamplers.0.conv.weight": "first_stage_model.encoder.down.2.downsample.conv.weight",
186
+ "encoder.down_blocks.2.resnets.0.conv1.bias": "first_stage_model.encoder.down.2.block.0.conv1.bias",
187
+ "encoder.down_blocks.2.resnets.0.conv1.weight": "first_stage_model.encoder.down.2.block.0.conv1.weight",
188
+ "encoder.down_blocks.2.resnets.0.conv2.bias": "first_stage_model.encoder.down.2.block.0.conv2.bias",
189
+ "encoder.down_blocks.2.resnets.0.conv2.weight": "first_stage_model.encoder.down.2.block.0.conv2.weight",
190
+ "encoder.down_blocks.2.resnets.0.conv_shortcut.bias": "first_stage_model.encoder.down.2.block.0.nin_shortcut.bias",
191
+ "encoder.down_blocks.2.resnets.0.conv_shortcut.weight": "first_stage_model.encoder.down.2.block.0.nin_shortcut.weight",
192
+ "encoder.down_blocks.2.resnets.0.norm1.bias": "first_stage_model.encoder.down.2.block.0.norm1.bias",
193
+ "encoder.down_blocks.2.resnets.0.norm1.weight": "first_stage_model.encoder.down.2.block.0.norm1.weight",
194
+ "encoder.down_blocks.2.resnets.0.norm2.bias": "first_stage_model.encoder.down.2.block.0.norm2.bias",
195
+ "encoder.down_blocks.2.resnets.0.norm2.weight": "first_stage_model.encoder.down.2.block.0.norm2.weight",
196
+ "encoder.down_blocks.2.resnets.1.conv1.bias": "first_stage_model.encoder.down.2.block.1.conv1.bias",
197
+ "encoder.down_blocks.2.resnets.1.conv1.weight": "first_stage_model.encoder.down.2.block.1.conv1.weight",
198
+ "encoder.down_blocks.2.resnets.1.conv2.bias": "first_stage_model.encoder.down.2.block.1.conv2.bias",
199
+ "encoder.down_blocks.2.resnets.1.conv2.weight": "first_stage_model.encoder.down.2.block.1.conv2.weight",
200
+ "encoder.down_blocks.2.resnets.1.norm1.bias": "first_stage_model.encoder.down.2.block.1.norm1.bias",
201
+ "encoder.down_blocks.2.resnets.1.norm1.weight": "first_stage_model.encoder.down.2.block.1.norm1.weight",
202
+ "encoder.down_blocks.2.resnets.1.norm2.bias": "first_stage_model.encoder.down.2.block.1.norm2.bias",
203
+ "encoder.down_blocks.2.resnets.1.norm2.weight": "first_stage_model.encoder.down.2.block.1.norm2.weight",
204
+ "encoder.down_blocks.3.resnets.0.conv1.bias": "first_stage_model.encoder.down.3.block.0.conv1.bias",
205
+ "encoder.down_blocks.3.resnets.0.conv1.weight": "first_stage_model.encoder.down.3.block.0.conv1.weight",
206
+ "encoder.down_blocks.3.resnets.0.conv2.bias": "first_stage_model.encoder.down.3.block.0.conv2.bias",
207
+ "encoder.down_blocks.3.resnets.0.conv2.weight": "first_stage_model.encoder.down.3.block.0.conv2.weight",
208
+ "encoder.down_blocks.3.resnets.0.norm1.bias": "first_stage_model.encoder.down.3.block.0.norm1.bias",
209
+ "encoder.down_blocks.3.resnets.0.norm1.weight": "first_stage_model.encoder.down.3.block.0.norm1.weight",
210
+ "encoder.down_blocks.3.resnets.0.norm2.bias": "first_stage_model.encoder.down.3.block.0.norm2.bias",
211
+ "encoder.down_blocks.3.resnets.0.norm2.weight": "first_stage_model.encoder.down.3.block.0.norm2.weight",
212
+ "encoder.down_blocks.3.resnets.1.conv1.bias": "first_stage_model.encoder.down.3.block.1.conv1.bias",
213
+ "encoder.down_blocks.3.resnets.1.conv1.weight": "first_stage_model.encoder.down.3.block.1.conv1.weight",
214
+ "encoder.down_blocks.3.resnets.1.conv2.bias": "first_stage_model.encoder.down.3.block.1.conv2.bias",
215
+ "encoder.down_blocks.3.resnets.1.conv2.weight": "first_stage_model.encoder.down.3.block.1.conv2.weight",
216
+ "encoder.down_blocks.3.resnets.1.norm1.bias": "first_stage_model.encoder.down.3.block.1.norm1.bias",
217
+ "encoder.down_blocks.3.resnets.1.norm1.weight": "first_stage_model.encoder.down.3.block.1.norm1.weight",
218
+ "encoder.down_blocks.3.resnets.1.norm2.bias": "first_stage_model.encoder.down.3.block.1.norm2.bias",
219
+ "encoder.down_blocks.3.resnets.1.norm2.weight": "first_stage_model.encoder.down.3.block.1.norm2.weight",
220
+ "encoder.mid_block.attentions.0.group_norm.bias": "first_stage_model.encoder.mid.attn_1.norm.bias",
221
+ "encoder.mid_block.attentions.0.group_norm.weight": "first_stage_model.encoder.mid.attn_1.norm.weight",
222
+ "encoder.mid_block.attentions.0.to_k.bias": "first_stage_model.encoder.mid.attn_1.k.bias",
223
+ "encoder.mid_block.attentions.0.to_k.weight": "first_stage_model.encoder.mid.attn_1.k.weight",
224
+ "encoder.mid_block.attentions.0.to_out.0.bias": "first_stage_model.encoder.mid.attn_1.proj_out.bias",
225
+ "encoder.mid_block.attentions.0.to_out.0.weight": "first_stage_model.encoder.mid.attn_1.proj_out.weight",
226
+ "encoder.mid_block.attentions.0.to_q.bias": "first_stage_model.encoder.mid.attn_1.q.bias",
227
+ "encoder.mid_block.attentions.0.to_q.weight": "first_stage_model.encoder.mid.attn_1.q.weight",
228
+ "encoder.mid_block.attentions.0.to_v.bias": "first_stage_model.encoder.mid.attn_1.v.bias",
229
+ "encoder.mid_block.attentions.0.to_v.weight": "first_stage_model.encoder.mid.attn_1.v.weight",
230
+ "encoder.mid_block.resnets.0.conv1.bias": "first_stage_model.encoder.mid.block_1.conv1.bias",
231
+ "encoder.mid_block.resnets.0.conv1.weight": "first_stage_model.encoder.mid.block_1.conv1.weight",
232
+ "encoder.mid_block.resnets.0.conv2.bias": "first_stage_model.encoder.mid.block_1.conv2.bias",
233
+ "encoder.mid_block.resnets.0.conv2.weight": "first_stage_model.encoder.mid.block_1.conv2.weight",
234
+ "encoder.mid_block.resnets.0.norm1.bias": "first_stage_model.encoder.mid.block_1.norm1.bias",
235
+ "encoder.mid_block.resnets.0.norm1.weight": "first_stage_model.encoder.mid.block_1.norm1.weight",
236
+ "encoder.mid_block.resnets.0.norm2.bias": "first_stage_model.encoder.mid.block_1.norm2.bias",
237
+ "encoder.mid_block.resnets.0.norm2.weight": "first_stage_model.encoder.mid.block_1.norm2.weight",
238
+ "encoder.mid_block.resnets.1.conv1.bias": "first_stage_model.encoder.mid.block_2.conv1.bias",
239
+ "encoder.mid_block.resnets.1.conv1.weight": "first_stage_model.encoder.mid.block_2.conv1.weight",
240
+ "encoder.mid_block.resnets.1.conv2.bias": "first_stage_model.encoder.mid.block_2.conv2.bias",
241
+ "encoder.mid_block.resnets.1.conv2.weight": "first_stage_model.encoder.mid.block_2.conv2.weight",
242
+ "encoder.mid_block.resnets.1.norm1.bias": "first_stage_model.encoder.mid.block_2.norm1.bias",
243
+ "encoder.mid_block.resnets.1.norm1.weight": "first_stage_model.encoder.mid.block_2.norm1.weight",
244
+ "encoder.mid_block.resnets.1.norm2.bias": "first_stage_model.encoder.mid.block_2.norm2.bias",
245
+ "encoder.mid_block.resnets.1.norm2.weight": "first_stage_model.encoder.mid.block_2.norm2.weight",
246
+ "post_quant_conv.bias": "first_stage_model.post_quant_conv.bias",
247
+ "post_quant_conv.weight": "first_stage_model.post_quant_conv.weight",
248
+ "quant_conv.bias": "first_stage_model.quant_conv.bias",
249
+ "quant_conv.weight": "first_stage_model.quant_conv.weight",
250
+ }
251
+
252
+ def SDXLVAEStateDictConverter_Original2Diffusers(state_dict):
253
+ state_dict_ = {name: state_dict[rename_dict[name]] for name in rename_dict if rename_dict[name] in state_dict}
254
+ for name in [
255
+ "encoder.mid_block.attentions.0.to_q.weight",
256
+ "encoder.mid_block.attentions.0.to_k.weight",
257
+ "encoder.mid_block.attentions.0.to_v.weight",
258
+ "encoder.mid_block.attentions.0.to_out.0.weight",
259
+ "decoder.mid_block.attentions.0.to_q.weight",
260
+ "decoder.mid_block.attentions.0.to_k.weight",
261
+ "decoder.mid_block.attentions.0.to_v.weight",
262
+ "decoder.mid_block.attentions.0.to_out.0.weight",
263
+ ]:
264
+ state_dict_[name] = state_dict_[name].squeeze()
265
+ return state_dict_
diffsynth/utils/state_dict_converters/stable_diffusion_text_encoder.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ def SDTextEncoderStateDictConverter(state_dict):
2
+ new_state_dict = {}
3
+ for key in state_dict:
4
+ if key.startswith("text_model.") and "position_ids" not in key:
5
+ new_key = "model." + key
6
+ new_state_dict[new_key] = state_dict[key]
7
+ return new_state_dict
diffsynth/utils/state_dict_converters/stable_diffusion_vae.py ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def SDVAEStateDictConverter(state_dict):
2
+ new_state_dict = {}
3
+ for key in state_dict:
4
+ if ".query." in key:
5
+ new_key = key.replace(".query.", ".to_q.")
6
+ new_state_dict[new_key] = state_dict[key]
7
+ elif ".key." in key:
8
+ new_key = key.replace(".key.", ".to_k.")
9
+ new_state_dict[new_key] = state_dict[key]
10
+ elif ".value." in key:
11
+ new_key = key.replace(".value.", ".to_v.")
12
+ new_state_dict[new_key] = state_dict[key]
13
+ elif ".proj_attn." in key:
14
+ new_key = key.replace(".proj_attn.", ".to_out.0.")
15
+ new_state_dict[new_key] = state_dict[key]
16
+ else:
17
+ new_state_dict[key] = state_dict[key]
18
+ return new_state_dict
diffsynth/utils/state_dict_converters/stable_diffusion_xl_text_encoder.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ def SDXLTextEncoder2StateDictConverter(state_dict):
4
+ new_state_dict = {}
5
+ for key in state_dict:
6
+ if key == "text_projection.weight":
7
+ val = state_dict[key]
8
+ new_state_dict["model.text_projection.weight"] = val.float() if val.dtype == torch.float16 else val
9
+ elif key.startswith("text_model.") and "position_ids" not in key:
10
+ new_key = "model." + key
11
+ val = state_dict[key]
12
+ new_state_dict[new_key] = val.float() if val.dtype == torch.float16 else val
13
+ return new_state_dict
diffsynth/utils/state_dict_converters/step1x_connector.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ def Qwen2ConnectorStateDictConverter(state_dict):
2
+ state_dict_ = {}
3
+ for name in state_dict:
4
+ if name.startswith("connector."):
5
+ name_ = name[len("connector."):]
6
+ state_dict_[name_] = state_dict[name]
7
+ return state_dict_
diffsynth/utils/state_dict_converters/wan_video_animate_adapter.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ def WanAnimateAdapterStateDictConverter(state_dict):
2
+ state_dict_ = {}
3
+ for name in state_dict:
4
+ if name.startswith("pose_patch_embedding.") or name.startswith("face_adapter") or name.startswith("face_encoder") or name.startswith("motion_encoder"):
5
+ state_dict_[name] = state_dict[name]
6
+ return state_dict_
diffsynth/utils/state_dict_converters/wan_video_dit.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def WanVideoDiTFromDiffusers(state_dict):
2
+ rename_dict = {
3
+ "blocks.0.attn1.norm_k.weight": "blocks.0.self_attn.norm_k.weight",
4
+ "blocks.0.attn1.norm_q.weight": "blocks.0.self_attn.norm_q.weight",
5
+ "blocks.0.attn1.to_k.bias": "blocks.0.self_attn.k.bias",
6
+ "blocks.0.attn1.to_k.weight": "blocks.0.self_attn.k.weight",
7
+ "blocks.0.attn1.to_out.0.bias": "blocks.0.self_attn.o.bias",
8
+ "blocks.0.attn1.to_out.0.weight": "blocks.0.self_attn.o.weight",
9
+ "blocks.0.attn1.to_q.bias": "blocks.0.self_attn.q.bias",
10
+ "blocks.0.attn1.to_q.weight": "blocks.0.self_attn.q.weight",
11
+ "blocks.0.attn1.to_v.bias": "blocks.0.self_attn.v.bias",
12
+ "blocks.0.attn1.to_v.weight": "blocks.0.self_attn.v.weight",
13
+ "blocks.0.attn2.norm_k.weight": "blocks.0.cross_attn.norm_k.weight",
14
+ "blocks.0.attn2.norm_q.weight": "blocks.0.cross_attn.norm_q.weight",
15
+ "blocks.0.attn2.to_k.bias": "blocks.0.cross_attn.k.bias",
16
+ "blocks.0.attn2.to_k.weight": "blocks.0.cross_attn.k.weight",
17
+ "blocks.0.attn2.to_out.0.bias": "blocks.0.cross_attn.o.bias",
18
+ "blocks.0.attn2.to_out.0.weight": "blocks.0.cross_attn.o.weight",
19
+ "blocks.0.attn2.to_q.bias": "blocks.0.cross_attn.q.bias",
20
+ "blocks.0.attn2.to_q.weight": "blocks.0.cross_attn.q.weight",
21
+ "blocks.0.attn2.to_v.bias": "blocks.0.cross_attn.v.bias",
22
+ "blocks.0.attn2.to_v.weight": "blocks.0.cross_attn.v.weight",
23
+ "blocks.0.attn2.add_k_proj.bias":"blocks.0.cross_attn.k_img.bias",
24
+ "blocks.0.attn2.add_k_proj.weight":"blocks.0.cross_attn.k_img.weight",
25
+ "blocks.0.attn2.add_v_proj.bias":"blocks.0.cross_attn.v_img.bias",
26
+ "blocks.0.attn2.add_v_proj.weight":"blocks.0.cross_attn.v_img.weight",
27
+ "blocks.0.attn2.norm_added_k.weight":"blocks.0.cross_attn.norm_k_img.weight",
28
+ "blocks.0.ffn.net.0.proj.bias": "blocks.0.ffn.0.bias",
29
+ "blocks.0.ffn.net.0.proj.weight": "blocks.0.ffn.0.weight",
30
+ "blocks.0.ffn.net.2.bias": "blocks.0.ffn.2.bias",
31
+ "blocks.0.ffn.net.2.weight": "blocks.0.ffn.2.weight",
32
+ "blocks.0.norm2.bias": "blocks.0.norm3.bias",
33
+ "blocks.0.norm2.weight": "blocks.0.norm3.weight",
34
+ "blocks.0.scale_shift_table": "blocks.0.modulation",
35
+ "condition_embedder.text_embedder.linear_1.bias": "text_embedding.0.bias",
36
+ "condition_embedder.text_embedder.linear_1.weight": "text_embedding.0.weight",
37
+ "condition_embedder.text_embedder.linear_2.bias": "text_embedding.2.bias",
38
+ "condition_embedder.text_embedder.linear_2.weight": "text_embedding.2.weight",
39
+ "condition_embedder.time_embedder.linear_1.bias": "time_embedding.0.bias",
40
+ "condition_embedder.time_embedder.linear_1.weight": "time_embedding.0.weight",
41
+ "condition_embedder.time_embedder.linear_2.bias": "time_embedding.2.bias",
42
+ "condition_embedder.time_embedder.linear_2.weight": "time_embedding.2.weight",
43
+ "condition_embedder.time_proj.bias": "time_projection.1.bias",
44
+ "condition_embedder.time_proj.weight": "time_projection.1.weight",
45
+ "condition_embedder.image_embedder.ff.net.0.proj.bias":"img_emb.proj.1.bias",
46
+ "condition_embedder.image_embedder.ff.net.0.proj.weight":"img_emb.proj.1.weight",
47
+ "condition_embedder.image_embedder.ff.net.2.bias":"img_emb.proj.3.bias",
48
+ "condition_embedder.image_embedder.ff.net.2.weight":"img_emb.proj.3.weight",
49
+ "condition_embedder.image_embedder.norm1.bias":"img_emb.proj.0.bias",
50
+ "condition_embedder.image_embedder.norm1.weight":"img_emb.proj.0.weight",
51
+ "condition_embedder.image_embedder.norm2.bias":"img_emb.proj.4.bias",
52
+ "condition_embedder.image_embedder.norm2.weight":"img_emb.proj.4.weight",
53
+ "patch_embedding.bias": "patch_embedding.bias",
54
+ "patch_embedding.weight": "patch_embedding.weight",
55
+ "scale_shift_table": "head.modulation",
56
+ "proj_out.bias": "head.head.bias",
57
+ "proj_out.weight": "head.head.weight",
58
+ }
59
+ state_dict_ = {}
60
+ for name in state_dict:
61
+ if name in rename_dict:
62
+ state_dict_[rename_dict[name]] = state_dict[name]
63
+ else:
64
+ name_ = ".".join(name.split(".")[:1] + ["0"] + name.split(".")[2:])
65
+ if name_ in rename_dict:
66
+ name_ = rename_dict[name_]
67
+ name_ = ".".join(name_.split(".")[:1] + [name.split(".")[1]] + name_.split(".")[2:])
68
+ state_dict_[name_] = state_dict[name]
69
+ return state_dict_
70
+
71
+
72
+ def WanVideoDiTStateDictConverter(state_dict):
73
+ state_dict_ = {}
74
+ for name in state_dict:
75
+ if name.startswith("vace"):
76
+ continue
77
+ if name.split(".")[0] in ["pose_patch_embedding", "face_adapter", "face_encoder", "motion_encoder"]:
78
+ continue
79
+ name_ = name
80
+ if name_.startswith("model."):
81
+ name_ = name_[len("model."):]
82
+ state_dict_[name_] = state_dict[name]
83
+ return state_dict_
diffsynth/utils/state_dict_converters/wan_video_image_encoder.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ def WanImageEncoderStateDictConverter(state_dict):
2
+ state_dict_ = {}
3
+ for name in state_dict:
4
+ if name.startswith("textual."):
5
+ continue
6
+ name_ = "model." + name
7
+ state_dict_[name_] = state_dict[name]
8
+ return state_dict_
diffsynth/utils/state_dict_converters/wan_video_mot.py ADDED
@@ -0,0 +1,78 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def WanVideoMotStateDictConverter(state_dict):
2
+ rename_dict = {
3
+ "blocks.0.attn1.norm_k.weight": "blocks.0.self_attn.norm_k.weight",
4
+ "blocks.0.attn1.norm_q.weight": "blocks.0.self_attn.norm_q.weight",
5
+ "blocks.0.attn1.to_k.bias": "blocks.0.self_attn.k.bias",
6
+ "blocks.0.attn1.to_k.weight": "blocks.0.self_attn.k.weight",
7
+ "blocks.0.attn1.to_out.0.bias": "blocks.0.self_attn.o.bias",
8
+ "blocks.0.attn1.to_out.0.weight": "blocks.0.self_attn.o.weight",
9
+ "blocks.0.attn1.to_q.bias": "blocks.0.self_attn.q.bias",
10
+ "blocks.0.attn1.to_q.weight": "blocks.0.self_attn.q.weight",
11
+ "blocks.0.attn1.to_v.bias": "blocks.0.self_attn.v.bias",
12
+ "blocks.0.attn1.to_v.weight": "blocks.0.self_attn.v.weight",
13
+ "blocks.0.attn2.norm_k.weight": "blocks.0.cross_attn.norm_k.weight",
14
+ "blocks.0.attn2.norm_q.weight": "blocks.0.cross_attn.norm_q.weight",
15
+ "blocks.0.attn2.to_k.bias": "blocks.0.cross_attn.k.bias",
16
+ "blocks.0.attn2.to_k.weight": "blocks.0.cross_attn.k.weight",
17
+ "blocks.0.attn2.to_out.0.bias": "blocks.0.cross_attn.o.bias",
18
+ "blocks.0.attn2.to_out.0.weight": "blocks.0.cross_attn.o.weight",
19
+ "blocks.0.attn2.to_q.bias": "blocks.0.cross_attn.q.bias",
20
+ "blocks.0.attn2.to_q.weight": "blocks.0.cross_attn.q.weight",
21
+ "blocks.0.attn2.to_v.bias": "blocks.0.cross_attn.v.bias",
22
+ "blocks.0.attn2.to_v.weight": "blocks.0.cross_attn.v.weight",
23
+ "blocks.0.attn2.add_k_proj.bias":"blocks.0.cross_attn.k_img.bias",
24
+ "blocks.0.attn2.add_k_proj.weight":"blocks.0.cross_attn.k_img.weight",
25
+ "blocks.0.attn2.add_v_proj.bias":"blocks.0.cross_attn.v_img.bias",
26
+ "blocks.0.attn2.add_v_proj.weight":"blocks.0.cross_attn.v_img.weight",
27
+ "blocks.0.attn2.norm_added_k.weight":"blocks.0.cross_attn.norm_k_img.weight",
28
+ "blocks.0.ffn.net.0.proj.bias": "blocks.0.ffn.0.bias",
29
+ "blocks.0.ffn.net.0.proj.weight": "blocks.0.ffn.0.weight",
30
+ "blocks.0.ffn.net.2.bias": "blocks.0.ffn.2.bias",
31
+ "blocks.0.ffn.net.2.weight": "blocks.0.ffn.2.weight",
32
+ "blocks.0.norm2.bias": "blocks.0.norm3.bias",
33
+ "blocks.0.norm2.weight": "blocks.0.norm3.weight",
34
+ "blocks.0.scale_shift_table": "blocks.0.modulation",
35
+ "condition_embedder.text_embedder.linear_1.bias": "text_embedding.0.bias",
36
+ "condition_embedder.text_embedder.linear_1.weight": "text_embedding.0.weight",
37
+ "condition_embedder.text_embedder.linear_2.bias": "text_embedding.2.bias",
38
+ "condition_embedder.text_embedder.linear_2.weight": "text_embedding.2.weight",
39
+ "condition_embedder.time_embedder.linear_1.bias": "time_embedding.0.bias",
40
+ "condition_embedder.time_embedder.linear_1.weight": "time_embedding.0.weight",
41
+ "condition_embedder.time_embedder.linear_2.bias": "time_embedding.2.bias",
42
+ "condition_embedder.time_embedder.linear_2.weight": "time_embedding.2.weight",
43
+ "condition_embedder.time_proj.bias": "time_projection.1.bias",
44
+ "condition_embedder.time_proj.weight": "time_projection.1.weight",
45
+ "condition_embedder.image_embedder.ff.net.0.proj.bias":"img_emb.proj.1.bias",
46
+ "condition_embedder.image_embedder.ff.net.0.proj.weight":"img_emb.proj.1.weight",
47
+ "condition_embedder.image_embedder.ff.net.2.bias":"img_emb.proj.3.bias",
48
+ "condition_embedder.image_embedder.ff.net.2.weight":"img_emb.proj.3.weight",
49
+ "condition_embedder.image_embedder.norm1.bias":"img_emb.proj.0.bias",
50
+ "condition_embedder.image_embedder.norm1.weight":"img_emb.proj.0.weight",
51
+ "condition_embedder.image_embedder.norm2.bias":"img_emb.proj.4.bias",
52
+ "condition_embedder.image_embedder.norm2.weight":"img_emb.proj.4.weight",
53
+ "patch_embedding.bias": "patch_embedding.bias",
54
+ "patch_embedding.weight": "patch_embedding.weight",
55
+ "scale_shift_table": "head.modulation",
56
+ "proj_out.bias": "head.head.bias",
57
+ "proj_out.weight": "head.head.weight",
58
+ }
59
+ mot_layers = (0, 4, 8, 12, 16, 20, 24, 28, 32, 36)
60
+ mot_layers_mapping = {i:n for n, i in enumerate(mot_layers)}
61
+ state_dict_ = {}
62
+ for name in state_dict:
63
+ if "_mot_ref" not in name:
64
+ continue
65
+ param = state_dict[name]
66
+ name = name.replace("_mot_ref", "")
67
+ if name in rename_dict:
68
+ state_dict_[rename_dict[name]] = param
69
+ else:
70
+ if name.split(".")[1].isdigit():
71
+ block_id = int(name.split(".")[1])
72
+ name = name.replace(str(block_id), str(mot_layers_mapping[block_id]))
73
+ name_ = ".".join(name.split(".")[:1] + ["0"] + name.split(".")[2:])
74
+ if name_ in rename_dict:
75
+ name_ = rename_dict[name_]
76
+ name_ = ".".join(name_.split(".")[:1] + [name.split(".")[1]] + name_.split(".")[2:])
77
+ state_dict_[name_] = param
78
+ return state_dict_
diffsynth/utils/state_dict_converters/wan_video_vace.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ def VaceWanModelDictConverter(state_dict):
2
+ state_dict_ = {name: state_dict[name] for name in state_dict if name.startswith("vace")}
3
+ return state_dict_
diffsynth/utils/state_dict_converters/wan_video_vae.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ def WanVideoVAEStateDictConverter(state_dict):
2
+ state_dict_ = {}
3
+ if 'model_state' in state_dict:
4
+ state_dict = state_dict['model_state']
5
+ for name in state_dict:
6
+ state_dict_['model.' + name] = state_dict[name]
7
+ return state_dict_
diffsynth/utils/state_dict_converters/wans2v_audio_encoder.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def WanS2VAudioEncoderStateDictConverter(state_dict):
2
+ rename_dict = {
3
+ "model.wav2vec2.encoder.pos_conv_embed.conv.weight_g": "model.wav2vec2.encoder.pos_conv_embed.conv.parametrizations.weight.original0",
4
+ "model.wav2vec2.encoder.pos_conv_embed.conv.weight_v": "model.wav2vec2.encoder.pos_conv_embed.conv.parametrizations.weight.original1",
5
+ }
6
+ state_dict_ = {}
7
+ for name in state_dict:
8
+ name_ = "model." + name
9
+ if name_ in rename_dict:
10
+ name_ = rename_dict[name_]
11
+ state_dict_[name_] = state_dict[name]
12
+ return state_dict_
diffsynth/utils/state_dict_converters/z_image_dit.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ def ZImageDiTStateDictConverter(state_dict):
2
+ state_dict_ = {name.replace("model.diffusion_model.", ""): state_dict[name] for name in state_dict}
3
+ return state_dict_
diffsynth/utils/state_dict_converters/z_image_text_encoder.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ def ZImageTextEncoderStateDictConverter(state_dict):
2
+ state_dict_ = {}
3
+ for name in state_dict:
4
+ if name != "lm_head.weight":
5
+ state_dict_[name] = state_dict[name]
6
+ return state_dict_
diffsynth/utils/tile/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .tile_worker import TileWorker
diffsynth/utils/tile/tile_worker.py ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from einops import repeat, rearrange
2
+ from tqdm import tqdm
3
+ import torch
4
+
5
+
6
+ class TileWorker:
7
+ def __init__(self):
8
+ pass
9
+
10
+ def build_mask(self, data, is_bound):
11
+ H, W = data.shape[:2]
12
+ h = repeat(torch.arange(H), "H -> H W", H=H, W=W)
13
+ w = repeat(torch.arange(W), "W -> H W", H=H, W=W)
14
+ border_width = (H + W) // 4
15
+ pad = torch.ones_like(h) * border_width
16
+ mask = torch.stack([
17
+ pad if is_bound[0] else h + 1,
18
+ pad if is_bound[1] else H - h,
19
+ pad if is_bound[2] else w + 1,
20
+ pad if is_bound[3] else W - w
21
+ ]).min(dim=0).values
22
+ mask = mask.clip(1, border_width)
23
+ mask = (mask / border_width).to(dtype=data.dtype, device=data.device)
24
+ mask = rearrange(mask, "H W -> H W 1")
25
+ return mask
26
+
27
+ def tiled_forward(self, forward_fn, channels, tile_size, tile_stride, tile_range, output_scale=1, device="cpu", dtype=torch.float32, border_width=None, progress_bar=tqdm):
28
+ # Prepare
29
+ H, W = tile_range
30
+ border_width = int(tile_stride*0.5) if border_width is None else border_width
31
+ weight = torch.zeros((H, W, 1), dtype=dtype, device=device)
32
+ values = torch.zeros((H, W, channels), dtype=dtype, device=device)
33
+
34
+ # Split tasks
35
+ tasks = []
36
+ for h in range(0, H, tile_stride):
37
+ for w in range(0, W, tile_stride):
38
+ if (h-tile_stride >= 0 and h-tile_stride+tile_size >= H) or (w-tile_stride >= 0 and w-tile_stride+tile_size >= W):
39
+ continue
40
+ h_, w_ = h + tile_size, w + tile_size
41
+ if h_ > H: h, h_ = H - tile_size, H
42
+ if w_ > W: w, w_ = W - tile_size, W
43
+ tasks.append((h, h_, w, w_))
44
+
45
+ # Run
46
+ for hl, hr, wl, wr in progress_bar(tasks):
47
+ # Forward
48
+ x = forward_fn(hl, hr, wl, wr).to(dtype=dtype, device=device)
49
+ mask = self.build_mask(x, is_bound=(hl==0, hr>=H, wl==0, wr>=W))
50
+ hl, hr = int(hl * output_scale), int(hr * output_scale)
51
+ wl, wr = int(wl * output_scale), int(wr * output_scale)
52
+ values[hl:hr, wl:wr] += x * mask
53
+ weight[hl:hr, wl:wr] += mask
54
+ values /= weight
55
+ return values
diffsynth/utils/xfuser/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .xdit_context_parallel import usp_attn_forward, usp_dit_forward, usp_vace_forward, get_sequence_parallel_world_size, get_sequence_parallel_rank, get_sp_group, initialize_usp, get_current_chunk, gather_all_chunks, all_to_all_4d, is_evenly_divisible
diffsynth/utils/xfuser/xdit_context_parallel.py ADDED
@@ -0,0 +1,219 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from typing import Optional
3
+ from einops import rearrange
4
+ from yunchang.kernels import AttnType
5
+ from yunchang.comm.all_to_all import SeqAllToAll4D
6
+ from xfuser.core.distributed import (get_sequence_parallel_rank,
7
+ get_sequence_parallel_world_size,
8
+ get_sp_group)
9
+ from xfuser.core.long_ctx_attention import xFuserLongContextAttention
10
+
11
+ from ... import IS_NPU_AVAILABLE
12
+ from ...core.device import parse_nccl_backend, parse_device_type
13
+ from ...core.gradient import gradient_checkpoint_forward
14
+
15
+
16
+ def initialize_usp(device_type):
17
+ import torch.distributed as dist
18
+ from xfuser.core.distributed import initialize_model_parallel, init_distributed_environment
19
+ dist.init_process_group(backend=parse_nccl_backend(device_type), init_method="env://")
20
+ init_distributed_environment(rank=dist.get_rank(), world_size=dist.get_world_size())
21
+ initialize_model_parallel(
22
+ sequence_parallel_degree=dist.get_world_size(),
23
+ ring_degree=1,
24
+ ulysses_degree=dist.get_world_size(),
25
+ )
26
+ getattr(torch, device_type).set_device(dist.get_rank())
27
+
28
+
29
+ def sinusoidal_embedding_1d(dim, position):
30
+ sinusoid = torch.outer(position.type(torch.float64), torch.pow(
31
+ 10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
32
+ x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
33
+ return x.to(position.dtype)
34
+
35
+ def pad_freqs(original_tensor, target_len):
36
+ seq_len, s1, s2 = original_tensor.shape
37
+ pad_size = target_len - seq_len
38
+ original_tensor_device = original_tensor.device
39
+ if original_tensor.device.type == "npu":
40
+ original_tensor = original_tensor.cpu()
41
+ padding_tensor = torch.ones(
42
+ pad_size,
43
+ s1,
44
+ s2,
45
+ dtype=original_tensor.dtype,
46
+ device=original_tensor.device)
47
+ padded_tensor = torch.cat([original_tensor, padding_tensor], dim=0).to(device=original_tensor_device)
48
+ return padded_tensor
49
+
50
+ def rope_apply(x, freqs, num_heads):
51
+ x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
52
+ s_per_rank = x.shape[1]
53
+
54
+ x_out = torch.view_as_complex(x.to(torch.float64).reshape(
55
+ x.shape[0], x.shape[1], x.shape[2], -1, 2))
56
+
57
+ sp_size = get_sequence_parallel_world_size()
58
+ sp_rank = get_sequence_parallel_rank()
59
+ freqs = pad_freqs(freqs, s_per_rank * sp_size)
60
+ freqs_rank = freqs[(sp_rank * s_per_rank):((sp_rank + 1) * s_per_rank), :, :]
61
+ freqs_rank = freqs_rank.to(torch.complex64) if freqs_rank.device.type == "npu" else freqs_rank
62
+ x_out = torch.view_as_real(x_out * freqs_rank).flatten(2)
63
+ return x_out.to(x.dtype)
64
+
65
+ def usp_dit_forward(self,
66
+ x: torch.Tensor,
67
+ timestep: torch.Tensor,
68
+ context: torch.Tensor,
69
+ clip_feature: Optional[torch.Tensor] = None,
70
+ y: Optional[torch.Tensor] = None,
71
+ use_gradient_checkpointing: bool = False,
72
+ use_gradient_checkpointing_offload: bool = False,
73
+ **kwargs,
74
+ ):
75
+ t = self.time_embedding(
76
+ sinusoidal_embedding_1d(self.freq_dim, timestep))
77
+ t_mod = self.time_projection(t).unflatten(1, (6, self.dim))
78
+ context = self.text_embedding(context)
79
+
80
+ if self.has_image_input:
81
+ x = torch.cat([x, y], dim=1) # (b, c_x + c_y, f, h, w)
82
+ clip_embdding = self.img_emb(clip_feature)
83
+ context = torch.cat([clip_embdding, context], dim=1)
84
+
85
+ x, (f, h, w) = self.patchify(x)
86
+
87
+ freqs = torch.cat([
88
+ self.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
89
+ self.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
90
+ self.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
91
+ ], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
92
+
93
+ # Context Parallel
94
+ chunks = torch.chunk(x, get_sequence_parallel_world_size(), dim=1)
95
+ pad_shape = chunks[0].shape[1] - chunks[-1].shape[1]
96
+ chunks = [torch.nn.functional.pad(chunk, (0, 0, 0, chunks[0].shape[1]-chunk.shape[1]), value=0) for chunk in chunks]
97
+ x = chunks[get_sequence_parallel_rank()]
98
+
99
+ for block in self.blocks:
100
+ if self.training:
101
+ x = gradient_checkpoint_forward(
102
+ block,
103
+ use_gradient_checkpointing,
104
+ use_gradient_checkpointing_offload,
105
+ x, context, t_mod, freqs
106
+ )
107
+ else:
108
+ x = block(x, context, t_mod, freqs)
109
+
110
+ x = self.head(x, t)
111
+
112
+ # Context Parallel
113
+ x = get_sp_group().all_gather(x, dim=1)
114
+ x = x[:, :-pad_shape] if pad_shape > 0 else x
115
+
116
+ # unpatchify
117
+ x = self.unpatchify(x, (f, h, w))
118
+ return x
119
+
120
+
121
+ def usp_vace_forward(
122
+ self, x, vace_context, context, t_mod, freqs,
123
+ use_gradient_checkpointing: bool = False,
124
+ use_gradient_checkpointing_offload: bool = False,
125
+ ):
126
+ # Compute full sequence length from the sharded x
127
+ full_seq_len = x.shape[1] * get_sequence_parallel_world_size()
128
+
129
+ # Embed vace_context via patch embedding
130
+ c = [self.vace_patch_embedding(u.unsqueeze(0)) for u in vace_context]
131
+ c = [u.flatten(2).transpose(1, 2) for u in c]
132
+ c = torch.cat([
133
+ torch.cat([u, u.new_zeros(1, full_seq_len - u.size(1), u.size(2))],
134
+ dim=1) for u in c
135
+ ])
136
+
137
+ # Chunk VACE context along sequence dim BEFORE processing through blocks
138
+ c = torch.chunk(c, get_sequence_parallel_world_size(), dim=1)[get_sequence_parallel_rank()]
139
+
140
+ # Process through vace_blocks (self_attn already monkey-patched to usp_attn_forward)
141
+ for block in self.vace_blocks:
142
+ c = gradient_checkpoint_forward(
143
+ block,
144
+ use_gradient_checkpointing,
145
+ use_gradient_checkpointing_offload,
146
+ c, x, context, t_mod, freqs
147
+ )
148
+
149
+ # Hints are already sharded per-rank
150
+ hints = torch.unbind(c)[:-1]
151
+ return hints
152
+
153
+
154
+ def usp_attn_forward(self, x, freqs):
155
+ q = self.norm_q(self.q(x))
156
+ k = self.norm_k(self.k(x))
157
+ v = self.v(x)
158
+
159
+ q = rope_apply(q, freqs, self.num_heads)
160
+ k = rope_apply(k, freqs, self.num_heads)
161
+ q = rearrange(q, "b s (n d) -> b s n d", n=self.num_heads)
162
+ k = rearrange(k, "b s (n d) -> b s n d", n=self.num_heads)
163
+ v = rearrange(v, "b s (n d) -> b s n d", n=self.num_heads)
164
+
165
+ attn_type = AttnType.FA
166
+ ring_impl_type = "basic"
167
+ if IS_NPU_AVAILABLE:
168
+ attn_type = AttnType.NPU
169
+ ring_impl_type = "basic_npu"
170
+ x = xFuserLongContextAttention(attn_type=attn_type, ring_impl_type=ring_impl_type)(
171
+ None,
172
+ query=q,
173
+ key=k,
174
+ value=v,
175
+ )
176
+ x = x.flatten(2)
177
+
178
+ del q, k, v
179
+ getattr(torch, parse_device_type(x.device)).empty_cache()
180
+ return self.o(x)
181
+
182
+
183
+ def get_current_chunk(x, dim=1):
184
+ chunks = torch.chunk(x, get_sequence_parallel_world_size(), dim=dim)
185
+ ndims = len(chunks[0].shape)
186
+ pad_list = [0] * (2 * ndims)
187
+ pad_end_index = 2 * (ndims - 1 - dim) + 1
188
+ max_size = chunks[0].size(dim)
189
+ chunks = [
190
+ torch.nn.functional.pad(
191
+ chunk,
192
+ tuple(pad_list[:pad_end_index] + [max_size - chunk.size(dim)] + pad_list[pad_end_index+1:]),
193
+ value=0
194
+ )
195
+ for chunk in chunks
196
+ ]
197
+ x = chunks[get_sequence_parallel_rank()]
198
+ return x
199
+
200
+
201
+ def gather_all_chunks(x, seq_len=None, dim=1):
202
+ x = get_sp_group().all_gather(x, dim=dim)
203
+ if seq_len is not None:
204
+ slices = [slice(None)] * x.ndim
205
+ slices[dim] = slice(0, seq_len)
206
+ x = x[tuple(slices)]
207
+ return x
208
+
209
+
210
+ def all_to_all_4d(x, scatter_dim, gather_dim):
211
+ world_size = get_sequence_parallel_world_size()
212
+ if world_size == 1:
213
+ return x
214
+ return SeqAllToAll4D.apply(get_sp_group().ulysses_group, x, scatter_dim, gather_dim)
215
+
216
+
217
+ def is_evenly_divisible(seq_len):
218
+ world_size = get_sequence_parallel_world_size()
219
+ return seq_len % world_size == 0
diffsynth/version.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ # Make sure to modify __release_datetime__ to release time when making official release.
2
+ __version__ = '2.1.5'
3
+ # default release datetime for branches under active development is set
4
+ # to be a time far-far-away-into-the-future
5
+ __release_datetime__ = '2099-10-13 08:56:12'
docs/en/.readthedocs.yaml ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # .readthedocs.yaml
2
+ # Read the Docs configuration file
3
+ # See https://docs.readthedocs.io/en/stable/config-file/v2.html for details
4
+
5
+ # Required
6
+ version: 2
7
+
8
+ # Set the OS, Python version and other tools you might need
9
+ build:
10
+ os: ubuntu-22.04
11
+ tools:
12
+ python: "3.10"
13
+
14
+ # Build documentation in the "docs/" directory with Sphinx
15
+ sphinx:
16
+ configuration: docs/en/conf.py
17
+
18
+ # Optionally build your docs in additional formats such as PDF and ePub
19
+ # formats:
20
+ # - pdf
21
+ # - epub
22
+
23
+ # Optional but recommended, declare the Python requirements required
24
+ # to build your documentation
25
+ # See https://docs.readthedocs.io/en/stable/guides/reproducible-builds.html
26
+ python:
27
+ install:
28
+ - requirements: docs/requirements.txt
docs/en/API_Reference/core/attention.md ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `diffsynth.core.attention`: Attention Mechanism Implementation
2
+
3
+ `diffsynth.core.attention` provides routing mechanisms for attention mechanism implementations, automatically selecting efficient attention implementations based on available packages in the `Python` environment and [environment variables](../../Pipeline_Usage/Environment_Variables.md#diffsynth_attention_implementation).
4
+
5
+ ## Attention Mechanism
6
+
7
+ The attention mechanism is a model structure proposed in the paper ["Attention Is All You Need"](https://arxiv.org/abs/1706.03762). In the original paper, the attention mechanism is implemented according to the following formula:
8
+
9
+ $$
10
+ \text{Attention}(Q, K, V) = \text{Softmax}\left(
11
+ \frac{QK^T}{\sqrt{d_k}}
12
+ \right)
13
+ V.
14
+ $$
15
+
16
+ In `PyTorch`, it can be implemented with the following code:
17
+ ```python
18
+ import torch
19
+
20
+ def attention(query, key, value):
21
+ scale_factor = 1 / query.size(-1)**0.5
22
+ attn_weight = query @ key.transpose(-2, -1) * scale_factor
23
+ attn_weight = torch.softmax(attn_weight, dim=-1)
24
+ return attn_weight @ value
25
+
26
+ query = torch.rand(32, 8, 128, 64, dtype=torch.bfloat16, device="cuda")
27
+ key = torch.rand(32, 8, 128, 64, dtype=torch.bfloat16, device="cuda")
28
+ value = torch.rand(32, 8, 128, 64, dtype=torch.bfloat16, device="cuda")
29
+ output_1 = attention(query, key, value)
30
+ ```
31
+
32
+ The dimensions of `query`, `key`, and `value` are $(b, n, s, d)$:
33
+ * $b$: Batch size
34
+ * $n$: Number of attention heads
35
+ * $s$: Sequence length
36
+ * $d$: Dimension of each attention head
37
+
38
+ This computation does not include any trainable parameters. Modern transformer architectures will pass through Linear layers before and after this computation, but the "attention mechanism" discussed in this article refers only to the computation in the above code, not including these calculations.
39
+
40
+ ## More Efficient Implementations
41
+
42
+ Note that the dimension of the Attention Score in the attention mechanism ( $\text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)$ in the formula, `attn_weight` in the code) is $(b, n, s, s)$, where the sequence length $s$ is typically very large, causing the time and space complexity of computation to reach quadratic level. Taking image generation models as an example, when the width and height of the image increase to 2 times, the sequence length increases to 4 times, and the computational load and memory requirements increase to 16 times. To avoid high computational costs, more efficient attention mechanism implementations are needed, including:
43
+ * Flash Attention 4: [GitHub](https://github.com/Dao-AILab/flash-attention), [Paper](https://arxiv.org/abs/2603.05451)
44
+ * Flash Attention 3: [GitHub](https://github.com/Dao-AILab/flash-attention), [Paper](https://arxiv.org/abs/2407.08608)
45
+ * Flash Attention 2: [GitHub](https://github.com/Dao-AILab/flash-attention), [Paper](https://arxiv.org/abs/2307.08691)
46
+ * Sage Attention: [GitHub](https://github.com/thu-ml/SageAttention), [Paper](https://arxiv.org/abs/2505.11594)
47
+ * xFormers: [GitHub](https://github.com/facebookresearch/xformers), [Documentation](https://facebookresearch.github.io/xformers/components/ops.html#module-xformers.ops)
48
+ * PyTorch: [GitHub](https://github.com/pytorch/pytorch), [Documentation](https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html)
49
+
50
+ To call attention implementations other than `PyTorch`, please follow the instructions on their GitHub pages to install the corresponding packages. `DiffSynth-Studio` will automatically route to the corresponding implementation based on available packages in the Python environment, or can be controlled through [environment variables](../../Pipeline_Usage/Environment_Variables.md#diffsynth_attention_implementation).
51
+
52
+ ```python
53
+ from diffsynth.core.attention import attention_forward
54
+ import torch
55
+
56
+ def attention(query, key, value):
57
+ scale_factor = 1 / query.size(-1)**0.5
58
+ attn_weight = query @ key.transpose(-2, -1) * scale_factor
59
+ attn_weight = torch.softmax(attn_weight, dim=-1)
60
+ return attn_weight @ value
61
+
62
+ query = torch.rand(32, 8, 128, 64, dtype=torch.bfloat16, device="cuda")
63
+ key = torch.rand(32, 8, 128, 64, dtype=torch.bfloat16, device="cuda")
64
+ value = torch.rand(32, 8, 128, 64, dtype=torch.bfloat16, device="cuda")
65
+ output_1 = attention(query, key, value)
66
+ output_2 = attention_forward(query, key, value)
67
+ print((output_1 - output_2).abs().mean())
68
+ ```
69
+
70
+ Please note that acceleration will introduce errors, but in most cases, the error is negligible.
71
+
72
+ ## Developer Guide
73
+
74
+ When integrating new models into `DiffSynth-Studio`, developers can decide whether to call `attention_forward` in `diffsynth.core.attention`, but we expect models to prioritize calling this module as much as possible, so that new attention mechanism implementations can take effect directly on these models.
75
+
76
+ ## Best Practices
77
+
78
+ **In most cases, we recommend directly using the native `PyTorch` implementation without installing any additional packages.** Although other attention mechanism implementations can accelerate, the acceleration effect is relatively limited, and in a few cases, compatibility and precision issues may arise.
79
+
80
+ In addition, efficient attention mechanism implementations will gradually be integrated into `PyTorch`. The `scaled_dot_product_attention` in `PyTorch` version 2.9.0 has already integrated Flash Attention 2. We still provide this interface in `DiffSynth-Studio` to allow some aggressive acceleration schemes to quickly move toward application, even though they still need time to be verified for stability.
docs/en/API_Reference/core/data.md ADDED
@@ -0,0 +1,151 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `diffsynth.core.data`: Data Processing Operators and Universal Dataset
2
+
3
+ ## Data Processing Operators
4
+
5
+ ### Available Data Processing Operators
6
+
7
+ `diffsynth.core.data` provides a series of data processing operators for data processing, including:
8
+
9
+ * Data format conversion operators
10
+ * `ToInt`: Convert to int format
11
+ * `ToFloat`: Convert to float format
12
+ * `ToStr`: Convert to str format
13
+ * `ToList`: Convert to list format, wrapping this data in a list
14
+ * `ToAbsolutePath`: Convert relative paths to absolute paths
15
+ * File loading operators
16
+ * `LoadImage`: Read image files
17
+ * `LoadVideo`: Read video files
18
+ * `LoadAudio`: Read audio files
19
+ * `LoadGIF`: Read GIF files
20
+ * `LoadTorchPickle`: Read binary files saved by [`torch.save`](https://docs.pytorch.org/docs/stable/generated/torch.save.html) [This operator may cause code injection attacks in binary files, please use with caution!]
21
+ * Media file processing operators
22
+ * `ImageCropAndResize`: Crop and resize images
23
+ * Meta operators
24
+ * `SequencialProcess`: Route each data in the sequence to an operator
25
+ * `RouteByExtensionName`: Route to specific operators by file extension
26
+ * `RouteByType`: Route to specific operators by data type
27
+
28
+ ### Operator Usage
29
+
30
+ Data operators are connected with the `>>` symbol to form data processing pipelines, for example:
31
+
32
+ ```python
33
+ from diffsynth.core.data.operators import *
34
+
35
+ data = "image.jpg"
36
+ data_pipeline = ToAbsolutePath(base_path="/data") >> LoadImage() >> ImageCropAndResize(max_pixels=512*512)
37
+ data = data_pipeline(data)
38
+ ```
39
+
40
+ After passing through each operator, the data is processed in sequence:
41
+
42
+ * `ToAbsolutePath(base_path="/data")`: `"/data/image.jpg"`
43
+ * `LoadImage()`: `<PIL.Image.Image image mode=RGB size=1024x1024 at 0x7F8E7AAEFC10>`
44
+ * `ImageCropAndResize(max_pixels=512*512)`: `<PIL.Image.Image image mode=RGB size=512x512 at 0x7F8E7A936F20>`
45
+
46
+ We can compose functionally complete data pipelines, for example, the default video data operator for the universal dataset is:
47
+
48
+ ```python
49
+ RouteByType(operator_map=[
50
+ (str, ToAbsolutePath(base_path) >> RouteByExtensionName(operator_map=[
51
+ (("jpg", "jpeg", "png", "webp"), LoadImage() >> ImageCropAndResize(height, width, max_pixels, height_division_factor, width_division_factor) >> ToList()),
52
+ (("gif",), LoadGIF(
53
+ num_frames, time_division_factor, time_division_remainder,
54
+ frame_processor=ImageCropAndResize(height, width, max_pixels, height_division_factor, width_division_factor),
55
+ )),
56
+ (("mp4", "avi", "mov", "wmv", "mkv", "flv", "webm"), LoadVideo(
57
+ num_frames, time_division_factor, time_division_remainder,
58
+ frame_processor=ImageCropAndResize(height, width, max_pixels, height_division_factor, width_division_factor),
59
+ )),
60
+ ])),
61
+ ])
62
+ ```
63
+
64
+ It includes the following logic:
65
+
66
+ * If the data is of type `str`
67
+ * If it's a `"jpg", "jpeg", "png", "webp"` type file
68
+ * Load this image
69
+ * Crop and scale to a specific resolution
70
+ * Pack into a list, treating it as a single-frame video
71
+ * If it's a `"gif"` type file
72
+ * Load the GIF file content
73
+ * Crop and scale each frame to a specific resolution
74
+ * If it's a `"mp4", "avi", "mov", "wmv", "mkv", "flv", "webm"` type file
75
+ * Load the video file content
76
+ * Crop and scale each frame to a specific resolution
77
+ * If the data is not of type `str`, an error is reported
78
+
79
+ ## Universal Dataset
80
+
81
+ `diffsynth.core.data` provides a unified dataset implementation. The dataset requires the following parameters:
82
+
83
+ * `base_path`: Root directory. If the dataset contains relative paths to image files, this field needs to be filled in to load the files pointed to by these paths
84
+ * `metadata_path`: Metadata directory, records the file paths of all metadata, supports `csv`, `json`, `jsonl` formats
85
+ * `repeat`: Data repetition count, defaults to 1, this parameter affects the number of training steps in an epoch
86
+ * `data_file_keys`: Data field names that need to be loaded, for example `(image, edit_image)`
87
+ * `main_data_operator`: Main loading operator, needs to assemble the data processing pipeline through data processing operators
88
+ * `special_operator_map`: Special operator mapping, operator mappings built for fields that require special processing
89
+
90
+ ### Metadata
91
+
92
+ The dataset's `metadata_path` points to a metadata file, supporting `csv`, `json`, `jsonl` formats. The following provides examples:
93
+
94
+ * `csv` format: High readability, does not support list data, small memory footprint
95
+
96
+ ```csv
97
+ image,prompt
98
+ image_1.jpg,"a dog"
99
+ image_2.jpg,"a cat"
100
+ ```
101
+
102
+ * `json` format: High readability, supports list data, large memory footprint
103
+
104
+ ```json
105
+ [
106
+ {
107
+ "image": "image_1.jpg",
108
+ "prompt": "a dog"
109
+ },
110
+ {
111
+ "image": "image_2.jpg",
112
+ "prompt": "a cat"
113
+ }
114
+ ]
115
+ ```
116
+
117
+ * `jsonl` format: Low readability, supports list data, small memory footprint
118
+
119
+ ```json
120
+ {"image": "image_1.jpg", "prompt": "a dog"}
121
+ {"image": "image_2.jpg", "prompt": "a cat"}
122
+ ```
123
+
124
+ How to choose the best metadata format?
125
+
126
+ * If the data volume is large, reaching tens of millions, since `json` file parsing requires additional memory, it's not available. Please use `csv` or `jsonl` format
127
+ * If the dataset contains list data, such as edit models that require multiple images as input, since `csv` format cannot store list format data, it's not available. Please use `json` or `jsonl` format
128
+
129
+ ### Data Loading Logic
130
+
131
+ When no additional settings are made, the dataset defaults to outputting data from the metadata set. Image and video file paths will be output in string format. To load these files, you need to set `data_file_keys`, `main_data_operator`, and `special_operator_map`.
132
+
133
+ In the data processing flow, processing is done according to the following logic:
134
+ * If the field is in `special_operator_map`, call the corresponding operator in `special_operator_map` for processing
135
+ * If the field is not in `special_operator_map`
136
+ * If the field is in `data_file_keys`, call the `main_data_operator` operator for processing
137
+ * If the field is not in `data_file_keys`, no processing is done
138
+
139
+ `special_operator_map` can be used to implement special data processing. For example, in the model [Wan-AI/Wan2.2-Animate-14B](https://www.modelscope.cn/models/Wan-AI/Wan2.2-Animate-14B), the input character face video `animate_face_video` is processed at a fixed resolution, inconsistent with the output video. Therefore, this field is processed by a dedicated operator:
140
+
141
+ ```python
142
+ special_operator_map={
143
+ "animate_face_video": ToAbsolutePath(args.dataset_base_path) >> LoadVideo(args.num_frames, 4, 1, frame_processor=ImageCropAndResize(512, 512, None, 16, 16)),
144
+ }
145
+ ```
146
+
147
+ ### Other Notes
148
+
149
+ When the data volume is too small, you can appropriately increase `repeat` to extend the training time of a single epoch, avoiding frequent model saving that generates considerable overhead.
150
+
151
+ When data volume * `repeat` exceeds $10^9$, we observe that the dataset speed becomes significantly slower. This seems to be a `PyTorch` bug, and we are not sure if newer versions of `PyTorch` have fixed this issue.
docs/en/API_Reference/core/gradient.md ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `diffsynth.core.gradient`: Gradient Checkpointing and Offload
2
+
3
+ `diffsynth.core.gradient` provides encapsulated gradient checkpointing and its Offload version for model training.
4
+
5
+ ## Gradient Checkpointing
6
+
7
+ Gradient checkpointing is a technique used to reduce memory usage during training. We provide an example to help you understand this technique. Here is a simple model structure:
8
+
9
+ ```python
10
+ import torch
11
+
12
+ class ToyModel(torch.nn.Module):
13
+ def __init__(self):
14
+ super().__init__()
15
+ self.activation = torch.nn.Sigmoid()
16
+
17
+ def forward(self, x):
18
+ return self.activation(x)
19
+
20
+ model = ToyModel()
21
+ x = torch.randn((2, 3))
22
+ y = model(x)
23
+ ```
24
+
25
+ In this model structure, the input parameter $x$ passes through the Sigmoid activation function to obtain the output value $y=\frac{1}{1+e^{-x}}$.
26
+
27
+ During the training process, assuming our loss function value is $\mathcal L$, when backpropagating gradients, we obtain $\frac{\partial \mathcal L}{\partial y}$. At this point, we need to calculate $\frac{\partial \mathcal L}{\partial x}$. It's not difficult to find that $\frac{\partial y}{\partial x}=y(1-y)$, and thus $\frac{\partial \mathcal L}{\partial x}=\frac{\partial \mathcal L}{\partial y}\frac{\partial y}{\partial x}=\frac{\partial \mathcal L}{\partial y}y(1-y)$. If we save the value of $y$ during the model's forward propagation and directly compute $y(1-y)$ during gradient backpropagation, this will avoid complex exp computations, speeding up the calculation. However, this requires additional memory to store the intermediate variable $y$.
28
+
29
+ When gradient checkpointing is not enabled, the training framework will default to storing all intermediate variables that assist gradient computation, thereby achieving optimal computational speed. When gradient checkpointing is enabled, intermediate variables are not stored, but the input parameter $x$ is still stored, reducing memory usage. During gradient backpropagation, these variables need to be recomputed, slowing down the calculation.
30
+
31
+ ## Enabling Gradient Checkpointing and Its Offload
32
+
33
+ `gradient_checkpoint_forward` in `diffsynth.core.gradient` implements gradient checkpointing and its Offload. Refer to the following code for calling:
34
+
35
+ ```python
36
+ import torch
37
+ from diffsynth.core.gradient import gradient_checkpoint_forward
38
+
39
+ class ToyModel(torch.nn.Module):
40
+ def __init__(self):
41
+ super().__init__()
42
+ self.activation = torch.nn.Sigmoid()
43
+
44
+ def forward(self, x):
45
+ return self.activation(x)
46
+
47
+ model = ToyModel()
48
+ x = torch.randn((2, 3))
49
+ y = gradient_checkpoint_forward(
50
+ model,
51
+ use_gradient_checkpointing=True,
52
+ use_gradient_checkpointing_offload=False,
53
+ x=x,
54
+ )
55
+ ```
56
+
57
+ * When `use_gradient_checkpointing=False` and `use_gradient_checkpointing_offload=False`, the computation process is exactly the same as the original computation, not affecting the model's inference and training. You can directly integrate it into your code.
58
+ * When `use_gradient_checkpointing=True` and `use_gradient_checkpointing_offload=False`, gradient checkpointing is enabled.
59
+ * When `use_gradient_checkpointing_offload=True`, gradient checkpointing is enabled, and all gradient checkpoint input parameters are stored in memory, further reducing memory usage and slowing down computation.
60
+
61
+ ## Best Practices
62
+
63
+ > Q: Where should gradient checkpointing be enabled?
64
+ >
65
+ > A: When enabling gradient checkpointing for the entire model, computational efficiency and memory usage are not optimal. We need to set fine-grained gradient checkpoints, but we don't want to add too much complicated code to the framework. Therefore, we recommend implementing it in the `model_fn` of `Pipeline`, for example, `model_fn_qwen_image` in `diffsynth/pipelines/qwen_image.py`, enabling gradient checkpointing at the Block level without modifying any code in the model structure.
66
+
67
+ > Q: When should gradient checkpointing be enabled?
68
+ >
69
+ > A: As model parameters become increasingly large, gradient checkpointing has become a necessary training technique. Gradient checkpointing usually needs to be enabled. Gradient checkpointing Offload should only be enabled in models where activation values occupy excessive memory (such as video generation models).
docs/en/API_Reference/core/loader.md ADDED
@@ -0,0 +1,141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `diffsynth.core.loader`: Model Download and Loading
2
+
3
+ This document introduces the model download and loading functionalities in `diffsynth.core.loader`.
4
+
5
+ ## ModelConfig
6
+
7
+ `ModelConfig` in `diffsynth.core.loader` is used to annotate model download sources, local paths, VRAM management configurations, and other information.
8
+
9
+ ### Downloading and Loading Models from Remote Sources
10
+
11
+ Taking the model [DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny](https://www.modelscope.cn/models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny) as an example, after filling in `model_id` and `origin_file_pattern` in `ModelConfig`, the model can be automatically downloaded. By default, it downloads to the `./models` path, which can be modified through the [environment variable DIFFSYNTH_MODEL_BASE_PATH](../../Pipeline_Usage/Environment_Variables.md#diffsynth_model_base_path).
12
+
13
+ By default, even if the model has already been downloaded, the program will still query the remote for any missing files. To completely disable remote requests, set the [environment variable DIFFSYNTH_SKIP_DOWNLOAD](../../Pipeline_Usage/Environment_Variables.md#diffsynth_skip_download) to `True`.
14
+
15
+ ```python
16
+ from diffsynth.core import ModelConfig
17
+
18
+ config = ModelConfig(
19
+ model_id="DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny",
20
+ origin_file_pattern="model.safetensors",
21
+ )
22
+ # Download models
23
+ config.download_if_necessary()
24
+ print(config.path)
25
+ ```
26
+
27
+ After calling `download_if_necessary`, the model will be automatically downloaded, and the path will be returned to `config.path`.
28
+
29
+ ### Loading Models from Local Paths
30
+
31
+ If loading models from local paths, you need to fill in `path`:
32
+
33
+ ```python
34
+ from diffsynth.core import ModelConfig
35
+
36
+ config = ModelConfig(path="models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors")
37
+ ```
38
+
39
+ If the model contains multiple shard files, input them in list form:
40
+
41
+ ```python
42
+ from diffsynth.core import ModelConfig
43
+
44
+ config = ModelConfig(path=[
45
+ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors",
46
+ "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors",
47
+ "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors",
48
+ "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors"
49
+ ])
50
+ ```
51
+
52
+ ### VRAM Management Configuration
53
+
54
+ `ModelConfig` also contains VRAM management configuration information. See [VRAM Management](../../Pipeline_Usage/VRAM_management.md#more-usage-methods) for details.
55
+
56
+ ## Model File Loading
57
+
58
+ `diffsynth.core.loader` provides a unified `load_state_dict` for loading state dicts from model files.
59
+
60
+ Loading a single model file:
61
+
62
+ ```python
63
+ from diffsynth.core import load_state_dict
64
+
65
+ state_dict = load_state_dict("models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors")
66
+ ```
67
+
68
+ Loading multiple model files (merged into one state dict):
69
+
70
+ ```python
71
+ from diffsynth.core import load_state_dict
72
+
73
+ state_dict = load_state_dict([
74
+ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors",
75
+ "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors",
76
+ "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors",
77
+ "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors"
78
+ ])
79
+ ```
80
+
81
+ ## Model Hash
82
+
83
+ Model hash is used to determine the model type. The hash value can be obtained through `hash_model_file`:
84
+
85
+ ```python
86
+ from diffsynth.core import hash_model_file
87
+
88
+ print(hash_model_file("models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors"))
89
+ ```
90
+
91
+ The hash value of multiple model files can also be calculated, which is equivalent to calculating the model hash value after merging the state dict:
92
+
93
+ ```python
94
+ from diffsynth.core import hash_model_file
95
+
96
+ print(hash_model_file([
97
+ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors",
98
+ "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors",
99
+ "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors",
100
+ "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors"
101
+ ]))
102
+ ```
103
+
104
+ The model hash value is only related to the keys and tensor shapes in the state dict of the model file, and is unrelated to the numerical values of the model parameters, file saving time, and other information. When calculating the model hash value of `.safetensors` format files, `hash_model_file` is almost instantly completed without reading the model parameters. However, when calculating the model hash value of `.bin`, `.pth`, `.ckpt`, and other binary files, all model parameters need to be read, so **we do not recommend developers to continue using these formats of files.**
105
+
106
+ By [writing model Config](../../Developer_Guide/Integrating_Your_Model.md#step-3-writing-model-config) and filling in model hash value and other information into `diffsynth/configs/model_configs.py`, developers can let `DiffSynth-Studio` automatically identify the model type and load it.
107
+
108
+ ## Model Loading
109
+
110
+ `load_model` is the external entry for loading models in `diffsynth.core.loader`. It will call [skip_model_initialization](../../API_Reference/core/vram.md#skipping-model-parameter-initialization) to skip model parameter initialization. If [Disk Offload](../../Pipeline_Usage/VRAM_management.md#disk-offload) is enabled, it calls [DiskMap](../../API_Reference/core/vram.md#state-dict-disk-mapping) for lazy loading. If Disk Offload is not enabled, it calls [load_state_dict](#model-file-loading) to load model parameters. If necessary, it will also call [state dict converter](../../Developer_Guide/Integrating_Your_Model.md#step-2-model-file-format-conversion) for model format conversion. Finally, it calls `model.eval()` to switch to inference mode.
111
+
112
+ Here is a usage example with Disk Offload enabled:
113
+
114
+ ```python
115
+ from diffsynth.core import load_model, enable_vram_management, AutoWrappedLinear, AutoWrappedModule
116
+ from diffsynth.models.qwen_image_dit import QwenImageDiT, RMSNorm
117
+ import torch
118
+
119
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
120
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
121
+
122
+ model = load_model(
123
+ QwenImageDiT,
124
+ model_path,
125
+ module_map={
126
+ torch.nn.Linear: AutoWrappedLinear,
127
+ RMSNorm: AutoWrappedModule,
128
+ },
129
+ vram_config={
130
+ "offload_dtype": "disk",
131
+ "offload_device": "disk",
132
+ "onload_dtype": "disk",
133
+ "onload_device": "disk",
134
+ "preparing_dtype": torch.bfloat16,
135
+ "preparing_device": "cuda",
136
+ "computation_dtype": torch.bfloat16,
137
+ "computation_device": "cuda",
138
+ },
139
+ vram_limit=0,
140
+ )
141
+ ```
docs/en/API_Reference/core/quant.md ADDED
@@ -0,0 +1,295 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `diffsynth.core.quant`: Model Quantization
2
+
3
+ This document introduces the low-level quantization interfaces in `diffsynth.core.quant`. Refer to it if you want to use these features in another codebase. If you only want to enable quantization in a `Pipeline`, see [Model Quantization](../../Pipeline_Usage/Quantization.md).
4
+
5
+ The module exports the following interfaces through `diffsynth.core.quant`, organized in three categories:
6
+
7
+ | Category | Interfaces |
8
+ | --- | --- |
9
+ | User interfaces | `QuantizeConfig`, `MixedQuantizeConfig`, `describe_quant_method`, `QUANT_METHODS` |
10
+ | Extension interfaces | `QuantBackend`, `BackendConfig`, `register_quant_backend`, `register_quant_method`, `QuantMethodSpec`, `QUANT_BACKENDS` |
11
+ | Verification tools | `check_differentiable`, `check_backend_contract` |
12
+
13
+ Quantization operates on the `nn.Linear` layers in a model: the framework traverses the model and replaces the matched `nn.Linear` layers with the backend's quantized Linears (all subclasses of `nn.Linear`, so LoRA injection, VRAM management, and other mechanisms recognize them without modification). A backend is only responsible for quantizing a single layer; model-level traversal and replacement is done by `QuantizeConfig`.
14
+
15
+ ## User Interfaces
16
+
17
+ ### QuantizeConfig
18
+
19
+ `QuantizeConfig` is both the quantization config and the operation entry point for any `nn.Module`.
20
+
21
+ Fields:
22
+
23
+ | Field | Type | Description |
24
+ | --- | --- | --- |
25
+ | `method` | `str` | Quantization method name, from `QUANT_METHODS`; determines the backend, scheme, and backend config. Required |
26
+ | `mode` | `str` | `"dynamic"` (default) keeps the backend-native quantized Linears, dequantizing at every forward; `"dequant_once"` restores plain fp `nn.Linear` right after the weights are quantized or loaded |
27
+ | `target_modules` | `list` | Only quantize the matched layers; `None` means no restriction |
28
+ | `exclude_modules` | `list` | Exclude the matched layers |
29
+ | `backend_config_kwargs` | `dict` | Parameters passed to the method's backend config factory, determining the quantization behavior, e.g. nf4's `blocksize` |
30
+ | `load_prequantized` | `bool` | The checkpoint already holds quantized weights; load them directly instead of quantizing online |
31
+
32
+ Matching rule for `target_modules` / `exclude_modules`: a layer matches if its full dotted name equals an entry, or ends with `"." + entry`. For example, `"img_mod.1"` matches `transformer_blocks.0.img_mod.1`.
33
+
34
+ Constructing a `QuantizeConfig` validates the backend dependencies and parameters, and raises immediately (with installation instructions) when they are not satisfied, rather than failing later at inference time.
35
+
36
+ Main methods:
37
+
38
+ #### `quantize_model(model, compute_device=None, model_device=None)`
39
+
40
+ Quantizes the matched `nn.Linear` layers in `model` in place, keeping each layer's existing dtype. Must be called **after** `load_state_dict`. Does nothing when `load_prequantized=True` (such a checkpoint is already quantized).
41
+
42
+ - `compute_device`: the device where quantization computation happens; `None` means quantize in place.
43
+ - `model_device`: the device where each layer is stored after quantization; `None` means leaving it on `compute_device`.
44
+
45
+ With an fp model on the CPU and `compute_device="cuda", model_device="cpu"`, quantization streams layer by layer, so the accelerator only ever holds one layer at a time:
46
+
47
+ ```python
48
+ import torch
49
+ from diffsynth.core.quant import QuantizeConfig
50
+
51
+ cfg = QuantizeConfig(method="bitsandbytes_nf4")
52
+ model.load_state_dict(fp_state_dict)
53
+ cfg.quantize_model(model, compute_device="cuda", model_device="cpu")
54
+ ```
55
+
56
+ #### `prepare_for_prequantized_load(model, compute_dtype=torch.bfloat16)`
57
+
58
+ Replaces the matched `nn.Linear` layers with empty quantized layers ("shells") matching the structure of a pre-quantized checkpoint. Must be called **before** `load_state_dict(assign=True)`. `compute_dtype` is the dtype the quantized layers dequantize to at forward time.
59
+
60
+ #### `unflatten_state_dict(state_dict, metadata)` / `flatten_state_dict(state_dict)`
61
+
62
+ Quantized weights are often composite structures of "packed tensors + quant state", while `.safetensors` can only store plain tensors. These two methods convert between the two forms.
63
+
64
+ - `unflatten_state_dict(state_dict, metadata)`: rebuilds composite quantized tensors from the flat tensors read out of a checkpoint; the result can be given to `load_state_dict(assign=True)`.
65
+ - `flatten_state_dict(state_dict)`: flattens a quantized model's state dict into plain tensors and string-only metadata, returning `(tensors, metadata)`, which can be passed directly to `safetensors.torch.save_file(tensors, path, metadata=metadata)`. Raises `NotImplementedError` if the backend does not declare `is_serializable`.
66
+
67
+ The complete flow for loading a pre-quantized checkpoint:
68
+
69
+ ```python
70
+ import torch
71
+ from diffsynth.core.quant import QuantizeConfig
72
+
73
+ cfg = QuantizeConfig(method="bitsandbytes_nf4", load_prequantized=True)
74
+ cfg.prepare_for_prequantized_load(model, compute_dtype=torch.bfloat16)
75
+ state_dict = cfg.unflatten_state_dict(state_dict, metadata)
76
+ model.load_state_dict(state_dict, assign=True)
77
+ ```
78
+
79
+ #### `dequantize_model(model, compute_dtype=torch.bfloat16, compute_device=None, model_device=None)`
80
+
81
+ Replaces all quantized Linears in the model with plain fp `nn.Linear`; the restored weights carry the quantization error. **Only takes effect when `mode="dequant_once"`**; otherwise returns directly. Can be called after either of the two flows above:
82
+
83
+ ```python
84
+ cfg.dequantize_model(model, compute_dtype=torch.bfloat16)
85
+ ```
86
+
87
+ #### `is_quantized_linear(module)`
88
+
89
+ Whether `module` is one of the quantized Linears produced by this config's backend.
90
+
91
+ #### `build_quantized_shell(module, compute_dtype)`
92
+
93
+ Builds an empty quantized Linear matching `module`'s shape and bias presence. Used to release a layer's weights while keeping it routable, and to stage a transient copy on the computation device — a companion interface for VRAM management.
94
+
95
+ ### MixedQuantizeConfig
96
+
97
+ Combines multiple `QuantizeConfig`s into one mixed quantization; each sub-config is responsible for a mutually disjoint set of layers. It exposes the same interface as a single `QuantizeConfig` (`quantize_model`, `prepare_for_prequantized_load`, `dequantize_model`, `flatten_state_dict`, `unflatten_state_dict`, `is_quantized_linear`, `build_quantized_shell`, plus the two read-only properties `method` / `mode`).
98
+
99
+ ```python
100
+ from diffsynth.core.quant import QuantizeConfig, MixedQuantizeConfig
101
+
102
+ mod_layers = ["img_mod.1", "txt_mod.1", "norm_out.linear", "img_in", "txt_in", "proj_out"]
103
+ cfg = MixedQuantizeConfig(configs=[
104
+ QuantizeConfig(method="bitsandbytes_nf4", exclude_modules=mod_layers),
105
+ QuantizeConfig(method="torchao_int8_w8a16", target_modules=mod_layers),
106
+ ])
107
+ cfg.quantize_model(model, compute_device="cuda")
108
+ ```
109
+
110
+ Fields and constraints:
111
+
112
+ - `configs`: a list of `QuantizeConfig`, executed in order. All sub-configs must share the same `mode`, and their `load_prequantized` must be `False`.
113
+ - `load_prequantized`: set on this wrapper when loading a mixed quantized checkpoint, not on the sub-configs.
114
+ - The layer sets matched by the sub-configs must be pairwise disjoint. `quantize_model` and `prepare_for_prequantized_load` verify this before touching the model, and raise on conflict, naming the overlapping layers.
115
+
116
+ `build_quantized_shell(module, compute_dtype, layer_name=None)` gains an extra `layer_name` parameter here: when multiple sub-configs share the same backend, the quantized Linears they produce are the same class, and ownership can only be determined by layer name.
117
+
118
+ ### describe_quant_method and QUANT_METHODS
119
+
120
+ `QUANT_METHODS` is a registry of `{method name: QuantMethodSpec}`. `QuantMethodSpec` has three fields: `backend` (backend name), `config_factory` (a callable turning `backend_config_kwargs` into the backend config), and `label` (a human-readable description).
121
+
122
+ Call `backends.load_all_backends()` before enumerating all methods:
123
+
124
+ ```python
125
+ from diffsynth.core.quant import QUANT_METHODS, backends
126
+
127
+ backends.load_all_backends()
128
+ print(sorted(QUANT_METHODS))
129
+ ```
130
+
131
+ `describe_quant_method(name)` prints a method's backend, description, and the accepted `backend_config_kwargs` with defaults (it loads the backend internally):
132
+
133
+ ```python
134
+ from diffsynth.core.quant import describe_quant_method
135
+
136
+ describe_quant_method("comfy_kitchen_int8_w8a8")
137
+ ```
138
+
139
+ ```
140
+ method: comfy_kitchen_int8_w8a8
141
+ backend: comfy_kitchen
142
+ detail: W8A8, int8 weight + int8 dynamic activation (ComfyUI int8_tensorwise)
143
+ backend config: diffsynth.core.quant.backends.comfy_kitchen.ComfyKitchenInt8Config
144
+ backend_config_kwargs (user-tunable):
145
+ per_channel = True
146
+ convrot = True
147
+ convrot_groupsize = 256
148
+ orig_dtype = torch.bfloat16
149
+ pinned by method (not overridable):
150
+ format = 'int8_tensorwise'
151
+ ```
152
+
153
+ `user-tunable` are the parameters that can be modified via `backend_config_kwargs`; `pinned by method` are fixed for the method and cannot be modified (e.g. `comfy_kitchen_int8_w8a8` and `comfy_kitchen_fp8_w8a8` share one backend and are distinguished by `format`). Passing an unaccepted key raises an error listing the available keys.
154
+
155
+ ## Extension Interface: Custom Backends
156
+
157
+ ### The QuantBackend Contract
158
+
159
+ `QuantBackend` is the adapter layer between the framework and a concrete quantization library (bitsandbytes / torchao / custom). Subclasses are registered into `QUANT_BACKENDS` via `register_quant_backend`, instantiated by `QuantizeConfig`, and injected with the method's backend config.
160
+
161
+ The quantized Linear produced by a backend must satisfy the following four contract clauses:
162
+
163
+ - **(a)** It is a drop-in replacement for `nn.Linear`: `forward(x)` performs dequantization + matmul internally.
164
+ - **(b)** `.to(...)` only moves devices, never re-types the packed weight / quant state: dtype casts (`.to(dtype)`, `.half()`, `.float()`, etc.) must leave their storage format and values intact.
165
+ - **(c)** `state_dict()` and `load_state_dict(assign=True)` round-trip (via `flatten_state_dict` / `unflatten_state_dict` when necessary).
166
+ - **(d)** (Training only) `forward` is differentiable with respect to its input, so gradients can pass through frozen quantized layers to reach LoRA branches. Declared statically by `capabilities()["is_differentiable"]` and verifiable at runtime with `check_differentiable`.
167
+
168
+ Clause (b) is necessary because VRAM management performs dtype/device conversions on the model; if a packed weight were accidentally cast to bf16, the quant state would be corrupted. See `Fp8Linear._apply` in `diffsynth/models/ideogram4_dit.py` for a reference: register the tensor names that need protection, and downgrade conversions that would change their dtype to device-only moves inside `_apply`.
169
+
170
+ Members to implement or override:
171
+
172
+ | Member | Description |
173
+ | --- | --- |
174
+ | `name` | Set automatically by `register_quant_backend` |
175
+ | `project_url` | The project page of the library this backend belongs to; `announce_environment()` prints it, pointing hardware compatibility issues upstream |
176
+ | `capabilities()` | Returns four boolean flags `is_serializable` / `is_differentiable` / `is_compileable` / `requires_calibration`, all defaulting to `False` |
177
+ | `validate_environment()` | Checks dependencies and hardware, raising an exception with installation instructions when missing. Called when constructing `QuantizeConfig` |
178
+ | `quantized_linear_classes()` | Declares the Linear classes this backend produces; they must all be subclasses of `torch.nn.Linear`. `is_quantized_linear` defaults to an `isinstance` check against them |
179
+ | `create_quantized_linear(linear, compute_device, model_device)` | Online quantization: turns an fp `nn.Linear` into a quantized Linear. If unimplemented, the backend does not support online quantization |
180
+ | `create_quantized_linear_shell(linear, compute_dtype)` | Builds an empty shell for loading pre-quantized checkpoints. If unimplemented, the backend does not support pre-quantized loading |
181
+ | `dequantize_to_linear(module, compute_dtype, compute_device, model_device)` | Restores a plain `nn.Linear`. If unimplemented, `mode="dequant_once"` is unavailable |
182
+ | `flatten_state_dict` / `unflatten_state_dict` | Conversion between quantized state dicts and flat tensors; must be implemented when `is_serializable=True` |
183
+
184
+ The base class provides clear error messages for unimplemented methods, so a backend supporting only some capabilities can implement just the ones it needs.
185
+
186
+ ### BackendConfig
187
+
188
+ `BackendConfig` is the base class for a backend's typed config. User-tunable parameters are written as ordinary dataclass fields; values pinned by the method are declared with `field(init=False, default=...)`, so they are both shown separately by `describe_quant_method` and impossible to modify via `backend_config_kwargs`.
189
+
190
+ The classmethod `from_kwargs(kwargs)` validates the keys passed in: unknown keys raise a `ValueError` listing all accepted keys. It is typically used directly as the `config_factory` of `register_quant_method`.
191
+
192
+ The bitsandbytes backend is a canonical example of this pattern — the shared 4bit parameters live in the base class, while `quant_type` is pinned by each method's subclass:
193
+
194
+ ```python
195
+ from dataclasses import dataclass, field
196
+ import torch
197
+ from diffsynth.core.quant import BackendConfig, register_quant_method
198
+
199
+
200
+ @dataclass
201
+ class BitsAndBytes4bitConfig(BackendConfig):
202
+ compress_statistics: bool = True
203
+ blocksize: int = None
204
+ quant_storage: torch.dtype = torch.uint8
205
+
206
+
207
+ @dataclass
208
+ class BitsAndBytesNF4Config(BitsAndBytes4bitConfig):
209
+ quant_type: str = field(init=False, default="nf4")
210
+
211
+
212
+ register_quant_method("bitsandbytes_nf4", "bitsandbytes", BitsAndBytesNF4Config.from_kwargs, label="4bit, nf4, weight-only")
213
+ ```
214
+
215
+ `config_factory` is not required to return a `BackendConfig`: if the backend directly consumes a third-party library's config object, you can pass any function that turns a `dict` into that object (the torchao backend does this, building `Int8WeightOnlyConfig` and the like directly).
216
+
217
+ ### register_quant_backend and register_quant_method
218
+
219
+ - `register_quant_backend(name)`: a class decorator that registers a backend class into `QUANT_BACKENDS` and sets its `name`.
220
+ - `register_quant_method(name, backend, config_factory, label="")`: registers a method name into `QUANT_METHODS`, specifying which backend it uses and how its backend config is built. One backend can register multiple methods, distinguished by pinned fields.
221
+
222
+ A complete skeleton of a minimal backend:
223
+
224
+ ```python
225
+ import torch
226
+ from diffsynth.core.quant import QuantBackend, register_quant_backend, register_quant_method
227
+
228
+
229
+ class MyQuantLinear(torch.nn.Linear):
230
+ """Custom quantized Linear; must satisfy contract clauses (a)-(d)."""
231
+
232
+
233
+ @register_quant_backend("my_backend")
234
+ class MyQuantBackend(QuantBackend):
235
+ project_url = "https://example.com/my-quant-lib"
236
+
237
+ def capabilities(self):
238
+ return {**super().capabilities(), "is_serializable": True, "is_differentiable": True}
239
+
240
+ def validate_environment(self):
241
+ ... # raise ImportError when dependencies are missing
242
+
243
+ def quantized_linear_classes(self):
244
+ return (MyQuantLinear,)
245
+
246
+ def create_quantized_linear(self, linear, compute_device=None, model_device=None):
247
+ ...
248
+
249
+ def create_quantized_linear_shell(self, linear, compute_dtype):
250
+ ...
251
+
252
+ def dequantize_to_linear(self, module, compute_dtype, compute_device=None, model_device=None):
253
+ ...
254
+
255
+
256
+ register_quant_method("my_method", "my_backend", lambda kwargs: dict(kwargs), label="my custom method")
257
+ ```
258
+
259
+ Once registered, it can be used just like a built-in method: `QuantizeConfig(method="my_method")`. If the backend is defined outside `diffsynth/core/quant/backends/` (e.g. alongside a model), it only needs to be imported before constructing `QuantizeConfig`.
260
+
261
+ ## Verification Tools
262
+
263
+ ### check_differentiable
264
+
265
+ ```python
266
+ check_differentiable(module, example_input=None, verbose=True) -> bool
267
+ ```
268
+
269
+ Checks whether gradients can pass through `module` to its input: runs a real backward pass from the output (`torch.autograd.grad`) and confirms a finite gradient arrives at the input. This is exactly what LoRA training requires from frozen (quantized) layers. The module is cast to bfloat16 in place and probed with a bfloat16 input; when `example_input` is `None`, a random input is constructed automatically for modules exposing `in_features`.
270
+
271
+ ```python
272
+ import torch
273
+ from diffsynth.core.quant import check_differentiable
274
+ from torchao.quantization import quantize_, Int8WeightOnlyConfig
275
+
276
+ linear = torch.nn.Linear(1024, 1024, dtype=torch.bfloat16, device="cuda")
277
+ quantize_(linear, Int8WeightOnlyConfig(version=2))
278
+ check_differentiable(linear)
279
+ ```
280
+
281
+ ### check_backend_contract
282
+
283
+ ```python
284
+ check_backend_contract(backend, in_features=512, out_features=512,
285
+ compute_dtype=torch.bfloat16, compute_device="cuda", verbose=True) -> bool
286
+ ```
287
+
288
+ An admission self-check for new backends: verifies that it declares its Linear classes, that both factory methods return instances of those classes, and that every declared class is a subclass of `torch.nn.Linear` (otherwise LoRA target detection and VRAM management cannot see it). It also checks that the checkpoint keys the backend actually writes all live under the layer name — a key pattern missing a scale would make Disk Offload silently load corrupted layers. Unsupported factory methods are skipped rather than counted as failures.
289
+
290
+ ```python
291
+ from diffsynth.core.quant import QUANT_BACKENDS, QUANT_METHODS, check_backend_contract
292
+
293
+ spec = QUANT_METHODS["bitsandbytes_nf4"]
294
+ check_backend_contract(QUANT_BACKENDS[spec.backend](spec.config_factory({})))
295
+ ```
docs/en/API_Reference/core/vram.md ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # `diffsynth.core.vram`: VRAM Management
2
+
3
+ This document introduces the underlying VRAM management functionalities in `diffsynth.core.vram`. If you wish to use these functionalities in other codebases, you can refer to this document.
4
+
5
+ ## Skipping Model Parameter Initialization
6
+
7
+ When loading models in `PyTorch`, model parameters default to occupying VRAM or memory and initializing parameters, but these parameters will be overwritten when loading pretrained weights, leading to redundant computations. `PyTorch` does not provide an interface to skip these redundant computations. We provide `skip_model_initialization` in `diffsynth.core.vram` to skip model parameter initialization.
8
+
9
+ Default model loading approach:
10
+
11
+ ```python
12
+ from diffsynth.core import load_state_dict
13
+ from diffsynth.models.qwen_image_controlnet import QwenImageBlockWiseControlNet
14
+
15
+ model = QwenImageBlockWiseControlNet() # Slow
16
+ path = "models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors"
17
+ state_dict = load_state_dict(path, device="cpu")
18
+ model.load_state_dict(state_dict, assign=True)
19
+ ```
20
+
21
+ Model loading approach that skips parameter initialization:
22
+
23
+ ```python
24
+ from diffsynth.core import load_state_dict, skip_model_initialization
25
+ from diffsynth.models.qwen_image_controlnet import QwenImageBlockWiseControlNet
26
+
27
+ with skip_model_initialization():
28
+ model = QwenImageBlockWiseControlNet() # Fast
29
+ path = "models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors"
30
+ state_dict = load_state_dict(path, device="cpu")
31
+ model.load_state_dict(state_dict, assign=True)
32
+ ```
33
+
34
+ In `DiffSynth-Studio`, all pretrained models follow this loading logic. After developers [integrate models](../../Developer_Guide/Integrating_Your_Model.md), they can directly load models quickly using this approach.
35
+
36
+ ## State Dict Disk Mapping
37
+
38
+ For pretrained weight files of a model, if we only need to read a set of parameters rather than all parameters, State Dict Disk Mapping can accelerate this process. We provide `DiskMap` in `diffsynth.core.vram` for on-demand loading of model parameters.
39
+
40
+ Default weight loading approach:
41
+
42
+ ```python
43
+ from diffsynth.core import load_state_dict
44
+
45
+ path = "models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors"
46
+ state_dict = load_state_dict(path, device="cpu") # Slow
47
+ print(state_dict["img_in.weight"])
48
+ ```
49
+
50
+ Using `DiskMap` to load only specific parameters:
51
+
52
+ ```python
53
+ from diffsynth.core import DiskMap
54
+
55
+ path = "models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny/model.safetensors"
56
+ state_dict = DiskMap(path, device="cpu") # Fast
57
+ print(state_dict["img_in.weight"])
58
+ ```
59
+
60
+ `DiskMap` is the basic component of Disk Offload in `DiffSynth-Studio`. After developers [configure fine-grained VRAM management schemes](../../Developer_Guide/Enabling_VRAM_management.md), they can directly enable Disk Offload.
61
+
62
+ `DiskMap` is a functionality implemented using the characteristics of `.safetensors` files. Therefore, when using `.bin`, `.pth`, `.ckpt`, and other binary files, model parameters are fully loaded, which causes Disk Offload to not support these formats of files. **We do not recommend developers to continue using these formats of files.**
63
+
64
+ ## Replacable Modules for VRAM Management
65
+
66
+ When `DiffSynth-Studio`'s VRAM management is enabled, the modules inside the model will be replaced with replacable modules in `diffsynth.core.vram.layers`. For usage, see [Fine-grained VRAM Management Scheme](../../Developer_Guide/Enabling_VRAM_management.md#writing-fine-grained-vram-management-schemes).
docs/en/Developer_Guide/Building_a_Pipeline.md ADDED
@@ -0,0 +1,254 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Building a Pipeline
2
+
3
+ After [integrating the required models for the Pipeline](../Developer_Guide/Integrating_Your_Model.md), you also need to build a `Pipeline` for model inference. This document provides a standardized process for building a `Pipeline`. Developers can also refer to existing `Pipeline` implementations for construction.
4
+
5
+ The `Pipeline` implementation is located in `diffsynth/pipelines`. Each `Pipeline` contains the following essential key components:
6
+
7
+ * `__init__`
8
+ * `from_pretrained`
9
+ * `__call__`
10
+ * `units`
11
+ * `model_fn`
12
+
13
+ ## `__init__`
14
+
15
+ In `__init__`, the `Pipeline` is initialized. Here is a simple implementation:
16
+
17
+ ```python
18
+ import torch
19
+ from PIL import Image
20
+ from typing import Union
21
+ from tqdm import tqdm
22
+ from ..diffusion import FlowMatchScheduler
23
+ from ..core import ModelConfig
24
+ from ..diffusion.base_pipeline import BasePipeline, PipelineUnit
25
+ from ..models.new_models import XXX_Model, YYY_Model, ZZZ_Model
26
+
27
+ class NewDiffSynthPipeline(BasePipeline):
28
+
29
+ def __init__(self, device="cuda", torch_dtype=torch.bfloat16):
30
+ super().__init__(device=device, torch_dtype=torch_dtype)
31
+ self.scheduler = FlowMatchScheduler()
32
+ self.text_encoder: XXX_Model = None
33
+ self.dit: YYY_Model = None
34
+ self.vae: ZZZ_Model = None
35
+ self.in_iteration_models = ("dit",)
36
+ self.units = [
37
+ NewDiffSynthPipelineUnit_xxx(),
38
+ ...
39
+ ]
40
+ self.model_fn = model_fn_new
41
+ ```
42
+
43
+ This includes the following parts:
44
+
45
+ * `scheduler`: Scheduler, used to control the coefficients in the iterative formula during inference, controlling the noise content at each step.
46
+ * `text_encoder`, `dit`, `vae`: Models. Since [Latent Diffusion](https://arxiv.org/abs/2112.10752) was proposed, this three-stage model architecture has become the mainstream Diffusion model architecture. However, this is not immutable, and any number of models can be added to the `Pipeline`.
47
+ * `in_iteration_models`: Iteration models. This tuple marks which models will be called during iteration.
48
+ * `units`: Pre-processing units for model iteration. See [`units`](#units) for details.
49
+ * `model_fn`: The `forward` function of the denoising model during iteration. See [`model_fn`](#model_fn) for details.
50
+
51
+ > Q: Model loading does not occur in `__init__`, why initialize each model as `None` here?
52
+ >
53
+ > A: By annotating the type of each model here, the code editor can provide code completion prompts based on each model, facilitating subsequent development.
54
+
55
+ ## `from_pretrained`
56
+
57
+ `from_pretrained` is responsible for loading the required models to make the `Pipeline` callable. Here is a simple implementation:
58
+
59
+ ```python
60
+ @staticmethod
61
+ def from_pretrained(
62
+ torch_dtype: torch.dtype = torch.bfloat16,
63
+ device: Union[str, torch.device] = "cuda",
64
+ model_configs: list[ModelConfig] = [],
65
+ vram_limit: float = None,
66
+ ):
67
+ # Initialize pipeline
68
+ pipe = NewDiffSynthPipeline(device=device, torch_dtype=torch_dtype)
69
+ model_pool = pipe.download_and_load_models(model_configs, vram_limit)
70
+
71
+ # Fetch models
72
+ pipe.text_encoder = model_pool.fetch_model("xxx_text_encoder")
73
+ pipe.dit = model_pool.fetch_model("yyy_dit")
74
+ pipe.vae = model_pool.fetch_model("zzz_vae")
75
+ # If necessary, load tokenizers here.
76
+
77
+ # VRAM Management
78
+ pipe.vram_management_enabled = pipe.check_vram_management_state()
79
+ return pipe
80
+ ```
81
+
82
+ Developers need to implement the logic for fetching models. The corresponding model names are the `"model_name"` in the [model Config filled in during model integration](../Developer_Guide/Integrating_Your_Model.md#step-3-writing-model-config).
83
+
84
+ Some models also need to load `tokenizer`. Extra `tokenizer_config` parameters can be added to `from_pretrained` as needed, and this part can be implemented after fetching the models.
85
+
86
+ ## `__call__`
87
+
88
+ `__call__` implements the entire generation process of the Pipeline. Below is a common generation process template. Developers can modify it based on their needs.
89
+
90
+ ```python
91
+ @torch.no_grad()
92
+ def __call__(
93
+ self,
94
+ prompt: str,
95
+ negative_prompt: str = "",
96
+ cfg_scale: float = 4.0,
97
+ input_image: Image.Image = None,
98
+ denoising_strength: float = 1.0,
99
+ height: int = 1328,
100
+ width: int = 1328,
101
+ seed: int = None,
102
+ rand_device: str = "cpu",
103
+ num_inference_steps: int = 30,
104
+ progress_bar_cmd = tqdm,
105
+ ):
106
+ # Scheduler
107
+ self.scheduler.set_timesteps(
108
+ num_inference_steps,
109
+ denoising_strength=denoising_strength
110
+ )
111
+
112
+ # Parameters
113
+ inputs_posi = {
114
+ "prompt": prompt,
115
+ }
116
+ inputs_nega = {
117
+ "negative_prompt": negative_prompt,
118
+ }
119
+ inputs_shared = {
120
+ "cfg_scale": cfg_scale,
121
+ "input_image": input_image,
122
+ "denoising_strength": denoising_strength,
123
+ "height": height,
124
+ "width": width,
125
+ "seed": seed,
126
+ "rand_device": rand_device,
127
+ "num_inference_steps": num_inference_steps,
128
+ }
129
+ for unit in self.units:
130
+ inputs_shared, inputs_posi, inputs_nega = self.unit_runner(unit, self, inputs_shared, inputs_posi, inputs_nega)
131
+
132
+ # Denoise
133
+ self.load_models_to_device(self.in_iteration_models)
134
+ models = {name: getattr(self, name) for name in self.in_iteration_models}
135
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
136
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
137
+
138
+ # Inference
139
+ noise_pred_posi = self.model_fn(**models, **inputs_shared, **inputs_posi, timestep=timestep, progress_id=progress_id)
140
+ if cfg_scale != 1.0:
141
+ noise_pred_nega = self.model_fn(**models, **inputs_shared, **inputs_nega, timestep=timestep, progress_id=progress_id)
142
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
143
+ else:
144
+ noise_pred = noise_pred_posi
145
+
146
+ # Scheduler
147
+ inputs_shared["latents"] = self.step(self.scheduler, progress_id=progress_id, noise_pred=noise_pred, **inputs_shared)
148
+
149
+ # Decode
150
+ self.load_models_to_device(['vae'])
151
+ image = self.vae.decode(inputs_shared["latents"], device=self.device)
152
+ image = self.vae_output_to_image(image)
153
+ self.load_models_to_device([])
154
+
155
+ return image
156
+ ```
157
+
158
+ ## `units`
159
+
160
+ `units` contains all the preprocessing processes, such as: width/height checking, prompt encoding, initial noise generation, etc. In the entire model preprocessing process, data is abstracted into three mutually exclusive parts, stored in corresponding dictionaries:
161
+
162
+ * `inputs_shared`: Shared inputs, parameters unrelated to [Classifier-Free Guidance](https://arxiv.org/abs/2207.12598) (CFG for short).
163
+ * `inputs_posi`: Positive side inputs for Classifier-Free Guidance, containing content related to positive prompts.
164
+ * `inputs_nega`: Negative side inputs for Classifier-Free Guidance, containing content related to negative prompts.
165
+
166
+ Pipeline Unit implementations include three types: direct mode, CFG separation mode, and takeover mode.
167
+
168
+ If some calculations are unrelated to CFG, direct mode can be used, for example, Qwen-Image's random noise initialization:
169
+
170
+ ```python
171
+ class QwenImageUnit_NoiseInitializer(PipelineUnit):
172
+ def __init__(self):
173
+ super().__init__(
174
+ input_params=("height", "width", "seed", "rand_device"),
175
+ output_params=("noise",),
176
+ )
177
+
178
+ def process(self, pipe: QwenImagePipeline, height, width, seed, rand_device):
179
+ noise = pipe.generate_noise((1, 16, height//8, width//8), seed=seed, rand_device=rand_device, rand_torch_dtype=pipe.torch_dtype)
180
+ return {"noise": noise}
181
+ ```
182
+
183
+ If some calculations are related to CFG and need to separately process positive and negative prompts, but the input parameters on both sides are the same, CFG separation mode can be used, for example, Qwen-image's prompt encoding:
184
+
185
+ ```python
186
+ class QwenImageUnit_PromptEmbedder(PipelineUnit):
187
+ def __init__(self):
188
+ super().__init__(
189
+ seperate_cfg=True,
190
+ input_params_posi={"prompt": "prompt"},
191
+ input_params_nega={"prompt": "negative_prompt"},
192
+ input_params=("edit_image",),
193
+ output_params=("prompt_emb", "prompt_emb_mask"),
194
+ onload_model_names=("text_encoder",)
195
+ )
196
+
197
+ def process(self, pipe: QwenImagePipeline, prompt, edit_image=None) -> dict:
198
+ pipe.load_models_to_device(self.onload_model_names)
199
+ # Do something
200
+ return {"prompt_emb": prompt_embeds, "prompt_emb_mask": encoder_attention_mask}
201
+ ```
202
+
203
+ If some calculations need global information, takeover mode is required, for example, Qwen-Image's entity partition control:
204
+
205
+ ```python
206
+ class QwenImageUnit_EntityControl(PipelineUnit):
207
+ def __init__(self):
208
+ super().__init__(
209
+ take_over=True,
210
+ input_params=("eligen_entity_prompts", "width", "height", "eligen_enable_on_negative", "cfg_scale"),
211
+ output_params=("entity_prompt_emb", "entity_masks", "entity_prompt_emb_mask"),
212
+ onload_model_names=("text_encoder",)
213
+ )
214
+
215
+ def process(self, pipe: QwenImagePipeline, inputs_shared, inputs_posi, inputs_nega):
216
+ # Do something
217
+ return inputs_shared, inputs_posi, inputs_nega
218
+ ```
219
+
220
+ The following are the parameter configurations required for Pipeline Unit:
221
+
222
+ * `seperate_cfg`: Whether to enable CFG separation mode
223
+ * `take_over`: Whether to enable takeover mode
224
+ * `input_params`: Shared input parameters
225
+ * `output_params`: Output parameters
226
+ * `input_params_posi`: Positive side input parameters
227
+ * `input_params_nega`: Negative side input parameters
228
+ * `onload_model_names`: Names of model components to be called
229
+
230
+ When designing `unit`, please try to follow these principles:
231
+
232
+ * Default fallback: For optional function `unit` input parameters, the default is `None` rather than `False` or other values. Please provide fallback processing for this default value.
233
+ * Parameter triggering: Some Adapter models may not be loaded, such as ControlNet. The corresponding `unit` should control triggering based on whether the parameter input is `None` rather than whether the model is loaded. For example, when the user inputs `controlnet_image` but does not load the ControlNet model, the code should give an error rather than ignore these input parameters and continue execution.
234
+ * Simplicity first: Use direct mode as much as possible, only use takeover mode when the function cannot be implemented.
235
+ * VRAM efficiency: When calling models in `unit`, please use `pipe.load_models_to_device(self.onload_model_names)` to activate the corresponding models. Do not call other models outside `onload_model_names`. After `unit` calculation is completed, do not manually release VRAM with `pipe.load_models_to_device([])`.
236
+
237
+ > Q: Some parameters are not called during the inference process, such as `output_params`. Is it still necessary to configure them?
238
+ >
239
+ > A: These parameters will not affect the inference process, but they will affect some experimental features. Therefore, we recommend configuring them properly. For example, "split training" - we can complete the preprocessing offline during training, but some model calculations that require gradient backpropagation cannot be split. These parameters are used to build computational graphs to infer which calculations can be split.
240
+
241
+ ## `model_fn`
242
+
243
+ `model_fn` is the unified `forward` interface during iteration. For models where the open-source ecosystem is not yet formed, you can directly use the denoising model's `forward`, for example:
244
+
245
+ ```python
246
+ def model_fn_new(dit=None, latents=None, timestep=None, prompt_emb=None, **kwargs):
247
+ return dit(latents, prompt_emb, timestep)
248
+ ```
249
+
250
+ For models with rich open-source ecosystems, `model_fn` usually contains complex and chaotic cross-model inference. Taking `diffsynth/pipelines/qwen_image.py` as an example, the additional calculations implemented in this function include: entity partition control, three types of ControlNet, Gradient Checkpointing, etc. Developers need to be extra careful when implementing this part to avoid conflicts between module functions.
251
+
252
+ ## Compilation Acceleration
253
+
254
+ To enable compilation acceleration, please refer to [Inference Acceleration](../Pipeline_Usage/Accelerated_Inference.md).
docs/en/Developer_Guide/Enabling_VRAM_management.md ADDED
@@ -0,0 +1,455 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Fine-Grained VRAM Management Scheme
2
+
3
+ This document introduces how to write reasonable fine-grained VRAM management schemes for models, and how to use the VRAM management functions in `DiffSynth-Studio` for other external code libraries. Before reading this document, please read the document [VRAM Management](../Pipeline_Usage/VRAM_management.md).
4
+
5
+ ## How Much VRAM Does a 20B Model Need?
6
+
7
+ Taking Qwen-Image's DiT model as an example, this model has reached 20B parameters. The following code will load this model and perform inference, requiring about 40G VRAM. This model obviously cannot run on consumer-grade GPUs with smaller VRAM.
8
+
9
+ ```python
10
+ from diffsynth.core import load_model
11
+ from diffsynth.models.qwen_image_dit import QwenImageDiT
12
+ from modelscope import snapshot_download
13
+ import torch
14
+
15
+ snapshot_download(
16
+ model_id="Qwen/Qwen-Image",
17
+ local_dir="models/Qwen/Qwen-Image",
18
+ allow_file_pattern="transformer/*"
19
+ )
20
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
21
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
22
+ inputs = {
23
+ "latents": torch.randn((1, 16, 128, 128), dtype=torch.bfloat16, device="cuda"),
24
+ "timestep": torch.zeros((1,), dtype=torch.bfloat16, device="cuda"),
25
+ "prompt_emb": torch.randn((1, 5, 3584), dtype=torch.bfloat16, device="cuda"),
26
+ "prompt_emb_mask": torch.ones((1, 5), dtype=torch.int64, device="cuda"),
27
+ "height": 1024,
28
+ "width": 1024,
29
+ }
30
+
31
+ model = load_model(QwenImageDiT, model_path, torch_dtype=torch.bfloat16, device="cuda")
32
+ with torch.no_grad():
33
+ output = model(**inputs)
34
+ ```
35
+
36
+ ## Writing Fine-Grained VRAM Management Scheme
37
+
38
+ To write a fine-grained VRAM management scheme, we need to use `print(model)` to observe and analyze the model structure:
39
+
40
+ ```
41
+ QwenImageDiT(
42
+ (pos_embed): QwenEmbedRope()
43
+ (time_text_embed): TimestepEmbeddings(
44
+ (time_proj): TemporalTimesteps()
45
+ (timestep_embedder): DiffusersCompatibleTimestepProj(
46
+ (linear_1): Linear(in_features=256, out_features=3072, bias=True)
47
+ (act): SiLU()
48
+ (linear_2): Linear(in_features=3072, out_features=3072, bias=True)
49
+ )
50
+ )
51
+ (txt_norm): RMSNorm()
52
+ (img_in): Linear(in_features=64, out_features=3072, bias=True)
53
+ (txt_in): Linear(in_features=3584, out_features=3072, bias=True)
54
+ (transformer_blocks): ModuleList(
55
+ (0-59): 60 x QwenImageTransformerBlock(
56
+ (img_mod): Sequential(
57
+ (0): SiLU()
58
+ (1): Linear(in_features=3072, out_features=18432, bias=True)
59
+ )
60
+ (img_norm1): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
61
+ (attn): QwenDoubleStreamAttention(
62
+ (to_q): Linear(in_features=3072, out_features=3072, bias=True)
63
+ (to_k): Linear(in_features=3072, out_features=3072, bias=True)
64
+ (to_v): Linear(in_features=3072, out_features=3072, bias=True)
65
+ (norm_q): RMSNorm()
66
+ (norm_k): RMSNorm()
67
+ (add_q_proj): Linear(in_features=3072, out_features=3072, bias=True)
68
+ (add_k_proj): Linear(in_features=3072, out_features=3072, bias=True)
69
+ (add_v_proj): Linear(in_features=3072, out_features=3072, bias=True)
70
+ (norm_added_q): RMSNorm()
71
+ (norm_added_k): RMSNorm()
72
+ (to_out): Sequential(
73
+ (0): Linear(in_features=3072, out_features=3072, bias=True)
74
+ )
75
+ (to_add_out): Linear(in_features=3072, out_features=3072, bias=True)
76
+ )
77
+ (img_norm2): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
78
+ (img_mlp): QwenFeedForward(
79
+ (net): ModuleList(
80
+ (0): ApproximateGELU(
81
+ (proj): Linear(in_features=3072, out_features=12288, bias=True)
82
+ )
83
+ (1): Dropout(p=0.0, inplace=False)
84
+ (2): Linear(in_features=12288, out_features=3072, bias=True)
85
+ )
86
+ )
87
+ (txt_mod): Sequential(
88
+ (0): SiLU()
89
+ (1): Linear(in_features=3072, out_features=18432, bias=True)
90
+ )
91
+ (txt_norm1): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
92
+ (txt_norm2): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
93
+ (txt_mlp): QwenFeedForward(
94
+ (net): ModuleList(
95
+ (0): ApproximateGELU(
96
+ (proj): Linear(in_features=3072, out_features=12288, bias=True)
97
+ )
98
+ (1): Dropout(p=0.0, inplace=False)
99
+ (2): Linear(in_features=12288, out_features=3072, bias=True)
100
+ )
101
+ )
102
+ )
103
+ )
104
+ (norm_out): AdaLayerNorm(
105
+ (linear): Linear(in_features=3072, out_features=6144, bias=True)
106
+ (norm): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
107
+ )
108
+ (proj_out): Linear(in_features=3072, out_features=64, bias=True)
109
+ )
110
+ ```
111
+
112
+ In VRAM management, we only care about layers containing parameters. In this model structure, `QwenEmbedRope`, `TemporalTimesteps`, `SiLU` and other Layers do not contain parameters. `LayerNorm` also does not contain parameters because `elementwise_affine=False` is set. Layers containing parameters are only `Linear` and `RMSNorm`.
113
+
114
+ `diffsynth.core.vram` provides two replacement modules for VRAM management:
115
+ * `AutoWrappedLinear`: Used to replace `Linear` layers
116
+ * `AutoWrappedModule`: Used to replace any other layer
117
+
118
+ Write a `module_map` to map `Linear` and `RMSNorm` in the model to the corresponding modules:
119
+
120
+ ```python
121
+ module_map={
122
+ torch.nn.Linear: AutoWrappedLinear,
123
+ RMSNorm: AutoWrappedModule,
124
+ }
125
+ ```
126
+
127
+ In addition, `vram_config` and `vram_limit` are also required, which have been introduced in [VRAM Management](../Pipeline_Usage/VRAM_management.md#more-usage-methods).
128
+
129
+ Call `enable_vram_management` to enable VRAM management. Note that the `device` when loading the model is `cpu`, consistent with `offload_device`:
130
+
131
+ ```python
132
+ from diffsynth.core import load_model, enable_vram_management, AutoWrappedLinear, AutoWrappedModule
133
+ from diffsynth.models.qwen_image_dit import QwenImageDiT, RMSNorm
134
+ import torch
135
+
136
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
137
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
138
+ inputs = {
139
+ "latents": torch.randn((1, 16, 128, 128), dtype=torch.bfloat16, device="cuda"),
140
+ "timestep": torch.zeros((1,), dtype=torch.bfloat16, device="cuda"),
141
+ "prompt_emb": torch.randn((1, 5, 3584), dtype=torch.bfloat16, device="cuda"),
142
+ "prompt_emb_mask": torch.ones((1, 5), dtype=torch.int64, device="cuda"),
143
+ "height": 1024,
144
+ "width": 1024,
145
+ }
146
+
147
+ model = load_model(QwenImageDiT, model_path, torch_dtype=torch.bfloat16, device="cpu")
148
+ enable_vram_management(
149
+ model,
150
+ module_map={
151
+ torch.nn.Linear: AutoWrappedLinear,
152
+ RMSNorm: AutoWrappedModule,
153
+ },
154
+ vram_config = {
155
+ "offload_dtype": torch.bfloat16,
156
+ "offload_device": "cpu",
157
+ "onload_dtype": torch.bfloat16,
158
+ "onload_device": "cpu",
159
+ "preparing_dtype": torch.bfloat16,
160
+ "preparing_device": "cuda",
161
+ "computation_dtype": torch.bfloat16,
162
+ "computation_device": "cuda",
163
+ },
164
+ vram_limit=0,
165
+ )
166
+ with torch.no_grad():
167
+ output = model(**inputs)
168
+ ```
169
+
170
+ The above code only requires 2G VRAM to run the `forward` of a 20B model.
171
+
172
+ ## Disk Offload
173
+
174
+ [Disk Offload](../Pipeline_Usage/VRAM_management.md#disk-offload) is a special VRAM management scheme that needs to be enabled during the model loading process, not after the model is loaded. Usually, when the above code can run smoothly, Disk Offload can be directly enabled:
175
+
176
+ ```python
177
+ from diffsynth.core import load_model, enable_vram_management, AutoWrappedLinear, AutoWrappedModule
178
+ from diffsynth.models.qwen_image_dit import QwenImageDiT, RMSNorm
179
+ import torch
180
+
181
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
182
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
183
+ inputs = {
184
+ "latents": torch.randn((1, 16, 128, 128), dtype=torch.bfloat16, device="cuda"),
185
+ "timestep": torch.zeros((1,), dtype=torch.bfloat16, device="cuda"),
186
+ "prompt_emb": torch.randn((1, 5, 3584), dtype=torch.bfloat16, device="cuda"),
187
+ "prompt_emb_mask": torch.ones((1, 5), dtype=torch.int64, device="cuda"),
188
+ "height": 1024,
189
+ "width": 1024,
190
+ }
191
+
192
+ model = load_model(
193
+ QwenImageDiT,
194
+ model_path,
195
+ module_map={
196
+ torch.nn.Linear: AutoWrappedLinear,
197
+ RMSNorm: AutoWrappedModule,
198
+ },
199
+ vram_config={
200
+ "offload_dtype": "disk",
201
+ "offload_device": "disk",
202
+ "onload_dtype": "disk",
203
+ "onload_device": "disk",
204
+ "preparing_dtype": torch.bfloat16,
205
+ "preparing_device": "cuda",
206
+ "computation_dtype": torch.bfloat16,
207
+ "computation_device": "cuda",
208
+ },
209
+ vram_limit=0,
210
+ )
211
+ with torch.no_grad():
212
+ output = model(**inputs)
213
+ ```
214
+
215
+ Disk Offload is an extremely special VRAM management scheme. It only supports `.safetensors` format files, not binary files such as `.bin`, `.pth`, `.ckpt`, and does not support [state dict converter](../Developer_Guide/Integrating_Your_Model.md#step-2-model-file-format-conversion) with Tensor reshape.
216
+
217
+ If there are situations where Disk Offload cannot run normally but non-Disk Offload can run normally, please submit an issue to us on GitHub.
218
+
219
+ ## Writing Default Configuration
220
+
221
+ To make it easier for users to use the VRAM management function, we write the fine-grained VRAM management configuration in `diffsynth/configs/vram_management_module_maps.py`. The configuration information for the above model is:
222
+
223
+ ```python
224
+ "diffsynth.models.qwen_image_dit.QwenImageDiT": {
225
+ "diffsynth.models.qwen_image_dit.RMSNorm": "diffsynth.core.vram.layers.AutoWrappedModule",
226
+ "torch.nn.Linear": "diffsynth.core.vram.layers.AutoWrappedLinear",
227
+ }
228
+ ```# Fine-Grained VRAM Management Scheme
229
+
230
+ This document introduces how to write reasonable fine-grained VRAM management schemes for models, and how to use the VRAM management functions in `DiffSynth-Studio` for other external code libraries. Before reading this document, please read the document [VRAM Management](../Pipeline_Usage/VRAM_management.md).
231
+
232
+ ## How Much VRAM Does a 20B Model Need?
233
+
234
+ Taking Qwen-Image's DiT model as an example, this model has reached 20B parameters. The following code will load this model and perform inference, requiring about 40G VRAM. This model obviously cannot run on consumer-grade GPUs with smaller VRAM.
235
+
236
+ ```python
237
+ from diffsynth.core import load_model
238
+ from diffsynth.models.qwen_image_dit import QwenImageDiT
239
+ from modelscope import snapshot_download
240
+ import torch
241
+
242
+ snapshot_download(
243
+ model_id="Qwen/Qwen-Image",
244
+ local_dir="models/Qwen/Qwen-Image",
245
+ allow_file_pattern="transformer/*"
246
+ )
247
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
248
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
249
+ inputs = {
250
+ "latents": torch.randn((1, 16, 128, 128), dtype=torch.bfloat16, device="cuda"),
251
+ "timestep": torch.zeros((1,), dtype=torch.bfloat16, device="cuda"),
252
+ "prompt_emb": torch.randn((1, 5, 3584), dtype=torch.bfloat16, device="cuda"),
253
+ "prompt_emb_mask": torch.ones((1, 5), dtype=torch.int64, device="cuda"),
254
+ "height": 1024,
255
+ "width": 1024,
256
+ }
257
+
258
+ model = load_model(QwenImageDiT, model_path, torch_dtype=torch.bfloat16, device="cuda")
259
+ with torch.no_grad():
260
+ output = model(**inputs)
261
+ ```
262
+
263
+ ## Writing Fine-Grained VRAM Management Scheme
264
+
265
+ To write a fine-grained VRAM management scheme, we need to use `print(model)` to observe and analyze the model structure:
266
+
267
+ ```
268
+ QwenImageDiT(
269
+ (pos_embed): QwenEmbedRope()
270
+ (time_text_embed): TimestepEmbeddings(
271
+ (time_proj): TemporalTimesteps()
272
+ (timestep_embedder): DiffusersCompatibleTimestepProj(
273
+ (linear_1): Linear(in_features=256, out_features=3072, bias=True)
274
+ (act): SiLU()
275
+ (linear_2): Linear(in_features=3072, out_features=3072, bias=True)
276
+ )
277
+ )
278
+ (txt_norm): RMSNorm()
279
+ (img_in): Linear(in_features=64, out_features=3072, bias=True)
280
+ (txt_in): Linear(in_features=3584, out_features=3072, bias=True)
281
+ (transformer_blocks): ModuleList(
282
+ (0-59): 60 x QwenImageTransformerBlock(
283
+ (img_mod): Sequential(
284
+ (0): SiLU()
285
+ (1): Linear(in_features=3072, out_features=18432, bias=True)
286
+ )
287
+ (img_norm1): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
288
+ (attn): QwenDoubleStreamAttention(
289
+ (to_q): Linear(in_features=3072, out_features=3072, bias=True)
290
+ (to_k): Linear(in_features=3072, out_features=3072, bias=True)
291
+ (to_v): Linear(in_features=3072, out_features=3072, bias=True)
292
+ (norm_q): RMSNorm()
293
+ (norm_k): RMSNorm()
294
+ (add_q_proj): Linear(in_features=3072, out_features=3072, bias=True)
295
+ (add_k_proj): Linear(in_features=3072, out_features=3072, bias=True)
296
+ (add_v_proj): Linear(in_features=3072, out_features=3072, bias=True)
297
+ (norm_added_q): RMSNorm()
298
+ (norm_added_k): RMSNorm()
299
+ (to_out): Sequential(
300
+ (0): Linear(in_features=3072, out_features=3072, bias=True)
301
+ )
302
+ (to_add_out): Linear(in_features=3072, out_features=3072, bias=True)
303
+ )
304
+ (img_norm2): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
305
+ (img_mlp): QwenFeedForward(
306
+ (net): ModuleList(
307
+ (0): ApproximateGELU(
308
+ (proj): Linear(in_features=3072, out_features=12288, bias=True)
309
+ )
310
+ (1): Dropout(p=0.0, inplace=False)
311
+ (2): Linear(in_features=12288, out_features=3072, bias=True)
312
+ )
313
+ )
314
+ (txt_mod): Sequential(
315
+ (0): SiLU()
316
+ (1): Linear(in_features=3072, out_features=18432, bias=True)
317
+ )
318
+ (txt_norm1): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
319
+ (txt_norm2): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
320
+ (txt_mlp): QwenFeedForward(
321
+ (net): ModuleList(
322
+ (0): ApproximateGELU(
323
+ (proj): Linear(in_features=3072, out_features=12288, bias=True)
324
+ )
325
+ (1): Dropout(p=0.0, inplace=False)
326
+ (2): Linear(in_features=12288, out_features=3072, bias=True)
327
+ )
328
+ )
329
+ )
330
+ )
331
+ (norm_out): AdaLayerNorm(
332
+ (linear): Linear(in_features=3072, out_features=6144, bias=True)
333
+ (norm): LayerNorm((3072,), eps=1e-06, elementwise_affine=False)
334
+ )
335
+ (proj_out): Linear(in_features=3072, out_features=64, bias=True)
336
+ )
337
+ ```
338
+
339
+ In VRAM management, we only care about layers containing parameters. In this model structure, `QwenEmbedRope`, `TemporalTimesteps`, `SiLU` and other Layers do not contain parameters. `LayerNorm` also does not contain parameters because `elementwise_affine=False` is set. Layers containing parameters are only `Linear` and `RMSNorm`.
340
+
341
+ `diffsynth.core.vram` provides two replacement modules for VRAM management:
342
+ * `AutoWrappedLinear`: Used to replace `Linear` layers
343
+ * `AutoWrappedModule`: Used to replace any other layer
344
+
345
+ Write a `module_map` to map `Linear` and `RMSNorm` in the model to the corresponding modules:
346
+
347
+ ```python
348
+ module_map={
349
+ torch.nn.Linear: AutoWrappedLinear,
350
+ RMSNorm: AutoWrappedModule,
351
+ }
352
+ ```
353
+
354
+ In addition, `vram_config` and `vram_limit` are also required, which have been introduced in [VRAM Management](../Pipeline_Usage/VRAM_management.md#more-usage-methods).
355
+
356
+ Call `enable_vram_management` to enable VRAM management. Note that the `device` when loading the model is `cpu`, consistent with `offload_device`:
357
+
358
+ ```python
359
+ from diffsynth.core import load_model, enable_vram_management, AutoWrappedLinear, AutoWrappedModule
360
+ from diffsynth.models.qwen_image_dit import QwenImageDiT, RMSNorm
361
+ import torch
362
+
363
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
364
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
365
+ inputs = {
366
+ "latents": torch.randn((1, 16, 128, 128), dtype=torch.bfloat16, device="cuda"),
367
+ "timestep": torch.zeros((1,), dtype=torch.bfloat16, device="cuda"),
368
+ "prompt_emb": torch.randn((1, 5, 3584), dtype=torch.bfloat16, device="cuda"),
369
+ "prompt_emb_mask": torch.ones((1, 5), dtype=torch.int64, device="cuda"),
370
+ "height": 1024,
371
+ "width": 1024,
372
+ }
373
+
374
+ model = load_model(QwenImageDiT, model_path, torch_dtype=torch.bfloat16, device="cpu")
375
+ enable_vram_management(
376
+ model,
377
+ module_map={
378
+ torch.nn.Linear: AutoWrappedLinear,
379
+ RMSNorm: AutoWrappedModule,
380
+ },
381
+ vram_config = {
382
+ "offload_dtype": torch.bfloat16,
383
+ "offload_device": "cpu",
384
+ "onload_dtype": torch.bfloat16,
385
+ "onload_device": "cpu",
386
+ "preparing_dtype": torch.bfloat16,
387
+ "preparing_device": "cuda",
388
+ "computation_dtype": torch.bfloat16,
389
+ "computation_device": "cuda",
390
+ },
391
+ vram_limit=0,
392
+ )
393
+ with torch.no_grad():
394
+ output = model(**inputs)
395
+ ```
396
+
397
+ The above code only requires 2G VRAM to run the `forward` of a 20B model.
398
+
399
+ ## Disk Offload
400
+
401
+ [Disk Offload](../Pipeline_Usage/VRAM_management.md#disk-offload) is a special VRAM management scheme that needs to be enabled during the model loading process, not after the model is loaded. Usually, when the above code can run smoothly, Disk Offload can be directly enabled:
402
+
403
+ ```python
404
+ from diffsynth.core import load_model, enable_vram_management, AutoWrappedLinear, AutoWrappedModule
405
+ from diffsynth.models.qwen_image_dit import QwenImageDiT, RMSNorm
406
+ import torch
407
+
408
+ prefix = "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model"
409
+ model_path = [prefix + f"-0000{i}-of-00009.safetensors" for i in range(1, 10)]
410
+ inputs = {
411
+ "latents": torch.randn((1, 16, 128, 128), dtype=torch.bfloat16, device="cuda"),
412
+ "timestep": torch.zeros((1,), dtype=torch.bfloat16, device="cuda"),
413
+ "prompt_emb": torch.randn((1, 5, 3584), dtype=torch.bfloat16, device="cuda"),
414
+ "prompt_emb_mask": torch.ones((1, 5), dtype=torch.int64, device="cuda"),
415
+ "height": 1024,
416
+ "width": 1024,
417
+ }
418
+
419
+ model = load_model(
420
+ QwenImageDiT,
421
+ model_path,
422
+ module_map={
423
+ torch.nn.Linear: AutoWrappedLinear,
424
+ RMSNorm: AutoWrappedModule,
425
+ },
426
+ vram_config={
427
+ "offload_dtype": "disk",
428
+ "offload_device": "disk",
429
+ "onload_dtype": "disk",
430
+ "onload_device": "disk",
431
+ "preparing_dtype": torch.bfloat16,
432
+ "preparing_device": "cuda",
433
+ "computation_dtype": torch.bfloat16,
434
+ "computation_device": "cuda",
435
+ },
436
+ vram_limit=0,
437
+ )
438
+ with torch.no_grad():
439
+ output = model(**inputs)
440
+ ```
441
+
442
+ Disk Offload is an extremely special VRAM management scheme. It only supports `.safetensors` format files, not binary files such as `.bin`, `.pth`, `.ckpt`, and does not support [state dict converter](../Developer_Guide/Integrating_Your_Model.md#step-2-model-file-format-conversion) with Tensor reshape.
443
+
444
+ If there are situations where Disk Offload cannot run normally but non-Disk Offload can run normally, please submit an issue to us on GitHub.
445
+
446
+ ## Writing Default Configuration
447
+
448
+ To make it easier for users to use the VRAM management function, we write the fine-grained VRAM management configuration in `diffsynth/configs/vram_management_module_maps.py`. The configuration information for the above model is:
449
+
450
+ ```python
451
+ "diffsynth.models.qwen_image_dit.QwenImageDiT": {
452
+ "diffsynth.models.qwen_image_dit.RMSNorm": "diffsynth.core.vram.layers.AutoWrappedModule",
453
+ "torch.nn.Linear": "diffsynth.core.vram.layers.AutoWrappedLinear",
454
+ }
455
+ ```
docs/en/Developer_Guide/Integrating_Quantization_Backend.md ADDED
@@ -0,0 +1,473 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Integrating a Quantization Backend
2
+
3
+ The quantization framework of `DiffSynth-Studio` lives in `diffsynth.core.quant` and ships with bitsandbytes, torchao, and comfy-kitchen backends (see [Model Quantization](../Pipeline_Usage/Quantization.md)). If you have your own quantization algorithm, or want to plug in another quantization library, you only need to implement a `QuantBackend` — online quantization, saving/loading pre-quantized checkpoints, mixed quantization, VRAM management, and quantization + LoRA training are all reused as-is.
4
+
5
+ This guide walks through the whole process with a toy backend: **INT9** — 9-bit symmetric weight-only quantization, genuinely stored at 9 bits per weight, with one fp32 scale per output channel. INT9 does not exist on any hardware; it is used here because it keeps the example short while still covering every interface you have to implement. For the full interface signatures and contracts, see the [`diffsynth.core.quant` API documentation](../API_Reference/core/quant.md#extension-interface-custom-backends).
6
+
7
+ ## Framework Structure
8
+
9
+ The framework has three layers:
10
+
11
+ - **`QuantizeConfig`**: the user-facing config and entry point, responsible for traversing the model, matching layers, and replacing `nn.Linear`. You never need to touch it.
12
+ - **`QuantBackend`**: the adapter layer, which only ever deals with a **single** `nn.Linear`: how to quantize it, how to build an empty shell, how to dequantize it, how to read and write its state dict. This is the part you implement.
13
+ - **The quantized Linear**: the module that actually holds the quantized weight and performs dequantization + matmul in `forward`.
14
+
15
+ The quantized Linear must satisfy four contract clauses:
16
+
17
+ - **(a)** It is a drop-in replacement for `nn.Linear`, with `forward(x)` doing dequantization + matmul internally. It must subclass `torch.nn.Linear`, otherwise LoRA injection and VRAM management cannot see it.
18
+ - **(b)** `.to(...)` only moves devices, never re-types the packed weight or quant state. VRAM management performs dtype conversions on the model; if a packed weight were cast to bf16, the quant state would be corrupted.
19
+ - **(c)** `state_dict()` and `load_state_dict(assign=True)` round-trip, via `flatten_state_dict` / `unflatten_state_dict` when necessary.
20
+ - **(d)** (Training only) `forward` is differentiable w.r.t. its input, so gradients can pass through frozen quantized layers to reach LoRA branches.
21
+
22
+ ## Step 1: Write the Quantized Linear
23
+
24
+ INT9's storage layout needs a little thought: there is no native 9-bit dtype, and simply putting the codes into an int16 tensor would still spend 16 bits per weight — exactly as much as bf16, so the quantization would save nothing. Each weight is therefore split in two: the low 8 bits go into a uint8 `weight` buffer, and the 9th (most significant) bit forms a separate bit plane where 8 weights are packed into one byte in `weight_msb`, plus one fp32 `weight_scale` per output channel. That is 9 bits per weight, 56% of bf16.
25
+
26
+ Two more details matter:
27
+
28
+ - Delete `nn.Linear`'s original `weight` parameter and register a buffer under the same name, so checkpoint keys stay `layer_name.weight`. Disk offload and mixed-quantization key ownership rely on this.
29
+ - Guard the dtype of the packed tensors by overriding `_apply`, i.e. contract clause (b). Every conversion (`.to()`, `.half()`, `.float()`, ...) funnels through `_apply`, so a conversion that would change the dtype is downgraded to a device-only move.
30
+
31
+ ```python
32
+ from dataclasses import dataclass, field
33
+
34
+ import torch
35
+ import torch.nn.functional as F
36
+
37
+ from diffsynth.core.quant import BackendConfig, QuantBackend, register_quant_backend, register_quant_method
38
+
39
+
40
+ def pack_msb(bits):
41
+ """Pack a 0/1 bit plane into one bit per weight, 8 weights per byte."""
42
+ flat = bits.reshape(-1)
43
+ padding = (-flat.numel()) % 8
44
+ if padding:
45
+ flat = torch.cat([flat, flat.new_zeros(padding)])
46
+ groups = flat.view(-1, 8)
47
+ packed = torch.zeros(groups.shape[0], dtype=torch.uint8, device=flat.device)
48
+ for index in range(8):
49
+ packed |= groups[:, index] << index
50
+ return packed
51
+
52
+
53
+ def unpack_msb(packed, numel):
54
+ bits = torch.stack([(packed >> index) & 1 for index in range(8)], dim=1)
55
+ return bits.reshape(-1)[:numel]
56
+
57
+
58
+ class Int9Linear(torch.nn.Linear):
59
+ """int9 weight: low 8 bits in the uint8 `weight`, the 9th bit packed into `weight_msb`,
60
+ plus one fp32 scale per output channel. 9 bits per weight, 56% of bf16."""
61
+
62
+ dtype_guarded_tensor_names = ("weight", "weight_msb", "weight_scale")
63
+
64
+ def __init__(self, in_features, out_features, bias, compute_dtype):
65
+ with torch.device("meta"):
66
+ super().__init__(in_features, out_features, bias=bias, dtype=compute_dtype)
67
+ del self.weight
68
+ self.register_buffer("weight", torch.empty(out_features, in_features, dtype=torch.uint8, device="meta"))
69
+ self.register_buffer("weight_msb", torch.empty((in_features * out_features + 7) // 8, dtype=torch.uint8, device="meta"))
70
+ self.register_buffer("weight_scale", torch.empty(out_features, dtype=torch.float32, device="meta"))
71
+ if self.bias is not None:
72
+ self.bias.requires_grad_(False)
73
+
74
+ def _apply(self, fn, recurse=True):
75
+ protected = {id(tensor) for name in self.dtype_guarded_tensor_names
76
+ if (tensor := getattr(self, name, None)) is not None}
77
+
78
+ def guard(tensor):
79
+ converted = fn(tensor)
80
+ if id(tensor) in protected and converted.dtype != tensor.dtype:
81
+ return tensor.to(device=converted.device)
82
+ return converted
83
+
84
+ return super()._apply(guard, recurse)
85
+
86
+ def dequantize_weight(self, dtype):
87
+ msb = unpack_msb(self.weight_msb, self.weight.numel()).view_as(self.weight)
88
+ codes = self.weight.to(torch.int16) | (msb.to(torch.int16) << 8)
89
+ return ((codes - 256).float() * self.weight_scale.unsqueeze(1)).to(dtype)
90
+
91
+ def forward(self, x):
92
+ bias = self.bias.to(x.dtype) if self.bias is not None else None
93
+ return F.linear(x, self.dequantize_weight(x.dtype), bias)
94
+ ```
95
+
96
+ Dequantization in `forward` uses ordinary tensor ops, so gradients flow back to the input `x` through `F.linear` and contract clause (d) holds automatically. Unpacking here is written bit by bit with PyTorch ops purely for clarity; a real backend fuses unpacking into the matmul kernel instead of materializing an fp weight on every forward.
97
+
98
+ There is an easy trap in `dequantize_weight`: the integer codes must be reconstructed in fp32. bf16 only carries 8 bits of significand, so integers above 256 are not representable exactly; casting the codes to bf16 before applying the scale rounds the 9th bit away and throws the accuracy gain out (measured: the advantage over int8 collapses from 2.25x to 1.15x). This applies to any format whose code width exceeds the significand of the compute dtype.
99
+
100
+ ## Step 2: Write the Backend
101
+
102
+ Every backend method operates on a single `nn.Linear`:
103
+
104
+ - `capabilities()`: declares what the backend supports; all four flags default to `False`. Saving quantized weights requires `is_serializable=True`, and quantization + LoRA training requires `is_differentiable=True`.
105
+ - `quantized_linear_classes()`: declares the Linear classes this backend produces; `is_quantized_linear` defaults to an `isinstance` check against them.
106
+ - `create_quantized_linear()`: online quantization, turning an fp `nn.Linear` into a quantized one. `compute_device` is where quantization runs and `model_device` is where the result is stored, so the two together stream the work layer by layer with only one layer on the accelerator at a time.
107
+ - `create_quantized_linear_shell()`: builds an empty shell, used for loading pre-quantized checkpoints and for disk offload. It is rebuilt on every offload cycle, so build it on the `meta` device and keep it cheap.
108
+ - `dequantize_to_linear()`: restores a plain `nn.Linear`, used by `mode="dequant_once"`.
109
+ - `flatten_state_dict` / `unflatten_state_dict`: conversion between the state dict and flat tensors. INT9's state dict already holds plain tensors, so the base class implementation is enough; only backends with composite tensors (tensor subclasses, nested quant state) such as bitsandbytes and torchao need to override them.
110
+
111
+ Unimplemented methods raise a descriptive exception from the base class, so a backend that only supports some capabilities can implement just what it needs. `self.config` is the backend config instance injected by the framework — the `Int9WeightOnlyConfig` written in the next step.
112
+
113
+ ```python
114
+ @register_quant_backend("toy_int9")
115
+ class Int9QuantBackend(QuantBackend):
116
+ project_url = "https://example.com/toy-int9"
117
+
118
+ def capabilities(self):
119
+ return {**super().capabilities(), "is_serializable": True, "is_differentiable": True}
120
+
121
+ def quantized_linear_classes(self):
122
+ return (Int9Linear,)
123
+
124
+ def create_quantized_linear(self, linear, compute_device=None, model_device=None):
125
+ weight = linear.weight.data
126
+ if compute_device is not None:
127
+ weight = weight.to(device=compute_device)
128
+ amax = weight.abs().amax(dim=1) if self.config.per_channel else weight.abs().amax().expand(weight.shape[0])
129
+ scale = (amax.float() / 255).clamp(min=1e-8)
130
+ codes = (weight.float() / scale.unsqueeze(1)).round().clamp(-256, 255).to(torch.int16) + 256
131
+
132
+ quant_linear = Int9Linear(linear.in_features, linear.out_features, bias=linear.bias is not None, compute_dtype=weight.dtype)
133
+ quant_linear.weight = (codes & 0xFF).to(torch.uint8)
134
+ quant_linear.weight_msb = pack_msb((codes >> 8).to(torch.uint8))
135
+ quant_linear.weight_scale = scale
136
+ if linear.bias is not None:
137
+ quant_linear.bias = torch.nn.Parameter(linear.bias.data.to(device=scale.device), requires_grad=False)
138
+ return quant_linear if model_device is None else quant_linear.to(device=model_device)
139
+
140
+ def create_quantized_linear_shell(self, linear, compute_dtype):
141
+ return Int9Linear(linear.in_features, linear.out_features, bias=linear.bias is not None, compute_dtype=compute_dtype)
142
+
143
+ def dequantize_to_linear(self, module, compute_dtype, compute_device=None, model_device=None):
144
+ if compute_device is not None:
145
+ module = module.to(device=compute_device)
146
+ fp_weight = module.dequantize_weight(compute_dtype)
147
+ linear = torch.nn.Linear(module.in_features, module.out_features, bias=module.bias is not None, device="meta")
148
+ linear.weight = torch.nn.Parameter(fp_weight, requires_grad=False)
149
+ if module.bias is not None:
150
+ linear.bias = torch.nn.Parameter(module.bias.data.to(dtype=compute_dtype, device=fp_weight.device), requires_grad=False)
151
+ return linear if model_device is None else linear.to(device=model_device)
152
+ ```
153
+
154
+ ## Step 3: Write the Backend Config
155
+
156
+ The backend config subclasses `BackendConfig`: user-tunable parameters are ordinary dataclass fields, while values pinned by the method are declared with `field(init=False, default=...)`. `describe_quant_method` reports the two groups separately, and `from_kwargs` raises when a user passes unknown `backend_config_kwargs`.
157
+
158
+ ```python
159
+ @dataclass
160
+ class Int9WeightOnlyConfig(BackendConfig):
161
+ per_channel: bool = True # user-tunable: per-channel or per-tensor
162
+ bits: int = field(init=False, default=9) # pinned by the method, not overridable
163
+ ```
164
+
165
+ ## Step 4: Register the Quantization Method
166
+
167
+ One backend can register several methods, distinguished by the fields pinned in its config (the bitsandbytes backend, for example, distinguishes nf4 from fp4 via `quant_type`). Method names should follow the `<backend>_<format>_w<weight bits>a<activation bits>` convention:
168
+
169
+ ```python
170
+ register_quant_method("toy_int9_w9a16", "toy_int9", Int9WeightOnlyConfig.from_kwargs, label="9bit, int9, weight-only (toy)")
171
+ ```
172
+
173
+ There are two ways to register a backend and its methods:
174
+
175
+ **Option 1: keep it in your own code (recommended, plug-and-play).** Put the code above in any module; as long as that module is imported before you construct `QuantizeConfig`, the method is already in `QUANT_METHODS` and can be used just like a built-in one, with no framework changes:
176
+
177
+ ```python
178
+ import my_project.toy_int9 # triggers register_quant_backend / register_quant_method
179
+
180
+ from diffsynth.core.quant import QuantizeConfig
181
+
182
+ quantize = QuantizeConfig(method="toy_int9_w9a16", backend_config_kwargs={"per_channel": True})
183
+ ```
184
+
185
+ **Option 2: ship it as a built-in backend (permanent).** Put the backend file under `diffsynth/core/quant/backends/` and register it in `_LAZY_BACKENDS` in `diffsynth/core/quant/backends/__init__.py`; the framework then imports it on demand and users do not need to import anything:
186
+
187
+ ```python
188
+ _LAZY_BACKENDS = {
189
+ "bitsandbytes": ".bitsandbytes",
190
+ "torchao": ".torchao",
191
+ "comfy_kitchen": ".comfy_kitchen",
192
+ "toy_int9": ".toy_int9",
193
+ }
194
+ ```
195
+
196
+ If your quantization algorithm or library is generally useful, you are welcome to submit it as a PR following Option 2, so that more users can benefit from it. A backend that depends on a third-party library should check its dependencies in `validate_environment()` with an installation hint, and point `project_url` at the upstream project.
197
+
198
+ ## Step 5: Self-Check
199
+
200
+ The framework provides two verification tools; run them right after integrating. `check_backend_contract` verifies that the backend declares its Linear classes, that both factory methods return instances of those classes, that all declared classes subclass `nn.Linear`, and that every checkpoint key the backend actually writes lives under the layer name (a missing scale would make disk offload silently load corrupted layers). Unsupported factory methods are skipped rather than counted as failures.
201
+
202
+ ```python
203
+ from diffsynth.core.quant import QUANT_BACKENDS, QUANT_METHODS, check_backend_contract, check_differentiable, describe_quant_method
204
+
205
+ describe_quant_method("toy_int9_w9a16")
206
+
207
+ spec = QUANT_METHODS["toy_int9_w9a16"]
208
+ check_backend_contract(QUANT_BACKENDS[spec.backend](spec.config_factory({})), compute_device="cpu")
209
+ ```
210
+
211
+ The output is as follows; `describe_quant_method` also confirms that the split between user-tunable and pinned parameters is what you intended:
212
+
213
+ ```
214
+ method: toy_int9_w9a16
215
+ backend: toy_int9
216
+ detail: 9bit, int9, weight-only (toy)
217
+ backend config: my_project.toy_int9.Int9WeightOnlyConfig
218
+ backend_config_kwargs (user-tunable):
219
+ per_channel = True
220
+ pinned by method (not overridable):
221
+ bits = 9
222
+ check_backend_contract (toy_int9):
223
+ [PASS] quantized_linear_classes() is non-empty: ['Int9Linear']
224
+ [PASS] Int9Linear subclasses torch.nn.Linear
225
+ [PASS] a plain nn.Linear is not reported as quantized
226
+ [PASS] create_quantized_linear_shell() returns a declared class, got Int9Linear
227
+ [PASS] the shell is recognized before load_state_dict (disk offload routing)
228
+ [PASS] create_quantized_linear() returns a declared class, got Int9Linear
229
+ [PASS] every stored key lives under the layer name; uncovered: []
230
+ => OK
231
+ ```
232
+
233
+ Next, check the numerical error, the real memory saving, the dtype guard of clause (b), and the differentiability of clause (d) on a small model:
234
+
235
+ ```python
236
+ import torch
237
+ from diffsynth.core.quant import QuantizeConfig, check_differentiable
238
+
239
+
240
+ class ToyModel(torch.nn.Module):
241
+ def __init__(self):
242
+ super().__init__()
243
+ self.fc1 = torch.nn.Linear(256, 512)
244
+ self.fc2 = torch.nn.Linear(512, 256, bias=False)
245
+
246
+ def forward(self, x):
247
+ return self.fc2(torch.nn.functional.silu(self.fc1(x)))
248
+
249
+
250
+ def footprint(model):
251
+ return sum(t.numel() * t.element_size() for t in list(model.parameters()) + list(model.buffers()))
252
+
253
+
254
+ torch.manual_seed(0)
255
+ model = ToyModel().to(torch.bfloat16)
256
+ x = torch.randn(4, 256, dtype=torch.bfloat16)
257
+ reference = model(x)
258
+ fp_bytes = footprint(model)
259
+
260
+ QuantizeConfig(method="toy_int9_w9a16").quantize_model(model, compute_device="cpu")
261
+ print("relative error:", ((model(x) - reference).norm() / reference.norm()).item())
262
+ print(f"footprint: {fp_bytes} -> {footprint(model)} bytes ({footprint(model) / fp_bytes:.3f} of bf16)")
263
+
264
+ model.to(torch.float32) # clause (b): packed dtypes must not change
265
+ print(model.fc1.weight.dtype, model.fc1.weight_msb.dtype, model.fc1.weight_scale.dtype, model.fc1.bias.dtype)
266
+
267
+ check_differentiable(model.fc1) # clause (d)
268
+ ```
269
+
270
+ ```
271
+ 2 nn.Linear layers quantized (method: toy_int9_w9a16).
272
+ relative error: 0.004150390625
273
+ footprint: 525312 -> 299008 bytes (0.569 of bf16)
274
+ torch.uint8 torch.uint8 torch.float32 torch.float32
275
+ check_differentiable (Int9Linear): OK -- gradients pass through the module to its input
276
+ ```
277
+
278
+ The measured footprint is 0.569 of bf16, slightly above 9/16 = 0.5625 because of the fp32 scales and the unquantized bias. If this ratio comes out close to 1, the packing format is not actually compressing the weights and you should revisit the storage layout in Step 1.
279
+
280
+ ### Inference on a Real Model: Z-Image
281
+
282
+ Once the small-model checks pass, the backend is ready for real models — a custom backend is used exactly like a built-in method. Import the module that registers it, then pass the method to `ModelConfig(quantize=...)`:
283
+
284
+ ```python
285
+ import torch
286
+
287
+ import my_project.toy_int9 # registers the toy_int9 backend and the toy_int9_w9a16 method
288
+ from diffsynth.core.quant import QuantizeConfig
289
+ from diffsynth.pipelines.z_image import ModelConfig, ZImagePipeline
290
+
291
+ pipe = ZImagePipeline.from_pretrained(
292
+ torch_dtype=torch.bfloat16,
293
+ device="cuda",
294
+ model_configs=[
295
+ ModelConfig(
296
+ model_id="Tongyi-MAI/Z-Image-Turbo",
297
+ origin_file_pattern="transformer/*.safetensors",
298
+ quantize=QuantizeConfig(method="toy_int9_w9a16"),
299
+ ),
300
+ ModelConfig(model_id="Tongyi-MAI/Z-Image-Turbo", origin_file_pattern="text_encoder/*.safetensors"),
301
+ ModelConfig(model_id="Tongyi-MAI/Z-Image-Turbo", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"),
302
+ ],
303
+ tokenizer_config=ModelConfig(model_id="Tongyi-MAI/Z-Image-Turbo", origin_file_pattern="tokenizer/"),
304
+ )
305
+
306
+ dit_bytes = sum(t.numel() * t.element_size() for t in list(pipe.dit.parameters()) + list(pipe.dit.buffers()))
307
+ print(f"dit weights: {dit_bytes / 1024 ** 3:.3f} GiB")
308
+
309
+ prompt = "A delicate portrait of an underwater girl, blue dress flowing, hair gently drifting, light and shadow clear, surrounded by bubbles, serene expression, exquisite details, dreamlike and beautiful."
310
+ image = pipe(prompt=prompt, seed=42, rand_device="cuda")
311
+ image.save("z_image_toy_int9.jpg")
312
+ ```
313
+
314
+ Measured DiT weight footprint on Z-Image Turbo (the 8-step Turbo generation works normally, with no visible quality difference from bf16):
315
+
316
+ | | DiT weights |
317
+ | --- | --- |
318
+ | bf16 | 11.464 GiB |
319
+ | `toy_int9_w9a16` | 6.456 GiB (0.563x) |
320
+
321
+ Note that peak memory and weight footprint are not the same thing: this toy materializes a temporary fp weight on every forward, so the peak saving is smaller than the storage saving. Measured on a synthetic model with 48 Linears, all weights resident on the GPU:
322
+
323
+ | | Weights | Forward peak |
324
+ | --- | --- | --- |
325
+ | bf16 | 1.500 GiB | 1.527 GiB |
326
+ | `toy_int9_w9a16` | 0.845 GiB (0.563x) | 1.036 GiB (0.678x) |
327
+
328
+ That temporary weight depends only on the **largest single layer** and does not grow with depth, so the deeper the model, the closer the peak saving gets to the weight ratio; a real backend fusing unpacking into the matmul kernel does not need it at all. To push the peak down further, stack [VRAM management](../Pipeline_Usage/VRAM_management.md) on top and move weights layer by layer (pass a `vram_config` to each `ModelConfig` above — measured peak drops to 2.1 GiB).
329
+
330
+ ### Accuracy: int9 vs int8
331
+
332
+ Does the extra bit actually buy accuracy? Quantize the same weight with the **identical** per-channel symmetric scheme at 8 and 9 bits, then compare the dequantized weight error and the layer output error. This is the general recipe for an accuracy regression on a new backend: hold everything else fixed and change only the bit width.
333
+
334
+ ```python
335
+ import torch
336
+ from my_project.toy_int9 import Int9QuantBackend, Int9WeightOnlyConfig
337
+
338
+
339
+ def quantize_int8(linear):
340
+ """The same per-channel symmetric scheme with one bit less: codes in [-128, 127]."""
341
+ weight = linear.weight.data
342
+ scale = (weight.abs().amax(dim=1).float() / 127).clamp(min=1e-8)
343
+ codes = (weight.float() / scale.unsqueeze(1)).round().clamp(-128, 127)
344
+ return (codes * scale.unsqueeze(1)).to(weight.dtype)
345
+
346
+
347
+ def relative_error(reference, value):
348
+ return ((value.float() - reference.float()).norm() / reference.float().norm()).item()
349
+
350
+
351
+ torch.manual_seed(0)
352
+ backend = Int9QuantBackend(Int9WeightOnlyConfig())
353
+ linear = torch.nn.Linear(2048, 2048, bias=False).to(torch.bfloat16)
354
+ fp_weight = linear.weight.data.clone()
355
+
356
+ int9_weight = backend.create_quantized_linear(linear).dequantize_weight(torch.bfloat16)
357
+ int8_weight = quantize_int8(linear)
358
+ error8, error9 = relative_error(fp_weight, int8_weight), relative_error(fp_weight, int9_weight)
359
+ print(f"weight error: int8 {error8:.6f} | int9 {error9:.6f} ({error8 / error9:.2f}x lower)")
360
+
361
+ x = torch.randn(64, 2048, dtype=torch.bfloat16)
362
+ reference = torch.nn.functional.linear(x, fp_weight)
363
+ out8 = relative_error(reference, torch.nn.functional.linear(x, int8_weight))
364
+ out9 = relative_error(reference, torch.nn.functional.linear(x, int9_weight))
365
+ print(f"output error: int8 {out8:.6f} | int9 {out9:.6f} ({out8 / out9:.2f}x lower)")
366
+ ```
367
+
368
+ ```
369
+ weight error: int8 0.004353 | int9 0.001937 (2.25x lower)
370
+ output error: int8 0.004947 | int9 0.002816 (1.76x lower)
371
+ ```
372
+
373
+ This matches the theory: going from 255 to 511 levels halves the quantization step, and for uniform quantization the error is proportional to the step, so the weight error drops to roughly half (2.25x measured). The end-to-end layer output gain is smaller (1.76x) because the activations themselves are bf16 and the matmul's own rounding noise eats part of the benefit — a reminder to evaluate bit-width gains at the actual compute precision, not only on the weights.
374
+
375
+ Finally, verify clause (c): save the quantized weights, load them back into shells, and confirm both produce identical outputs.
376
+
377
+ ```python
378
+ from safetensors.torch import load_file, save_file
379
+
380
+ save_config = QuantizeConfig(method="toy_int9_w9a16")
381
+ tensors, metadata = save_config.flatten_state_dict(model.state_dict())
382
+ save_file(tensors, "toy_int9.safetensors", metadata=metadata)
383
+
384
+ loaded = ToyModel().to(torch.bfloat16)
385
+ load_config = QuantizeConfig(method="toy_int9_w9a16", load_prequantized=True)
386
+ load_config.prepare_for_prequantized_load(loaded, compute_dtype=torch.bfloat16)
387
+ loaded.load_state_dict(load_config.unflatten_state_dict(load_file("toy_int9.safetensors"), metadata), assign=True)
388
+ print("reload match:", torch.equal(loaded(x.float()), model(x.float())))
389
+ ```
390
+
391
+ ```
392
+ reload match: True
393
+ ```
394
+
395
+ ### Combining with Disk Offload
396
+
397
+ Disk offload, part of [VRAM management](../Pipeline_Usage/VRAM_management.md), places the strictest demands on a quantization backend: the resident model keeps only `meta` shells, and each layer's tensors are streamed back from disk at forward time and dropped right after. It relies on two things:
398
+
399
+ - It only supports **pre-quantized checkpoints**, so `load_prequantized=True` is required and `prepare_for_prequantized_load` must first swap the target layers for shells.
400
+ - Which tensors a layer needs is resolved by a prefix scan over the checkpoint keys using the layer's dotted name, and the result is loaded with a strict `load_state_dict(assign=True)`. So the only requirement on a backend is that every tensor lives under `{layer_name}.` — flat siblings like `layer.weight_scale` and nested quant state like bnb's both work. A missing or extra key raises instead of silently loading a corrupted layer.
401
+
402
+ ```python
403
+ import torch
404
+ from safetensors.torch import save_file
405
+
406
+ from diffsynth.core.loader.model import load_metadata_from_safetensors
407
+ from diffsynth.core.quant import QuantizeConfig
408
+ from diffsynth.core.vram.disk_map import DiskMap
409
+ from diffsynth.core.vram.layers import AutoWrappedLinear, enable_vram_management_recursively
410
+
411
+ resident = ToyModel().to(torch.bfloat16)
412
+ x = torch.randn(2, 256, dtype=torch.bfloat16, device="cuda")
413
+
414
+ save_config = QuantizeConfig(method="toy_int9_w9a16")
415
+ save_config.quantize_model(resident, compute_device="cuda")
416
+ resident = resident.to("cuda")
417
+ reference = resident(x)
418
+
419
+ tensors, metadata = save_config.flatten_state_dict(resident.state_dict())
420
+ save_file({key: value.cpu() for key, value in tensors.items()}, "toy_int9.safetensors", metadata=metadata)
421
+
422
+ fresh = ToyModel().to(torch.bfloat16)
423
+ load_config = QuantizeConfig(method="toy_int9_w9a16", load_prequantized=True)
424
+ load_config.prepare_for_prequantized_load(fresh, compute_dtype=torch.bfloat16)
425
+ enable_vram_management_recursively(
426
+ fresh,
427
+ module_map={torch.nn.Linear: AutoWrappedLinear},
428
+ vram_config={
429
+ "offload_dtype": "disk", "offload_device": "disk",
430
+ "onload_dtype": "disk", "onload_device": "disk",
431
+ "preparing_dtype": torch.bfloat16, "preparing_device": "cuda",
432
+ "computation_dtype": torch.bfloat16, "computation_device": "cuda",
433
+ },
434
+ disk_map=DiskMap(["toy_int9.safetensors"], "cuda", torch_dtype=None),
435
+ quantize=load_config,
436
+ metadata=load_metadata_from_safetensors("toy_int9.safetensors"),
437
+ )
438
+
439
+ for name, module in fresh.named_modules():
440
+ if getattr(module, "disk_offload", False):
441
+ print(f"{name}: {module._disk_required_keys()}")
442
+
443
+ resident_bytes = sum(t.numel() * t.element_size() for t in list(resident.parameters()) + list(resident.buffers()))
444
+ offloaded_bytes = sum(t.numel() * t.element_size() for t in list(fresh.parameters()) + list(fresh.buffers()) if not t.is_meta)
445
+ print(f"resident {resident_bytes} bytes -> in memory after disk offload {offloaded_bytes} bytes")
446
+ print("output matches:", torch.equal(fresh(x), reference), "| repeatable:", torch.equal(fresh(x), reference))
447
+ ```
448
+
449
+ Measured on the same `ToyModel` (`torch_dtype=None` on `DiskMap` is essential — it guarantees the packed tensors are not re-typed while being read):
450
+
451
+ ```
452
+ 2 nn.Linear layers replaced for loading the pre-quantized checkpoint (method: toy_int9_w9a16).
453
+ fc1: ['fc1.bias', 'fc1.weight', 'fc1.weight_msb', 'fc1.weight_scale']
454
+ fc2: ['fc2.weight', 'fc2.weight_msb', 'fc2.weight_scale']
455
+ resident 299008 bytes -> in memory after disk offload 0 bytes
456
+ output matches: True | repeatable: True
457
+ ```
458
+
459
+ Each layer's `weight` / `weight_msb` / `weight_scale` / `bias` is correctly attributed to that layer, the resident footprint drops to 0 bytes (everything is a `meta` shell), the output is bit-identical to the resident quantized model, and repeated forwards stay stable — so rebuilding shells and streaming from disk has no side effects.
460
+
461
+ On a real model, use the standard workflows from [Model Quantization](../Pipeline_Usage/Quantization.md) for end-to-end validation: pass `QuantizeConfig(method="toy_int9_w9a16")` to `ModelConfig(quantize=...)` for online-quantized inference, save the quantized weights with `save_quantized_model` and load them back after registering the hash, and inject LoRA into the quantized model for training.
462
+
463
+ ## Integration Checklist
464
+
465
+ - The packing format really shrinks the weights: the measured footprint ratio should be close to the theoretical bit-width ratio, not close to 1.
466
+ - The accuracy gain is verified: compared against the same scheme with one bit less, the error really goes down; otherwise precision is being lost somewhere in the dequantization path.
467
+ - The quantized Linear subclasses `torch.nn.Linear`, and all `state_dict` keys live under the layer name.
468
+ - `_apply` guards the dtype of every packed tensor and quant state.
469
+ - `capabilities()` matches reality: declaring `is_serializable` requires a round-tripping state dict, and declaring `is_differentiable` requires passing `check_differentiable`.
470
+ - `create_quantized_linear` honors `compute_device` / `model_device`, so layer-by-layer streaming quantization works.
471
+ - Disk offload works: the shell is built on `meta` and cheap to rebuild, every stored tensor lives under the layer's dotted name, and `unflatten_state_dict` tolerates being called with a single layer's subdict plus whole-file metadata.
472
+ - When depending on a third-party library, `validate_environment()` gives a clear installation hint and `project_url` points at the upstream project.
473
+ - `check_backend_contract` passes completely.
docs/en/Developer_Guide/Integrating_Your_Model.md ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Integrating Model Architecture
2
+
3
+ This document introduces how to integrate models into the `DiffSynth-Studio` framework for use by modules such as `Pipeline`.
4
+
5
+ ## Step 1: Integrate Model Architecture Code
6
+
7
+ All model architecture implementations in `DiffSynth-Studio` are unified in `diffsynth/models`. Each `.py` code file implements a model architecture, and all models are loaded through `ModelPool` in `diffsynth/models/model_loader.py`. When integrating new model architectures, please create a new `.py` file under this path.
8
+
9
+ ```shell
10
+ diffsynth/models/
11
+ ├── general_modules.py
12
+ ├── model_loader.py
13
+ ├── qwen_image_controlnet.py
14
+ ├── qwen_image_dit.py
15
+ ├── qwen_image_text_encoder.py
16
+ ├── qwen_image_vae.py
17
+ └── ...
18
+ ```
19
+
20
+ In most cases, we recommend integrating models in native `PyTorch` code form, with the model architecture class directly inheriting from `torch.nn.Module`, for example:
21
+
22
+ ```python
23
+ import torch
24
+
25
+ class NewDiffSynthModel(torch.nn.Module):
26
+ def __init__(self, dim=1024):
27
+ super().__init__()
28
+ self.linear = torch.nn.Linear(dim, dim)
29
+ self.activation = torch.nn.Sigmoid()
30
+
31
+ def forward(self, x):
32
+ x = self.linear(x)
33
+ x = self.activation(x)
34
+ return x
35
+ ```
36
+
37
+ If the model architecture implementation contains additional dependencies, we strongly recommend removing them, otherwise this will cause heavy package dependency issues. In our existing models, Qwen-Image's Blockwise ControlNet is integrated in this way. The code is lightweight, please refer to `diffsynth/models/qwen_image_controlnet.py`.
38
+
39
+ If the model has been integrated by Huggingface Library ([`transformers`](https://huggingface.co/docs/transformers/main/index), [`diffusers`](https://huggingface.co/docs/diffusers/main/index), etc.), we can integrate the model in a simpler way:
40
+
41
+ <details>
42
+ <summary>Integrating Huggingface Library Style Model Architecture Code</summary>
43
+
44
+ The loading method for these models in Huggingface Library is:
45
+
46
+ ```python
47
+ from transformers import XXX_Model
48
+
49
+ model = XXX_Model.from_pretrained("path_to_your_model")
50
+ ```
51
+
52
+ `DiffSynth-Studio` does not support loading models through `from_pretrained` because this conflicts with VRAM management and other functions. Please rewrite the model architecture in the following format:
53
+
54
+ ```python
55
+ import torch
56
+
57
+ class DiffSynth_XXX_Model(torch.nn.Module):
58
+ def __init__(self):
59
+ super().__init__()
60
+ from transformers import XXX_Config, XXX_Model
61
+ config = XXX_Config(**{
62
+ "architectures": ["XXX_Model"],
63
+ "other_configs": "Please copy and paste the other configs here.",
64
+ })
65
+ self.model = XXX_Model(config)
66
+
67
+ def forward(self, x):
68
+ outputs = self.model(x)
69
+ return outputs
70
+ ```
71
+
72
+ Where `XXX_Config` is the Config class corresponding to the model. For example, the Config class for `Qwen2_5_VLModel` is `Qwen2_5_VLConfig`, which can be found by consulting its source code. The content inside Config can usually be found in the `config.json` file in the model library. `DiffSynth-Studio` will not read the `config.json` file, so the content needs to be copied and pasted into the code.
73
+
74
+ In rare cases, version updates of `transformers` and `diffusers` may cause some models to be unable to import. Therefore, if possible, we still recommend using the model integration method in Step 1.1.
75
+
76
+ In our existing models, Qwen-Image's Text Encoder is integrated in this way. The code is lightweight, please refer to `diffsynth/models/qwen_image_text_encoder.py`.
77
+
78
+ </details>
79
+
80
+ ## Step 2: Model File Format Conversion
81
+
82
+ Due to the variety of model file formats provided by developers in the open-source community, we sometimes need to convert model file formats to form correctly formatted [state dict](https://docs.pytorch.org/tutorials/recipes/recipes/what_is_state_dict.html). This is common in the following situations:
83
+
84
+ * Model files built by different code libraries, for example [Wan-AI/Wan2.1-T2V-1.3B](https://www.modelscope.cn/models/Wan-AI/Wan2.1-T2V-1.3B) and [Wan-AI/Wan2.1-T2V-1.3B-Diffusers](https://www.modelscope.cn/models/Wan-AI/Wan2.1-T2V-1.3B-Diffusers).
85
+ * Models modified during integration, for example, the Text Encoder of [Qwen/Qwen-Image](https://www.modelscope.cn/models/Qwen/Qwen-Image) adds a `model.` prefix in `diffsynth/models/qwen_image_text_encoder.py`.
86
+ * Model files containing multiple models, for example, the VACE Adapter and base DiT model of [Wan-AI/Wan2.1-VACE-14B](https://www.modelscope.cn/models/Wan-AI/Wan2.1-VACE-14B) are mixed and stored in the same set of model files.
87
+
88
+ In our development philosophy, we hope to respect the wishes of model authors as much as possible. If we repackage the model files, for example [Comfy-Org/Qwen-Image_ComfyUI](https://www.modelscope.cn/models/Comfy-Org/Qwen-Image_ComfyUI), although we can call the model more conveniently, traffic (model page views and downloads, etc.) will be directed elsewhere, and the original author of the model will also lose the power to delete the model. Therefore, we have added the `diffsynth/utils/state_dict_converters` module to the framework for file format conversion during model loading.
89
+
90
+ This part of logic is very simple. Taking Qwen-Image's Text Encoder as an example, only 10 lines of code are needed:
91
+
92
+ ```python
93
+ def QwenImageTextEncoderStateDictConverter(state_dict):
94
+ state_dict_ = {}
95
+ for k in state_dict:
96
+ v = state_dict[k]
97
+ if k.startswith("visual."):
98
+ k = "model." + k
99
+ elif k.startswith("model."):
100
+ k = k.replace("model.", "model.language_model.")
101
+ state_dict_[k] = v
102
+ return state_dict_
103
+ ```
104
+
105
+ ## Step 3: Writing Model Config
106
+
107
+ Model Config is located in `diffsynth/configs/model_configs.py`, used to identify model types and load them. The following fields need to be filled in:
108
+
109
+ * `model_hash`: Model file hash value, which can be obtained through the `hash_model_file` function. This hash value is only related to the keys and tensor shapes in the model file's state dict, and is unrelated to other information in the file.
110
+ * `model_name`: Model name, used for `Pipeline` to identify the required model. If different structured models play the same role in `Pipeline`, the same `model_name` can be used. When integrating new models, just ensure that `model_name` is different from other existing functional models. The corresponding model is fetched through `model_name` in the `Pipeline`'s `from_pretrained`.
111
+ * `model_class`: Model architecture import path, pointing to the model architecture class implemented in Step 1, for example `diffsynth.models.qwen_image_text_encoder.QwenImageTextEncoder`.
112
+ * `state_dict_converter`: Optional parameter. If model file format conversion is needed, the import path of the model conversion logic needs to be filled in, for example `diffsynth.utils.state_dict_converters.qwen_image_text_encoder.QwenImageTextEncoderStateDictConverter`.
113
+ * `extra_kwargs`: Optional parameter. If additional parameters need to be passed when initializing the model, these parameters need to be filled in. For example, models [DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny](https://www.modelscope.cn/models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Canny) and [DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Inpaint](https://www.modelscope.cn/models/DiffSynth-Studio/Qwen-Image-Blockwise-ControlNet-Inpaint) both adopt the `QwenImageBlockWiseControlNet` structure in `diffsynth/models/qwen_image_controlnet.py`, but the latter also needs additional configuration `additional_in_dim=4`. Therefore, this configuration information needs to be filled in the `extra_kwargs` field.
114
+
115
+ We provide a piece of code to quickly understand how models are loaded through this configuration information:
116
+
117
+ ```python
118
+ from diffsynth.core import hash_model_file, load_state_dict, skip_model_initialization
119
+ from diffsynth.models.qwen_image_text_encoder import QwenImageTextEncoder
120
+ from diffsynth.utils.state_dict_converters.qwen_image_text_encoder import QwenImageTextEncoderStateDictConverter
121
+ import torch
122
+
123
+ model_hash = "8004730443f55db63092006dd9f7110e"
124
+ model_name = "qwen_image_text_encoder"
125
+ model_class = QwenImageTextEncoder
126
+ state_dict_converter = QwenImageTextEncoderStateDictConverter
127
+ extra_kwargs = {}
128
+
129
+ model_path = [
130
+ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors",
131
+ "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors",
132
+ "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors",
133
+ "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors",
134
+ ]
135
+ if hash_model_file(model_path) == model_hash:
136
+ with skip_model_initialization():
137
+ model = model_class(**extra_kwargs)
138
+ state_dict = load_state_dict(model_path, torch_dtype=torch.bfloat16, device="cuda")
139
+ state_dict = state_dict_converter(state_dict)
140
+ model.load_state_dict(state_dict, assign=True)
141
+ print("Done!")
142
+ ```
143
+
144
+ > Q: The logic of the above code looks very simple, why is this part of code in `DiffSynth-Studio` extremely complex?
145
+ >
146
+ > A: Because we provide aggressive VRAM management functions that are coupled with the model loading logic, this leads to the complexity of the framework structure. We have tried our best to simplify the interface exposed to developers.
147
+
148
+ The `model_hash` in `diffsynth/configs/model_configs.py` is not uniquely existing. Multiple models may exist in the same model file. For this situation, please use multiple model Configs to load each model separately, and write the corresponding `state_dict_converter` to separate the parameters required by each model.
149
+
150
+ ## Step 4: Verifying Whether the Model Can Be Recognized and Loaded
151
+
152
+ After model integration, the following code can be used to verify whether the model can be correctly recognized and loaded. The following code will attempt to load the model into memory:
153
+
154
+ ```python
155
+ from diffsynth.models.model_loader import ModelPool
156
+
157
+ model_pool = ModelPool()
158
+ model_pool.auto_load_model(
159
+ [
160
+ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors",
161
+ "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors",
162
+ "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors",
163
+ "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors",
164
+ ],
165
+ )
166
+ ```
167
+
168
+ If the model can be recognized and loaded, you will see the following output:
169
+
170
+ ```
171
+ Loading models from: [
172
+ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors",
173
+ "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors",
174
+ "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors",
175
+ "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors"
176
+ ]
177
+ Loaded model: {
178
+ "model_name": "qwen_image_text_encoder",
179
+ "model_class": "diffsynth.models.qwen_image_text_encoder.QwenImageTextEncoder",
180
+ "extra_kwargs": null
181
+ }
182
+ ```
183
+
184
+ ## Step 5: Writing Model VRAM Management Scheme
185
+
186
+ `DiffSynth-Studio` supports complex VRAM management. See [Enabling VRAM Management](../Developer_Guide/Enabling_VRAM_management.md) for details.
docs/en/Developer_Guide/Training_Diffusion_Models.md ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Integrating Model Training
2
+
3
+ After [integrating models](../Developer_Guide/Integrating_Your_Model.md) and [implementing Pipeline](../Developer_Guide/Building_a_Pipeline.md), the next step is to integrate model training functionality.
4
+
5
+ ## Training-Inference Consistent Pipeline Modification
6
+
7
+ To ensure strict consistency between training and inference processes, we will use most of the inference code during training, but still need to make minor modifications.
8
+
9
+ First, add extra logic during inference to switch the image-to-image/video-to-video logic based on the `scheduler` state. Taking Qwen-Image as an example:
10
+
11
+ ```python
12
+ class QwenImageUnit_InputImageEmbedder(PipelineUnit):
13
+ def __init__(self):
14
+ super().__init__(
15
+ input_params=("input_image", "noise", "tiled", "tile_size", "tile_stride"),
16
+ output_params=("latents", "input_latents"),
17
+ onload_model_names=("vae",)
18
+ )
19
+
20
+ def process(self, pipe: QwenImagePipeline, input_image, noise, tiled, tile_size, tile_stride):
21
+ if input_image is None:
22
+ return {"latents": noise, "input_latents": None}
23
+ pipe.load_models_to_device(['vae'])
24
+ image = pipe.preprocess_image(input_image).to(device=pipe.device, dtype=pipe.torch_dtype)
25
+ input_latents = pipe.vae.encode(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
26
+ if pipe.scheduler.training:
27
+ return {"latents": noise, "input_latents": input_latents}
28
+ else:
29
+ latents = pipe.scheduler.add_noise(input_latents, noise, timestep=pipe.scheduler.timesteps[0])
30
+ return {"latents": latents, "input_latents": input_latents}
31
+ ```
32
+
33
+ Then, enable Gradient Checkpointing in `model_fn`, which will significantly reduce the VRAM required for training at the cost of computational speed. This is not mandatory, but we strongly recommend doing so.
34
+
35
+ Taking Qwen-Image as an example, before modification:
36
+
37
+ ```python
38
+ text, image = block(
39
+ image=image,
40
+ text=text,
41
+ temb=conditioning,
42
+ image_rotary_emb=image_rotary_emb,
43
+ attention_mask=attention_mask,
44
+ )
45
+ ```
46
+
47
+ After modification:
48
+
49
+ ```python
50
+ from ..core import gradient_checkpoint_forward
51
+
52
+ text, image = gradient_checkpoint_forward(
53
+ block,
54
+ use_gradient_checkpointing,
55
+ use_gradient_checkpointing_offload,
56
+ image=image,
57
+ text=text,
58
+ temb=conditioning,
59
+ image_rotary_emb=image_rotary_emb,
60
+ attention_mask=attention_mask,
61
+ )
62
+ ```
63
+
64
+ ## Writing Training Scripts
65
+
66
+ `DiffSynth-Studio` does not strictly encapsulate the training framework, but exposes the script content to developers. This approach makes it more convenient to modify training scripts to implement additional functions. Developers can refer to existing training scripts, such as `examples/qwen_image/model_training/train.py`, for modification to adapt to new model training.
docs/en/Diffusion_Templates/Introducing_Diffusion_Templates.md ADDED
@@ -0,0 +1,76 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Diffusion Templates
2
+
3
+ Diffusion Templates is a controllable generation plugin framework for Diffusion models in DiffSynth-Studio, providing additional controllable generation capabilities for base models.
4
+
5
+ * Open-source code: [DiffSynth-Studio](https://github.com/modelscope/DiffSynth-Studio)
6
+ * Technical report: [arXiv](https://arxiv.org/abs/2604.24351)
7
+ * Project page: [GitHub](https://modelscope.github.io/diffusion-templates-web/)
8
+ * Documentation reference
9
+ * Introduction to Diffusion Templates: [English Version](https://diffsynth-studio-doc.readthedocs.io/en/latest/Diffusion_Templates/Introducing_Diffusion_Templates.html)、[中文版](https://diffsynth-studio-doc.readthedocs.io/zh-cn/latest/Diffusion_Templates/Introducing_Diffusion_Templates.html)
10
+ * Detailed Architecture of Diffusion Templates: [English Version](https://diffsynth-studio-doc.readthedocs.io/en/latest/Diffusion_Templates/Understanding_Diffusion_Templates.html)、[中文版](https://diffsynth-studio-doc.readthedocs.io/zh-cn/latest/Diffusion_Templates/Understanding_Diffusion_Templates.html)
11
+ * Template Model Inference: [English Version](https://diffsynth-studio-doc.readthedocs.io/en/latest/Diffusion_Templates/Template_Model_Inference.html)、[中文版](https://diffsynth-studio-doc.readthedocs.io/zh-cn/latest/Diffusion_Templates/Template_Model_Inference.html)
12
+ * Template Model Training: [English Version](https://diffsynth-studio-doc.readthedocs.io/en/latest/Diffusion_Templates/Template_Model_Training.html)、[中文版](https://diffsynth-studio-doc.readthedocs.io/zh-cn/latest/Diffusion_Templates/Template_Model_Training.html)
13
+ * Online demo: [ModelScope](https://modelscope.cn/studios/DiffSynth-Studio/Diffusion-Templates)
14
+ * Model collection: [ModelScope](https://modelscope.cn/collections/DiffSynth-Studio/KleinBase4B-Templates)、[ModelScope International](https://modelscope.ai/collections/DiffSynth-Studio/KleinBase4B-Templates)、[HuggingFace](https://huggingface.co/collections/DiffSynth-Studio/kleinbase4b-templates)
15
+
16
+ |Model Name|ModelScope|ModelScope International|HuggingFace|Inference Code|Low VRAM Inference Code|Training Code|Training Validation Code|
17
+ |-|-|-|-|-|-|-|-|
18
+ |Structure Control|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-ControlNet)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-ControlNet)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-ControlNet)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-ControlNet.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-ControlNet.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-ControlNet.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-ControlNet.py)|
19
+ |Brightness Adjustment|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Brightness)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Brightness)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Brightness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Brightness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Brightness.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Brightness.py)|
20
+ |Color Adjustment|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-SoftRGB)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-SoftRGB)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-SoftRGB)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-SoftRGB.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-SoftRGB.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-SoftRGB.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-SoftRGB.py)|
21
+ |Image Editing|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Edit)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Edit)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Edit)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Edit.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Edit.py)|
22
+ |Super-Resolution|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Upscaler)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Upscaler)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Upscaler)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Upscaler.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Upscaler.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Upscaler.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Upscaler.py)|
23
+ |Sharpness Enhancement|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Sharpness)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Sharpness)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Sharpness)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Sharpness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Sharpness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Sharpness.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Sharpness.py)|
24
+ |Aesthetic Alignment|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Aesthetic)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Aesthetic)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Aesthetic)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Aesthetic.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Aesthetic.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Aesthetic.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Aesthetic.py)|
25
+ |Local Redrawing|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Inpaint)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Inpaint)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Inpaint)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Inpaint.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Inpaint.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Inpaint.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Inpaint.py)|
26
+ |Content Reference|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-ContentRef)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-ContentRef)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-ContentRef)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-ContentRef.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-ContentRef.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-ContentRef.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-ContentRef.py)|
27
+ |Age Control|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Age)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-Age)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-Age)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Age.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Age.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Age.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Age.py)|
28
+ |Panda Meme (Easter Egg Model)|[link](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-PandaMeme)|[link](https://modelscope.ai/models/DiffSynth-Studio/Template-KleinBase4B-PandaMeme)|[link](https://huggingface.co/DiffSynth-Studio/Template-KleinBase4B-PandaMeme)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-PandaMeme.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-PandaMeme.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-PandaMeme.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-PandaMeme.py)|
29
+
30
+ * Dataset: [ModelScope](https://modelscope.cn/collections/DiffSynth-Studio/ImagePulseV2)、[ModelScope International](https://modelscope.cn/collections/DiffSynth-Studio/ImagePulseV2)、[HuggingFace](https://huggingface.co/collections/DiffSynth-Studio/imagepulsev2)
31
+
32
+ |Dataset Name|ModelScope|ModelScope International|HuggingFace|
33
+ |-|-|-|-|
34
+ |Text-to-Image|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-TextImage)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-TextImage)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-TextImage)|
35
+ |Local Redrawing|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Inpaint)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Inpaint)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Inpaint)|
36
+ |Background Replacement|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Background)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Background)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Background)|
37
+ |Clothing Replacement|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Clothes)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Clothes)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Clothes)|
38
+ |Pose Adjustment|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Pose)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Pose)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Pose)|
39
+ |Foreground Modification|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Change)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Change)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Change)|
40
+ |Local Addition/Removal|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-AddRemove)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-AddRemove)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-AddRemove)|
41
+ |Super-Resolution|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Upscale)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Upscale)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Upscale)|
42
+ |Portrait Generation|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-TextImage-Human)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-TextImage-Human)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-TextImage-Human)|
43
+ |Random Cropping|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Crop)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Crop)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Crop)|
44
+ |Lighting Adjustment|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Light)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Light)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Light)|
45
+ |Scene Structure|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Structure)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Structure)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Structure)|
46
+ |Facial Expression Editing|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-HumanFace)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-HumanFace)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-HumanFace)|
47
+ |View Angle Adjustment|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Angle)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Angle)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Angle)|
48
+ |Style Transfer|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Style)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Style)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Style)|
49
+ |Multi-Resolution|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-TextImage-MultiResolution)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-TextImage-MultiResolution)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-TextImage-MultiResolution)|
50
+ |Multi-Image Merge|[link](https://modelscope.cn/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Merge)|[link](https://modelscope.ai/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Merge)|[link](https://huggingface.co/datasets/DiffSynth-Studio/ImagePulseV2-Edit-Merge)|
51
+
52
+ ## Model Performance Overview
53
+
54
+ * Super-Resolution + Sharpness Enhancement: Generate ultra-high-resolution images
55
+
56
+ |Low Resolution Input|High Resolution Output|
57
+ |-|-|
58
+ |![](https://github.com/user-attachments/assets/53f378f7-0dc5-44cd-bc39-032d0b1d0208)|![](https://github.com/user-attachments/assets/135bab89-6d76-4d5c-ae5e-44b2826b5c50)|
59
+
60
+ * Structure Control + Aesthetic Alignment + Sharpness Enhancement: Fully-equipped ControlNet
61
+
62
+ |Structure Control Image|Output Image|
63
+ |-|-|
64
+ |![](https://github.com/user-attachments/assets/1feeb13f-f8a7-40df-958c-90463ef5eaf4)|![](https://github.com/user-attachments/assets/ea406387-9695-4efd-b0cb-980686474ab7)|
65
+
66
+ * Structure Control + Image Editing + Color Adjustment: Artistic Style Creation at Will
67
+
68
+ |Structure Control Image|Editing Input Image|Output Image|
69
+ |-|-|-|
70
+ |![](https://github.com/user-attachments/assets/1feeb13f-f8a7-40df-958c-90463ef5eaf4)|![](https://github.com/user-attachments/assets/4866e14b-0ac7-4099-aab5-86048a645cb7)|![](https://github.com/user-attachments/assets/0fd613a5-885b-44b0-83db-9dd08859cc24)|
71
+
72
+ * Brightness Control + Image Editing + Local Redrawing: Cross-dimensional Elements in Images
73
+
74
+ |Reference Image|Redrawing Area|Output Image|
75
+ |-|-|-|
76
+ |![](https://github.com/user-attachments/assets/4866e14b-0ac7-4099-aab5-86048a645cb7)|![](https://github.com/user-attachments/assets/52148a91-7c03-4042-944a-4c3182abe889)|![](https://github.com/user-attachments/assets/3e4cbc26-f6b5-4cc7-a017-d0e0165703ca)|
docs/en/Diffusion_Templates/Template_Model_Inference.md ADDED
@@ -0,0 +1,333 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Template Model Inference
2
+
3
+ ## Enabling Template Models on Base Model Pipelines
4
+
5
+ Using the base model [black-forest-labs/FLUX.2-klein-base-4B](https://modelscope.cn/models/black-forest-labs/FLUX.2-klein-base-4B) as an example, when generating images using only the base model:
6
+
7
+ ```python
8
+ from diffsynth.diffusion.template import TemplatePipeline
9
+ from diffsynth.pipelines.flux2_image import Flux2ImagePipeline, ModelConfig
10
+ import torch
11
+
12
+ # Load base model
13
+ pipe = Flux2ImagePipeline.from_pretrained(
14
+ torch_dtype=torch.bfloat16,
15
+ device="cuda",
16
+ model_configs=[
17
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="text_encoder/*.safetensors"),
18
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-base-4B", origin_file_pattern="transformer/*.safetensors"),
19
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"),
20
+ ],
21
+ tokenizer_config=ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="tokenizer/"),
22
+ )
23
+ # Generate an image
24
+ image = pipe(
25
+ prompt="a cat",
26
+ seed=0, cfg_scale=4,
27
+ height=1024, width=1024,
28
+ )
29
+ image.save("image.png")
30
+ ```
31
+
32
+ The Template model [DiffSynth-Studio/Template-KleinBase4B-Brightness](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness) can control image brightness during generation. Through the `TemplatePipeline` model, it can be loaded from ModelScope (via `ModelConfig(model_id="xxx/xxx")`) or from a local path (via `ModelConfig(path="xxx")`). Inputting `scale=0.8` increases image brightness. Note that in the code, input parameters for `pipe` must be transferred to `template_pipeline`, and `template_inputs` should be added.
33
+
34
+ ```python
35
+ # Load Template model
36
+ template_pipeline = TemplatePipeline.from_pretrained(
37
+ torch_dtype=torch.bfloat16,
38
+ device="cuda",
39
+ model_configs=[
40
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Brightness")
41
+ ],
42
+ )
43
+ # Generate an image
44
+ image = template_pipeline(
45
+ pipe,
46
+ prompt="a cat",
47
+ seed=0, cfg_scale=4,
48
+ height=1024, width=1024,
49
+ template_inputs=[{"scale": 0.8}],
50
+ )
51
+ image.save("image_0.8.png")
52
+ ```
53
+
54
+ ## CFG Enhancement for Template Models
55
+
56
+ Template models can enable CFG (Classifier-Free Guidance) to make control effects more pronounced. For example, with the model [DiffSynth-Studio/Template-KleinBase4B-Brightness](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness), adding `negative_template_inputs` to the TemplatePipeline input parameters and setting its scale to 0.5 will generate images with more noticeable brightness variations by contrasting both sides.
57
+
58
+ ```python
59
+ # Generate an image with CFG
60
+ image = template_pipeline(
61
+ pipe,
62
+ prompt="a cat",
63
+ seed=0, cfg_scale=4,
64
+ height=1024, width=1024,
65
+ template_inputs=[{"scale": 0.8}],
66
+ negative_template_inputs=[{"scale": 0.5}],
67
+ )
68
+ image.save("image_0.8_cfg.png")
69
+ ```
70
+
71
+ ## Low VRAM Support
72
+
73
+ Template models currently do not support the main framework's VRAM management, but lazy loading can be used - loading Template models only when needed for inference. This significantly reduces VRAM requirements when enabling multiple Template models, with peak VRAM usage being that of a single Template model. Add parameter `lazy_loading=True` to enable.
74
+
75
+ ```python
76
+ template_pipeline = TemplatePipeline.from_pretrained(
77
+ torch_dtype=torch.bfloat16,
78
+ device="cuda",
79
+ model_configs=[
80
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Brightness")
81
+ ],
82
+ lazy_loading=True,
83
+ )
84
+ ```
85
+
86
+ The base model's Pipeline and Template Pipeline are completely independent and can enable VRAM management on demand.
87
+
88
+ When Template model outputs contain LoRA in Template Cache, you need to enable VRAM management for the base model's Pipeline or enable LoRA hot loading (using the code below), otherwise LoRA weights will be fused repeatedly.
89
+
90
+ ```python
91
+ pipe.dit = pipe.enable_lora_hot_loading(pipe.dit)
92
+ ```
93
+
94
+ ## Enabling Multiple Template Models
95
+
96
+ `TemplatePipeline` can load multiple Template models. During inference, use `model_id` in `template_inputs` to distinguish inputs for each Template model.
97
+
98
+ After enabling VRAM management for the base model's Pipeline and lazy loading for Template Pipeline, you can load any number of Template models.
99
+
100
+ ```python
101
+ from diffsynth.diffusion.template import TemplatePipeline
102
+ from diffsynth.pipelines.flux2_image import Flux2ImagePipeline, ModelConfig
103
+ from modelscope import dataset_snapshot_download
104
+ import torch
105
+ from PIL import Image
106
+
107
+ vram_config = {
108
+ "offload_dtype": "disk",
109
+ "offload_device": "disk",
110
+ "onload_dtype": torch.bfloat16,
111
+ "onload_device": "cuda",
112
+ "preparing_dtype": torch.bfloat16,
113
+ "preparing_device": "cuda",
114
+ "computation_dtype": torch.bfloat16,
115
+ "computation_device": "cuda",
116
+ }
117
+ pipe = Flux2ImagePipeline.from_pretrained(
118
+ torch_dtype=torch.bfloat16,
119
+ device="cuda",
120
+ model_configs=[
121
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-base-4B", origin_file_pattern="transformer/*.safetensors", **vram_config),
122
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="text_encoder/*.safetensors", **vram_config),
123
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"),
124
+ ],
125
+ tokenizer_config=ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="tokenizer/"),
126
+ )
127
+ pipe.dit = pipe.enable_lora_hot_loading(pipe.dit)
128
+ template = TemplatePipeline.from_pretrained(
129
+ torch_dtype=torch.bfloat16,
130
+ device="cuda",
131
+ lazy_loading=True,
132
+ model_configs=[
133
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Brightness"),
134
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-ControlNet"),
135
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Edit"),
136
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Upscaler"),
137
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-SoftRGB"),
138
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Sharpness"),
139
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Inpaint"),
140
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Aesthetic"),
141
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-ContentRef"),
142
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-Age"),
143
+ ModelConfig(model_id="DiffSynth-Studio/Template-KleinBase4B-PandaMeme"),
144
+ ],
145
+ )
146
+ ```
147
+
148
+ ### Super-Resolution + Sharpness Enhancement
149
+
150
+ Combining [DiffSynth-Studio/Template-KleinBase4B-Upscaler](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Upscaler) and [DiffSynth-Studio/Template-KleinBase4B-Sharpness](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Sharpness) can upscale blurry images while improving detail clarity.
151
+
152
+ ```python
153
+ image = template(
154
+ pipe,
155
+ prompt="A cat is sitting on a stone.",
156
+ seed=0, cfg_scale=4, num_inference_steps=50,
157
+ template_inputs = [
158
+ {
159
+ "model_id": 3,
160
+ "image": Image.open("data/examples/templates/image_lowres_100.jpg"),
161
+ "prompt": "A cat is sitting on a stone.",
162
+ },
163
+ {
164
+ "model_id": 5,
165
+ "scale": 1,
166
+ },
167
+ ],
168
+ negative_template_inputs = [
169
+ {
170
+ "model_id": 3,
171
+ "image": Image.open("data/examples/templates/image_lowres_100.jpg"),
172
+ "prompt": "",
173
+ },
174
+ {
175
+ "model_id": 5,
176
+ "scale": 0,
177
+ },
178
+ ],
179
+ )
180
+ image.save("image_Upscaler_Sharpness.png")
181
+ ```
182
+
183
+ | Low Resolution Input | High Resolution Output |
184
+ |----------------------|------------------------|
185
+ | ![](https://github.com/user-attachments/assets/53f378f7-0dc5-44cd-bc39-032d0b1d0208) | ![](https://github.com/user-attachments/assets/135bab89-6d76-4d5c-ae5e-44b2826b5c50) |
186
+
187
+ ### Structure Control + Aesthetic Alignment + Sharpness Enhancement
188
+
189
+ [DiffSynth-Studio/Template-KleinBase4B-ControlNet](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-ControlNet) controls composition, [DiffSynth-Studio/Template-KleinBase4B-Aesthetic](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Aesthetic) fills in details, and [DiffSynth-Studio/Template-KleinBase4B-Sharpness](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Sharpness) ensures clarity. Combining these three Template models produces exquisite images.
190
+
191
+ ```python
192
+ image = template(
193
+ pipe,
194
+ prompt="A cat is sitting on a stone, bathed in bright sunshine.",
195
+ seed=0, cfg_scale=4, num_inference_steps=50,
196
+ template_inputs = [
197
+ {
198
+ "model_id": 1,
199
+ "image": Image.open("data/examples/templates/image_depth.jpg"),
200
+ "prompt": "A cat is sitting on a stone, bathed in bright sunshine.",
201
+ },
202
+ {
203
+ "model_id": 7,
204
+ "lora_ids": list(range(1, 180, 2)),
205
+ "lora_scales": 2.0,
206
+ "merge_type": "mean",
207
+ },
208
+ {
209
+ "model_id": 5,
210
+ "scale": 0.8,
211
+ },
212
+ ],
213
+ negative_template_inputs = [
214
+ {
215
+ "model_id": 1,
216
+ "image": Image.open("data/examples/templates/image_depth.jpg"),
217
+ "prompt": "",
218
+ },
219
+ {
220
+ "model_id": 7,
221
+ "lora_ids": list(range(1, 180, 2)),
222
+ "lora_scales": 2.0,
223
+ "merge_type": "mean",
224
+ },
225
+ {
226
+ "model_id": 5,
227
+ "scale": 0,
228
+ },
229
+ ],
230
+ )
231
+ image.save("image_Controlnet_Aesthetic_Sharpness.png")
232
+ ```
233
+
234
+ | Structure Control Image | Output Image |
235
+ |-------------------------|--------------|
236
+ | ![](https://github.com/user-attachments/assets/1feeb13f-f8a7-40df-958c-90463ef5eaf4) | ![](https://github.com/user-attachments/assets/ea406387-9695-4efd-b0cb-980686474ab7) |
237
+
238
+ ### Structure Control + Image Editing + Color Adjustment
239
+
240
+ [DiffSynth-Studio/Template-KleinBase4B-ControlNet](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-ControlNet) controls composition, [DiffSynth-Studio/Template-KleinBase4B-Edit](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Edit) preserves original image details like fur texture, and [DiffSynth-Studio/Template-KleinBase4B-SoftRGB](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-SoftRGB) controls color tones, creating an artistic masterpiece.
241
+
242
+ ```python
243
+ image = template(
244
+ pipe,
245
+ prompt="A cat is sitting on a stone. Colored ink painting.",
246
+ seed=0, cfg_scale=4, num_inference_steps=50,
247
+ template_inputs = [
248
+ {
249
+ "model_id": 1,
250
+ "image": Image.open("data/examples/templates/image_depth.jpg"),
251
+ "prompt": "A cat is sitting on a stone. Colored ink painting.",
252
+ },
253
+ {
254
+ "model_id": 2,
255
+ "image": Image.open("data/examples/templates/image_reference.jpg"),
256
+ "prompt": "Convert the image style to colored ink painting.",
257
+ },
258
+ {
259
+ "model_id": 4,
260
+ "R": 0.9,
261
+ "G": 0.5,
262
+ "B": 0.3,
263
+ },
264
+ ],
265
+ negative_template_inputs = [
266
+ {
267
+ "model_id": 1,
268
+ "image": Image.open("data/examples/templates/image_depth.jpg"),
269
+ "prompt": "",
270
+ },
271
+ {
272
+ "model_id": 2,
273
+ "image": Image.open("data/examples/templates/image_reference.jpg"),
274
+ "prompt": "",
275
+ },
276
+ ],
277
+ )
278
+ image.save("image_Controlnet_Edit_SoftRGB.png")
279
+ ```
280
+
281
+ | Structure Control Image | Editing Input Image | Output Image |
282
+ |-------------------------|---------------------|--------------|
283
+ | ![](https://github.com/user-attachments/assets/1feeb13f-f8a7-40df-958c-90463ef5eaf4) | ![](https://github.com/user-attachments/assets/4866e14b-0ac7-4099-aab5-86048a645cb7) | ![](https://github.com/user-attachments/assets/0fd613a5-885b-44b0-83db-9dd08859cc24) |
284
+
285
+ ### Brightness Control + Image Editing + Local Redrawing
286
+
287
+ [DiffSynth-Studio/Template-KleinBase4B-Brightness](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness) generates bright scenes, [DiffSynth-Studio/Template-KleinBase4B-Edit](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Edit) references original image layout, and [DiffSynth-Studio/Template-KleinBase4B-Inpaint](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Inpaint) keeps background unchanged, generating cross-dimensional content.
288
+
289
+ ```python
290
+ image = template(
291
+ pipe,
292
+ prompt="A cat is sitting on a stone. Flat anime style.",
293
+ seed=0, cfg_scale=4, num_inference_steps=50,
294
+ template_inputs = [
295
+ {
296
+ "model_id": 0,
297
+ "scale": 0.6,
298
+ },
299
+ {
300
+ "model_id": 2,
301
+ "image": Image.open("data/examples/templates/image_reference.jpg"),
302
+ "prompt": "Convert the image style to flat anime style.",
303
+ },
304
+ {
305
+ "model_id": 6,
306
+ "image": Image.open("data/examples/templates/image_reference.jpg"),
307
+ "mask": Image.open("data/examples/templates/image_mask_1.jpg"),
308
+ "force_inpaint": True,
309
+ },
310
+ ],
311
+ negative_template_inputs = [
312
+ {
313
+ "model_id": 0,
314
+ "scale": 0.5,
315
+ },
316
+ {
317
+ "model_id": 2,
318
+ "image": Image.open("data/examples/templates/image_reference.jpg"),
319
+ "prompt": "",
320
+ },
321
+ {
322
+ "model_id": 6,
323
+ "image": Image.open("data/examples/templates/image_reference.jpg"),
324
+ "mask": Image.open("data/examples/templates/image_mask_1.jpg"),
325
+ },
326
+ ],
327
+ )
328
+ image.save("image_Brightness_Edit_Inpaint.png")
329
+ ```
330
+
331
+ | Reference Image | Redrawing Area | Output Image |
332
+ |------------------|----------------|--------------|
333
+ | ![](https://github.com/user-attachments/assets/4866e14b-0ac7-4099-aab5-86048a645cb7) | ![](https://github.com/user-attachments/assets/52148a91-7c03-4042-944a-4c3182abe889) | ![](https://github.com/user-attachments/assets/3e4cbc26-f6b5-4cc7-a017-d0e0165703ca) |
docs/en/Diffusion_Templates/Template_Model_Training.md ADDED
@@ -0,0 +1,344 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Template Model Training
2
+
3
+ DiffSynth-Studio currently provides comprehensive Template training support for [black-forest-labs/FLUX.2-klein-base-4B](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-klein-base-4B), with more model adaptations coming soon.
4
+
5
+ ## Continuing Training from Pretrained Models
6
+
7
+ To continue training from our pretrained models, refer to the table in [FLUX.2](../Model_Details/FLUX2.md#model-overview) to find the corresponding training script.
8
+
9
+ ## Building New Template Models
10
+
11
+ ### Template Model Component Format
12
+
13
+ A Template model binds to a model repository (or local folder) containing a code file `model.py` as the entry point. Here's the template for `model.py`:
14
+
15
+ ```python
16
+ import torch
17
+
18
+ class CustomizedTemplateModel(torch.nn.Module):
19
+ def __init__(self):
20
+ super().__init__()
21
+
22
+ @torch.no_grad()
23
+ def process_inputs(self, xxx, **kwargs):
24
+ yyy = xxx
25
+ return {"yyy": yyy}
26
+
27
+ def forward(self, yyy, **kwargs):
28
+ zzz = yyy
29
+ return {"zzz": zzz}
30
+
31
+ class DataProcessor:
32
+ def __call__(self, www, **kwargs):
33
+ xxx = www
34
+ return {"xxx": xxx}
35
+
36
+ TEMPLATE_MODEL = CustomizedTemplateModel
37
+ TEMPLATE_MODEL_PATH = "model.safetensors"
38
+ TEMPLATE_DATA_PROCESSOR = DataProcessor
39
+ ```
40
+
41
+ During Template model inference, Template Input passes through `TEMPLATE_MODEL`'s `process_inputs` and `forward` to generate Template Cache.
42
+
43
+ ```mermaid
44
+ flowchart LR;
45
+ i@{shape: text, label: "Template Input"}-->p[process_inputs];
46
+ subgraph TEMPLATE_MODEL
47
+ p[process_inputs]-->f[forward]
48
+ end
49
+ f[forward]-->c@{shape: text, label: "Template Cache"};
50
+ ```
51
+
52
+ During Template model training, Template Input comes from the dataset through `TEMPLATE_DATA_PROCESSOR`.
53
+
54
+ ```mermaid
55
+ flowchart LR;
56
+ d@{shape: text, label: "Dataset"}-->dp[TEMPLATE_DATA_PROCESSOR]-->p[process_inputs];
57
+ subgraph TEMPLATE_MODEL
58
+ p[process_inputs]-->f[forward]
59
+ end
60
+ f[forward]-->c@{shape: text, label: "Template Cache"};
61
+ ```
62
+
63
+ #### `TEMPLATE_MODEL`
64
+
65
+ `TEMPLATE_MODEL` implements the Template model logic, inheriting from `torch.nn.Module` with required `process_inputs` and `forward` methods. These two methods form the complete Template model inference process, split into two stages to better support [two-stage split training](https://diffsynth-studio-doc.readthedocs.io/en/latest/Training/Split_Training.html).
66
+
67
+ * `process_inputs` must use `@torch.no_grad()` for gradient-free computation
68
+ * `forward` must contain all gradient computations required for training
69
+
70
+ Both methods should accept `**kwargs` for compatibility. Reserved parameters include:
71
+
72
+ * To interact with the base model Pipeline (e.g., call text encoder), add `pipe` parameter to method inputs
73
+ * To enable Gradient Checkpointing, add `use_gradient_checkpointing` and `use_gradient_checkpointing_offload` to `forward` inputs
74
+ * Multiple Template models use `model_id` to distinguish Template Inputs - do not use this field in method parameters
75
+
76
+ #### `TEMPLATE_MODEL_PATH` (Optional)
77
+
78
+ `TEMPLATE_MODEL_PATH` specifies the relative path to pretrained weights. For example:
79
+
80
+ ```python
81
+ TEMPLATE_MODEL_PATH = "model.safetensors"
82
+ ```
83
+
84
+ For multi-file models:
85
+
86
+ ```python
87
+ TEMPLATE_MODEL_PATH = [
88
+ "model-00001-of-00003.safetensors",
89
+ "model-00002-of-00003.safetensors",
90
+ "model-00003-of-00003.safetensors",
91
+ ]
92
+ ```
93
+
94
+ Set to `None` for random initialization:
95
+
96
+ ```python
97
+ TEMPLATE_MODEL_PATH = None
98
+ ```
99
+
100
+ #### `TEMPLATE_DATA_PROCESSOR` (Optional)
101
+
102
+ To train Template models with DiffSynth-Studio, datasets should contain `template_inputs` fields in `metadata.json`. These fields pass through `TEMPLATE_DATA_PROCESSOR` to generate inputs for Template model methods.
103
+
104
+ For example, the brightness control model [DiffSynth-Studio/Template-KleinBase4B-Brightness](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness) takes `scale` as input:
105
+
106
+ ```json
107
+ [
108
+ {
109
+ "image": "images/image_1.jpg",
110
+ "prompt": "a cat",
111
+ "template_inputs": {"scale": 0.2}
112
+ },
113
+ {
114
+ "image": "images/image_2.jpg",
115
+ "prompt": "a dog",
116
+ "template_inputs": {"scale": 0.6}
117
+ }
118
+ ]
119
+ ```
120
+
121
+ ```python
122
+ class DataProcessor:
123
+ def __call__(self, scale, **kwargs):
124
+ return {"scale": scale}
125
+
126
+ TEMPLATE_DATA_PROCESSOR = DataProcessor
127
+ ```
128
+
129
+ Or calculate scale from image paths:
130
+
131
+ ```json
132
+ [
133
+ {
134
+ "image": "images/image_1.jpg",
135
+ "prompt": "a cat",
136
+ "template_inputs": {"image": "/path/to/your/dataset/images/image_1.jpg"}
137
+ }
138
+ ]
139
+ ```
140
+
141
+ ```python
142
+ class DataProcessor:
143
+ def __call__(self, image, **kwargs):
144
+ image = Image.open(image)
145
+ image = np.array(image)
146
+ return {"scale": image.astype(np.float32).mean() / 255}
147
+
148
+ TEMPLATE_DATA_PROCESSOR = DataProcessor
149
+ ```
150
+
151
+ ### Training Template Models
152
+
153
+ A Template model is "trainable" if its Template Cache variables are fully decoupled from the base model Pipeline - these variables should reach `model_fn` without participating in any Pipeline Unit calculations.
154
+
155
+ For training with [black-forest-labs/FLUX.2-klein-base-4B](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-klein-base-4B), use these training script parameters:
156
+
157
+ * `--extra_inputs`: Additional inputs. Use `template_inputs` for text-to-image models, `edit_image,template_inputs` for image editing models
158
+ * `--template_model_id_or_path`: Template model ID or local path (use `:` suffix for ModelScope IDs, e.g., `"DiffSynth-Studio/Template-KleinBase4B-Brightness:"`)
159
+ * `--remove_prefix_in_ckpt`: State dict prefix to remove when saving models (use `"pipe.template_model."`)
160
+ * `--trainable_models`: Trainable components (use `"template_model"` for full model, or `"template_model.xxx,template_model.yyy"` for specific components)
161
+
162
+ Example training script:
163
+
164
+ ```shell
165
+ accelerate launch examples/flux2/model_training/train.py \
166
+ --dataset_base_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness \
167
+ --dataset_metadata_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness/metadata.jsonl \
168
+ --extra_inputs "template_inputs" \
169
+ --max_pixels 1048576 \
170
+ --dataset_repeat 50 \
171
+ --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \
172
+ --template_model_id_or_path "examples/flux2/model_training/scripts/brightness" \
173
+ --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \
174
+ --learning_rate 1e-4 \
175
+ --num_epochs 2 \
176
+ --remove_prefix_in_ckpt "pipe.template_model." \
177
+ --output_path "./models/train/Template-KleinBase4B-Brightness_example" \
178
+ --trainable_models "template_model" \
179
+ --use_gradient_checkpointing \
180
+ --find_unused_parameters
181
+ ```
182
+
183
+ ### Interacting with Base Model Pipeline Components
184
+
185
+ Template models can interact with base model Pipelines. For example, using the text encoder:
186
+
187
+ ```python
188
+ class CustomizedTemplateModel(torch.nn.Module):
189
+ def __init__(self):
190
+ super().__init__()
191
+ self.xxx = xxx()
192
+
193
+ @torch.no_grad()
194
+ def process_inputs(self, text, pipe, **kwargs):
195
+ input_ids = pipe.tokenizer(text)
196
+ text_emb = pipe.text_encoder(input_ids)
197
+ return {"text_emb": text_emb}
198
+
199
+ def forward(self, text_emb, pipe, **kwargs):
200
+ kv_cache = self.xxx(text_emb)
201
+ return {"kv_cache": kv_cache}
202
+
203
+ TEMPLATE_MODEL = CustomizedTemplateModel
204
+ ```
205
+
206
+ ### Using Non-Trainable Components
207
+
208
+ For models with pretrained components:
209
+
210
+ ```python
211
+ class CustomizedTemplateModel(torch.nn.Module):
212
+ def __init__(self):
213
+ super().__init__()
214
+ self.image_encoder = XXXEncoder.from_pretrained(xxx)
215
+ self.mlp = MLP()
216
+
217
+ @torch.no_grad()
218
+ def process_inputs(self, image, **kwargs):
219
+ emb = self.image_encoder(image)
220
+ return {"emb": emb}
221
+
222
+ def forward(self, emb, **kwargs):
223
+ kv_cache = self.mlp(emb)
224
+ return {"kv_cache": kv_cache}
225
+
226
+ TEMPLATE_MODEL = CustomizedTemplateModel
227
+ ```
228
+
229
+ Set `--trainable_models template_model.mlp` to train only the MLP component.
230
+
231
+ ### Training on Low VRAM Devices
232
+
233
+ The framework supports splitting Template model training into two stages: the first stage performs gradient-free computation, and the second stage performs gradient updates. For more information, refer to the documentation: [Two-stage Split Training](https://diffsynth-studio-doc.readthedocs.io/en/latest/Training/Split_Training.html). Here's a sample script:
234
+
235
+ ```shell
236
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "flux2/Template-KleinBase4B-Brightness/*" --local_dir ./data/diffsynth_example_dataset
237
+
238
+ accelerate launch examples/flux2/model_training/train.py \
239
+ --dataset_base_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness \
240
+ --dataset_metadata_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness/metadata.jsonl \
241
+ --extra_inputs "template_inputs" \
242
+ --max_pixels 1048576 \
243
+ --dataset_repeat 1 \
244
+ --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \
245
+ --template_model_id_or_path "DiffSynth-Studio/Template-KleinBase4B-Brightness:" \
246
+ --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \
247
+ --learning_rate 1e-4 \
248
+ --num_epochs 2 \
249
+ --remove_prefix_in_ckpt "pipe.template_model." \
250
+ --output_path "./models/train/Template-KleinBase4B-Brightness_full_cache" \
251
+ --trainable_models "template_model" \
252
+ --use_gradient_checkpointing \
253
+ --find_unused_parameters \
254
+ --task "sft:data_process"
255
+
256
+ accelerate launch examples/flux2/model_training/train.py \
257
+ --dataset_base_path "./models/train/Template-KleinBase4B-Brightness_full_cache" \
258
+ --extra_inputs "template_inputs" \
259
+ --max_pixels 1048576 \
260
+ --dataset_repeat 50 \
261
+ --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors" \
262
+ --template_model_id_or_path "DiffSynth-Studio/Template-KleinBase4B-Brightness:" \
263
+ --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \
264
+ --learning_rate 1e-4 \
265
+ --num_epochs 2 \
266
+ --remove_prefix_in_ckpt "pipe.template_model." \
267
+ --output_path "./models/train/Template-KleinBase4B-Brightness_full" \
268
+ --trainable_models "template_model" \
269
+ --use_gradient_checkpointing \
270
+ --find_unused_parameters \
271
+ --task "sft:train"
272
+ ```
273
+
274
+ Two-stage split training can reduce VRAM requirements and improve training speed. The training process is lossless in precision, but requires significant disk space for storing cache files.
275
+
276
+ To further reduce VRAM requirements, you can enable fp8 precision by adding the parameters `--fp8_models "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors"` and `--fp8_models "black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors"` to the two-stage training. Note that fp8 precision can only be enabled on non-trainable model components and introduces minor errors.
277
+
278
+ ### Uploading Template Models
279
+
280
+ After training, follow these steps to upload Template models to ModelScope for wider distribution.
281
+
282
+ 1. Set model path in `model.py`:
283
+ ```python
284
+ TEMPLATE_MODEL_PATH = "model.safetensors"
285
+ ```
286
+
287
+ 2. Upload using ModelScope CLI:
288
+ ```shell
289
+ modelscope upload user_name/your_model_id /path/to/your/model.py model.py --token ms-xxx
290
+ ```
291
+
292
+ 3. Package model files:
293
+ ```python
294
+ from diffsynth.diffusion.template import load_template_model, load_state_dict
295
+ from safetensors.torch import save_file
296
+ import torch
297
+
298
+ model = load_template_model("path/to/your/template/model", torch_dtype=torch.bfloat16, device="cpu")
299
+ state_dict = load_state_dict("path/to/your/ckpt/epoch-1.safetensors", torch_dtype=torch.bfloat16, device="cpu")
300
+ state_dict.update(model.state_dict())
301
+ save_file(state_dict, "model.safetensors")
302
+ ```
303
+
304
+ 4. Upload model file:
305
+ ```shell
306
+ modelscope upload user_name/your_model_id /path/to/your/model/epoch-1.safetensors model.safetensors --token ms-xxx
307
+ ```
308
+
309
+ 5. Verify inference:
310
+ ```python
311
+ from diffsynth.diffusion.template import TemplatePipeline
312
+ from diffsynth.pipelines.flux2_image import Flux2ImagePipeline, ModelConfig
313
+ import torch
314
+
315
+ # Load base model
316
+ pipe = Flux2ImagePipeline.from_pretrained(
317
+ torch_dtype=torch.bfloat16,
318
+ device="cuda",
319
+ model_configs=[
320
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="text_encoder/*.safetensors"),
321
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-base-4B", origin_file_pattern="transformer/*.safetensors"),
322
+ ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"),
323
+ ],
324
+ tokenizer_config=ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="tokenizer/"),
325
+ )
326
+
327
+ # Load Template model
328
+ template_pipeline = TemplatePipeline.from_pretrained(
329
+ torch_dtype=torch.bfloat16,
330
+ device="cuda",
331
+ model_configs=[
332
+ ModelConfig(model_id="user_name/your_model_id")
333
+ ],
334
+ )
335
+
336
+ # Generate image
337
+ image = template_pipeline(
338
+ pipe,
339
+ prompt="a cat",
340
+ seed=0, cfg_scale=4,
341
+ height=1024, width=1024,
342
+ template_inputs=[{xxx}],
343
+ )
344
+ image.save("image.png")
docs/en/Diffusion_Templates/Understanding_Diffusion_Templates.md ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Diffusion Templates Architecture Details
2
+
3
+ The Diffusion Templates framework is a controllable generation plugin framework in DiffSynth-Studio that provides additional controllable generation capabilities for Diffusion models.
4
+
5
+ ## Framework Structure
6
+
7
+ The Diffusion Templates framework structure is shown below:
8
+
9
+ ```mermaid
10
+ flowchart TD;
11
+ subgraph Template Pipeline
12
+ si@{shape: text, label: "Template Input"}-->i1@{shape: text, label: "Template Input 1"};
13
+ si@{shape: text, label: "Template Input"}-->i2@{shape: text, label: "Template Input 2"};
14
+ si@{shape: text, label: "Template Input"}-->i3@{shape: text, label: "Template Input 3"};
15
+ i1@{shape: text, label: "Template Input 1"}-->m1[Template Model 1]-->c1@{shape: text, label: "Template Cache 1"};
16
+ i2@{shape: text, label: "Template Input 2"}-->m2[Template Model 2]-->c2@{shape: text, label: "Template Cache 2"};
17
+ i3@{shape: text, label: "Template Input 3"}-->m3[Template Model 3]-->c3@{shape: text, label: "Template Cache 3"};
18
+ c1-->c@{shape: text, label: "Template Cache"};
19
+ c2-->c;
20
+ c3-->c;
21
+ end
22
+ i@{shape: text, label: "Model Input"}-->m[Diffusion Pipeline]-->o@{shape: text, label: "Model Output"};
23
+ c-->m;
24
+ ```
25
+
26
+ The framework contains these module designs:
27
+
28
+ * **Template Input**: Template model input. Format: Python dictionary with fields determined by each Template model (e.g., `{"scale": 0.8}`)
29
+ * **Template Model**: Template model, loadable from ModelScope (`ModelConfig(model_id="xxx/xxx")`) or local path (`ModelConfig(path="xxx")`)
30
+ * **Template Cache**: Template model output. Format: Python dictionary with fields matching base model Pipeline input parameters
31
+ * **Template Pipeline**: Module for managing multiple Template models. Handles model loading and cache integration
32
+
33
+ When the Diffusion Templates framework is disabled, base model components (Text Encoder, DiT, VAE) are loaded into the Diffusion Pipeline. Model Input (prompt, height, width) produces Model Output (e.g., images).
34
+
35
+ When enabled, Template models are loaded into the Template Pipeline. The Template Pipeline outputs Template Cache (a subset of Diffusion Pipeline input parameters) for subsequent processing in the Diffusion Pipeline. This enables controllable generation by intercepting part of the Diffusion Pipeline's input parameters.
36
+
37
+ ## Model Capability Medium
38
+
39
+ Template Cache is defined as a subset of Diffusion Pipeline input parameters, ensuring framework generality. We restrict Template model inputs to only be Diffusion Pipeline parameters. The KV-Cache is particularly suitable as a Diffusion medium:
40
+
41
+ * Proven effective in LLM Skills (prompts are converted to KV-Cache)
42
+ * Has "high permission" in Diffusion models - can directly control image generation
43
+ * Supports sequence-level concatenation for multiple Template models
44
+ * Requires minimal development (add pipeline parameter and integrate to model)
45
+
46
+ Other potential Template mediums:
47
+ * **Residual**: Used in ControlNet for point-to-point control, but has resolution limitations and potential conflicts when merging
48
+ * **LoRA**: Treated as input parameters rather than model components
49
+
50
+ **Currently, we only support KV-Cache and LoRA as Template Cache mediums in FLUX.2 Pipeline, with plans to support more models and mediums in the future.**
51
+
52
+ ## Template Model Format
53
+
54
+ A Template model has this structure:
55
+
56
+ ```
57
+ Template_Model
58
+ ├── model.py
59
+ └── model.safetensors
60
+ ```
61
+
62
+ Where `model.py` is the entry point and `model.safetensors` contains model weights. For implementation details, see [Template Model Training](Template_Model_Training.md) or [existing Template models](https://modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness).
docs/en/Makefile ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Minimal makefile for Sphinx documentation
2
+ #
3
+
4
+ # You can set these variables from the command line, and also
5
+ # from the environment for the first two.
6
+ SPHINXOPTS ?=
7
+ SPHINXBUILD ?= sphinx-build
8
+ SOURCEDIR = .
9
+ BUILDDIR = _build
10
+
11
+ # Put it first so that "make" without argument is like "make help".
12
+ help:
13
+ @$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
14
+
15
+ .PHONY: help Makefile
16
+
17
+ # Catch-all target: route all unknown targets to Sphinx using the new
18
+ # "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS).
19
+ %: Makefile
20
+ @$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
docs/en/Model_Details/ACE-Step.md ADDED
@@ -0,0 +1,166 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ACE-Step
2
+
3
+ ACE-Step 1.5 is an open-source music generation model based on DiT architecture, supporting text-to-music, audio cover, repainting and other functionalities, running efficiently on consumer-grade hardware.
4
+
5
+ ## Installation
6
+
7
+ Before performing model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more information on installation, please refer to [Setup Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ Running the following code will load the [ACE-Step/Ace-Step1.5](https://www.modelscope.cn/models/ACE-Step/Ace-Step1.5) model for inference. VRAM management is enabled, the framework automatically controls parameter loading based on available VRAM, requiring a minimum of 3GB VRAM.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.ace_step import AceStepPipeline, ModelConfig
23
+ from diffsynth.utils.data.audio import save_audio
24
+ import torch
25
+
26
+
27
+ vram_config = {
28
+ "offload_dtype": torch.bfloat16,
29
+ "offload_device": "cpu",
30
+ "onload_dtype": torch.bfloat16,
31
+ "onload_device": "cpu",
32
+ "preparing_dtype": torch.bfloat16,
33
+ "preparing_device": "cuda",
34
+ "computation_dtype": torch.bfloat16,
35
+ "computation_device": "cuda",
36
+ }
37
+
38
+
39
+ pipe = AceStepPipeline.from_pretrained(
40
+ torch_dtype=torch.bfloat16,
41
+ device="cuda",
42
+ model_configs=[
43
+ ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="acestep-v15-turbo/model.safetensors", **vram_config),
44
+ ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="Qwen3-Embedding-0.6B/model.safetensors", **vram_config),
45
+ ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="vae/diffusion_pytorch_model.safetensors", **vram_config),
46
+ ],
47
+ text_tokenizer_config=ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="Qwen3-Embedding-0.6B/"),
48
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 0.5,
49
+ )
50
+
51
+ prompt = "An explosive, high-energy pop-rock track with a strong anime theme song feel. The song kicks off with a catchy, synthesized brass fanfare over a driving rock beat with punchy drums and a solid bassline. A powerful, clear male vocal enters with a theatrical and energetic delivery, soaring through the verses and hitting powerful high notes in the chorus. The arrangement is dense and dynamic, featuring rhythmic electric guitar chords, brief instrumental breaks with synth flourishes, and a consistent, danceable groove throughout. The overall mood is triumphant, adventurous, and exhilarating."
52
+ lyrics = '[Intro - Synth Brass Fanfare]\n\n[Verse 1]\n黑夜里的风吹过耳畔\n甜蜜时光转瞬即万\n脚步飘摇在星光上\n心追节奏心跳狂乱\n耳边传来电吉他呼唤\n手指轻触碰点流点燃\n梦在云端任它蔓延\n疯狂跳跃自由无间\n\n[Chorus]\n心电感应在震动间\n拥抱未来勇敢冒险\n那旋律在心中无限\n世界变得如此耀眼\n\n[Instrumental Break - Synth Brass Melody]\n\n[Verse 2]\n鼓点撞击黑夜的底端\n跳动节拍连接你我俩\n在这里让灵魂发光\n燃尽所有不留遗憾\n\n[Instrumental Break - Synth Brass Melody]\n\n[Bridge]\n光影交错彼此的视线\n霓虹之下夜空的蔚蓝\n月光洒下温热心田\n追逐梦想它不会遥远\n\n[Chorus]\n心电感应在震动间\n拥抱未来勇敢冒险\n那旋律在心中无限\n世界变得如此耀眼\n\n[Outro - Instrumental with Synth Brass Melody]\n[Song ends abruptly]'
53
+ audio = pipe(
54
+ prompt=prompt,
55
+ lyrics=lyrics,
56
+ duration=160,
57
+ bpm=100,
58
+ keyscale="B minor",
59
+ timesignature="4",
60
+ vocal_language="zh",
61
+ seed=42,
62
+ )
63
+
64
+ save_audio(audio, pipe.vae.sampling_rate, "acestep-v15-turbo.wav")
65
+ ```
66
+
67
+ ## Model Overview
68
+
69
+ |Model ID|Inference|Low VRAM Inference|Full Training|Full Training Validation|LoRA Training|LoRA Training Validation|
70
+ |-|-|-|-|-|-|-|
71
+ |[ACE-Step/Ace-Step1.5](https://www.modelscope.cn/models/ACE-Step/Ace-Step1.5)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/Ace-Step1.5.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/Ace-Step1.5.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/Ace-Step1.5.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/Ace-Step1.5.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/Ace-Step1.5.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/Ace-Step1.5.py)|
72
+ |[ACE-Step/acestep-v15-turbo-shift1](https://www.modelscope.cn/models/ACE-Step/acestep-v15-turbo-shift1)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-turbo-shift1.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-turbo-shift1.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-turbo-shift1.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-turbo-shift1.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-turbo-shift1.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-turbo-shift1.py)|
73
+ |[ACE-Step/acestep-v15-turbo-shift3](https://www.modelscope.cn/models/ACE-Step/acestep-v15-turbo-shift3)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-turbo-shift3.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-turbo-shift3.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-turbo-shift3.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-turbo-shift3.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-turbo-shift3.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-turbo-shift3.py)|
74
+ |[ACE-Step/acestep-v15-turbo-continuous](https://www.modelscope.cn/models/ACE-Step/acestep-v15-turbo-continuous)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-turbo-continuous.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-turbo-continuous.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-turbo-continuous.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-turbo-continuous.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-turbo-continuous.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-turbo-continuous.py)|
75
+ |[ACE-Step/acestep-v15-base](https://www.modelscope.cn/models/ACE-Step/acestep-v15-base)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-base.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-base.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-base.py)|
76
+ |[ACE-Step/acestep-v15-base: CoverTask](https://www.modelscope.cn/models/ACE-Step/acestep-v15-base)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-base-CoverTask.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-base-CoverTask.py)|—|—|—|—|
77
+ |[ACE-Step/acestep-v15-base: RepaintTask](https://www.modelscope.cn/models/ACE-Step/acestep-v15-base)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-base-RepaintTask.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-base-RepaintTask.py)|—|—|—|—|
78
+ |[ACE-Step/acestep-v15-sft](https://www.modelscope.cn/models/ACE-Step/acestep-v15-sft)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-sft.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-sft.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-sft.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-sft.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-sft.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-sft.py)|
79
+ |[ACE-Step/acestep-v15-xl-base](https://www.modelscope.cn/models/ACE-Step/acestep-v15-xl-base)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-xl-base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-xl-base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-xl-base.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-xl-base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-xl-base.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-xl-base.py)|
80
+ |[ACE-Step/acestep-v15-xl-sft](https://www.modelscope.cn/models/ACE-Step/acestep-v15-xl-sft)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-xl-sft.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-xl-sft.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-xl-sft.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-xl-sft.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-xl-sft.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-xl-sft.py)|
81
+ |[ACE-Step/acestep-v15-xl-turbo](https://www.modelscope.cn/models/ACE-Step/acestep-v15-xl-turbo)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep-v15-xl-turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep-v15-xl-turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep-v15-xl-turbo.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep-v15-xl-turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/lora/acestep-v15-xl-turbo.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_lora/acestep-v15-xl-turbo.py)|
82
+ |[DiffSynth-Studio/acestep15xlsft-lora-music](https://www.modelscope.cn/models/DiffSynth-Studio/acestep15xlsft-lora-music)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference/acestep15xlsft-vocals2music.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_inference_low_vram/acestep15xlsft-vocals2music.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/full/acestep15xlsft-vocals2music.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ace_step/model_training/validate_full/acestep15xlsft-vocals2music.py)|-|-|
83
+
84
+ ## Model Inference
85
+
86
+ The model is loaded via `AceStepPipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models) for details.
87
+
88
+ The input parameters for `AceStepPipeline` inference include:
89
+
90
+ * `prompt`: Text description of the music.
91
+ * `cfg_scale`: Classifier-free guidance scale, defaults to 1.0.
92
+ * `lyrics`: Lyrics text.
93
+ * `task_type`: Task type,可选 values include `"text2music"` (text-to-music), `"cover"` (audio cover), `"repaint"` (repainting), defaults to `"text2music"`.
94
+ * `reference_audios`: List of reference audio tensors for timbre reference.
95
+ * `src_audio`: Source audio tensor for cover or repaint tasks.
96
+ * `denoising_strength`: Denoising strength, controlling how much the output is influenced by source audio, defaults to 1.0.
97
+ * `audio_cover_strength`: Audio cover step ratio, controlling how many steps use cover condition in cover tasks, defaults to 1.0.
98
+ * `audio_code_string`: Input audio code string for cover tasks with discrete audio codes.
99
+ * `repainting_ranges`: List of repainting time ranges (tuples of floats, in seconds) for repaint tasks.
100
+ * `repainting_strength`: Repainting intensity, controlling the degree of change in repainted areas, defaults to 1.0.
101
+ * `duration`: Audio duration in seconds, defaults to 60.
102
+ * `bpm`: Beats per minute, defaults to 100.
103
+ * `keyscale`: Musical key scale, defaults to "B minor".
104
+ * `timesignature`: Time signature, defaults to "4".
105
+ * `vocal_language`: Vocal language, defaults to "unknown".
106
+ * `seed`: Random seed.
107
+ * `rand_device`: Device for noise generation, defaults to "cpu".
108
+ * `num_inference_steps`: Number of inference steps, defaults to 8.
109
+ * `shift`: Timestep shift parameter for the scheduler, defaults to 3.0.
110
+
111
+ ## Model Training
112
+
113
+ Models in the ace_step series are trained uniformly via `examples/ace_step/model_training/train.py`. The script parameters include:
114
+
115
+ * General Training Parameters
116
+ * Dataset Configuration
117
+ * `--dataset_base_path`: Root directory of the dataset.
118
+ * `--dataset_metadata_path`: Path to the dataset metadata file.
119
+ * `--dataset_repeat`: Number of dataset repeats per epoch.
120
+ * `--dataset_num_workers`: Number of processes per DataLoader.
121
+ * `--data_file_keys`: Field names to load from metadata, typically paths to image or video files, separated by `,`.
122
+ * Model Loading Configuration
123
+ * `--model_paths`: Paths to load models from, in JSON format.
124
+ * `--model_id_with_origin_paths`: Model IDs with original paths, separated by commas.
125
+ * `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
126
+ * `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
127
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
128
+ * Basic Training Configuration
129
+ * `--learning_rate`: Learning rate.
130
+ * `--num_epochs`: Number of epochs.
131
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
132
+ * `--find_unused_parameters`: Whether unused parameters exist in DDP training.
133
+ * `--weight_decay`: Weight decay magnitude.
134
+ * `--task`: Training task, defaults to `sft`.
135
+ * Output Configuration
136
+ * `--output_path`: Path to save the model.
137
+ * `--remove_prefix_in_ckpt`: Remove prefix in the model's state dict.
138
+ * `--save_steps`: Interval in training steps to save the model.
139
+ * LoRA Configuration
140
+ * `--lora_base_model`: Which model to add LoRA to.
141
+ * `--lora_target_modules`: Which layers to add LoRA to.
142
+ * `--lora_rank`: Rank of LoRA.
143
+ * `--lora_checkpoint`: Path to LoRA checkpoint.
144
+ * `--preset_lora_path`: Path to preset LoRA checkpoint for LoRA differential training.
145
+ * `--preset_lora_model`: Which model to integrate preset LoRA into, e.g., `dit`.
146
+ * Gradient Configuration
147
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
148
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to CPU memory.
149
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
150
+ * Resolution Configuration
151
+ * `--height`: Height of the image/video. Leave empty to enable dynamic resolution.
152
+ * `--width`: Width of the image/video. Leave empty to enable dynamic resolution.
153
+ * `--max_pixels`: Maximum pixel area, images larger than this will be scaled down during dynamic resolution.
154
+ * `--num_frames`: Number of frames for video (video generation models only).
155
+ * ACE-Step Specific Parameters
156
+ * `--tokenizer_path`: Tokenizer path, in format model_id:origin_pattern.
157
+ * `--silence_latent_path`: Silence latent path, in format model_id:origin_pattern.
158
+ * `--initialize_model_on_cpu`: Whether to initialize models on CPU.
159
+
160
+ ### Example Dataset
161
+
162
+ ```shell
163
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
164
+ ```
165
+
166
+ We provide recommended training scripts for each model, please refer to the table in "Model Overview" above. For guidance on writing model training scripts, see [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, see [Training Framework Overview](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).
docs/en/Model_Details/Anima.md ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Anima
2
+
3
+ Anima is an image generation model trained and open-sourced by CircleStone Labs and Comfy Org.
4
+
5
+ ## Installation
6
+
7
+ Before using this project for model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more installation information, please refer to [Install Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ The following code demonstrates how to quickly load the [circlestone-labs/Anima](https://www.modelscope.cn/models/circlestone-labs/Anima) model for inference. VRAM management is enabled by default, allowing the framework to automatically control model parameter loading based on available VRAM. Minimum 8GB VRAM required.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.anima_image import AnimaImagePipeline, ModelConfig
23
+ import torch
24
+
25
+ vram_config = {
26
+ "offload_dtype": "disk",
27
+ "offload_device": "disk",
28
+ "onload_dtype": "disk",
29
+ "onload_device": "disk",
30
+ "preparing_dtype": torch.bfloat16,
31
+ "preparing_device": "cuda",
32
+ "computation_dtype": torch.bfloat16,
33
+ "computation_device": "cuda",
34
+ }
35
+ pipe = AnimaImagePipeline.from_pretrained(
36
+ torch_dtype=torch.bfloat16,
37
+ device="cuda",
38
+ model_configs=[
39
+ ModelConfig(model_id="circlestone-labs/Anima", origin_file_pattern="split_files/diffusion_models/anima-preview.safetensors", **vram_config),
40
+ ModelConfig(model_id="circlestone-labs/Anima", origin_file_pattern="split_files/text_encoders/qwen_3_06b_base.safetensors", **vram_config),
41
+ ModelConfig(model_id="circlestone-labs/Anima", origin_file_pattern="split_files/vae/qwen_image_vae.safetensors", **vram_config),
42
+ ],
43
+ tokenizer_config=ModelConfig(model_id="Qwen/Qwen3-0.6B", origin_file_pattern="./"),
44
+ tokenizer_t5xxl_config=ModelConfig(model_id="stabilityai/stable-diffusion-3.5-large", origin_file_pattern="tokenizer_3/"),
45
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 0.5,
46
+ )
47
+ prompt = "Masterpiece, best quality, solo, long hair, wavy hair, silver hair, blue eyes, blue dress, medium breasts, dress, underwater, air bubble, floating hair, refraction, portrait."
48
+ negative_prompt = "worst quality, low quality, monochrome, zombie, interlocked fingers, Aissist, cleavage, nsfw,"
49
+ image = pipe(prompt, seed=0, num_inference_steps=50)
50
+ image.save("image.jpg")
51
+ ```
52
+
53
+ ## Model Overview
54
+
55
+ |Model ID|Inference|Low VRAM Inference|Full Training|Validation after Full Training|LoRA Training|Validation after LoRA Training|
56
+ |-|-|-|-|-|-|-|
57
+ |[circlestone-labs/Anima](https://www.modelscope.cn/models/circlestone-labs/Anima)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_inference/anima-preview.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_inference_low_vram/anima-preview.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_training/full/anima-preview.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_training/validate_full/anima-preview.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_training/lora/anima-preview.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_training/validate_lora/anima-preview.py)|
58
+
59
+ Special training scripts:
60
+
61
+ * Differential LoRA Training: [doc](../Training/Differential_LoRA.md)
62
+ * FP8 Precision Training: [doc](../Training/FP8_Precision.md)
63
+ * Two-Stage Split Training: [doc](../Training/Split_Training.md)
64
+ * End-to-End Direct Distillation: [doc](../Training/Direct_Distill.md)
65
+
66
+ ## Model Inference
67
+
68
+ Models are loaded through `AnimaImagePipeline.from_pretrained`, see [Model Inference](../Pipeline_Usage/Model_Inference.md#loading-models) for details.
69
+
70
+ Input parameters for `AnimaImagePipeline` inference include:
71
+
72
+ * `prompt`: Text description of the desired image content.
73
+ * `negative_prompt`: Content to exclude from the generated image (default: `""`).
74
+ * `cfg_scale`: Classifier-free guidance parameter (default: 4.0).
75
+ * `input_image`: Input image for image-to-image generation (default: `None`).
76
+ * `denoising_strength`: Controls similarity to input image (default: 1.0).
77
+ * `height`: Image height (must be multiple of 16, default: 1024).
78
+ * `width`: Image width (must be multiple of 16, default: 1024).
79
+ * `seed`: Random seed (default: `None`).
80
+ * `rand_device`: Device for random noise generation (default: `"cpu"`).
81
+ * `num_inference_steps`: Inference steps (default: 30).
82
+ * `sigma_shift`: Scheduler sigma offset (default: `None`).
83
+ * `progress_bar_cmd`: Progress bar implementation (default: `tqdm.tqdm`).
84
+
85
+ For VRAM constraints, enable [VRAM Management](../Pipeline_Usage/VRAM_management.md). Recommended low-VRAM configurations are provided in the "Model Overview" table above.
86
+
87
+ ## Model Training
88
+
89
+ Anima models are trained through [`examples/anima/model_training/train.py`](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/anima/model_training/train.py) with parameters including:
90
+
91
+ * General Training Parameters
92
+ * Dataset Configuration
93
+ * `--dataset_base_path`: Dataset root directory.
94
+ * `--dataset_metadata_path`: Metadata file path.
95
+ * `--dataset_repeat`: Dataset repetition per epoch.
96
+ * `--dataset_num_workers`: Dataloader worker count.
97
+ * `--data_file_keys`: Metadata fields to load (comma-separated).
98
+ * Model Loading
99
+ * `--model_paths`: Model paths (JSON format).
100
+ * `--model_id_with_origin_paths`: Model IDs with origin paths (e.g., `"anima-team/anima-1B:text_encoder/*.safetensors"`).
101
+ * `--extra_inputs`: Additional pipeline inputs (e.g., `controlnet_inputs` for ControlNet).
102
+ * `--fp8_models`: FP8-formatted models (same format as `--model_paths`).
103
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
104
+ * Training Configuration
105
+ * `--learning_rate`: Learning rate.
106
+ * `--num_epochs`: Training epochs.
107
+ * `--trainable_models`: Trainable components (e.g., `dit`, `vae`, `text_encoder`).
108
+ * `--find_unused_parameters`: Handle unused parameters in DDP training.
109
+ * `--weight_decay`: Weight decay value.
110
+ * `--task`: Training task (default: `sft`).
111
+ * Output Configuration
112
+ * `--output_path`: Model output directory.
113
+ * `--remove_prefix_in_ckpt`: Remove state dict prefixes.
114
+ * `--save_steps`: Model saving interval.
115
+ * LoRA Configuration
116
+ * `--lora_base_model`: Target model for LoRA.
117
+ * `--lora_target_modules`: Target modules for LoRA.
118
+ * `--lora_rank`: LoRA rank.
119
+ * `--lora_checkpoint`: LoRA checkpoint path.
120
+ * `--preset_lora_path`: Preloaded LoRA checkpoint path.
121
+ * `--preset_lora_model`: Model to merge LoRA with (e.g., `dit`).
122
+ * Gradient Configuration
123
+ * `--use_gradient_checkpointing`: Enable gradient checkpointing.
124
+ * `--use_gradient_checkpointing_offload`: Offload checkpointing to CPU.
125
+ * `--gradient_accumulation_steps`: Gradient accumulation steps.
126
+ * Image Resolution
127
+ * `--height`: Image height (empty for dynamic resolution).
128
+ * `--width`: Image width (empty for dynamic resolution).
129
+ * `--max_pixels`: Maximum pixel area for dynamic resolution.
130
+ * Anima-Specific Parameters
131
+ * `--tokenizer_path`: Tokenizer path for text-to-image models.
132
+ * `--tokenizer_t5xxl_path`: T5-XXL tokenizer path.
133
+
134
+ We provide a sample image dataset for testing:
135
+
136
+ ```shell
137
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
138
+ ```
139
+
140
+ For training script details, refer to [Model Training](../Pipeline_Usage/Model_Training.md). For advanced training techniques, see [Training Framework Documentation](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/zh/Training/).
docs/en/Model_Details/Boogu-Image.md ADDED
@@ -0,0 +1,148 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Boogu-Image
2
+
3
+ Boogu-Image supports text-to-image, image-to-image, and instruction-guided image editing.
4
+
5
+ ## Installation
6
+
7
+ Before performing model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more information on installation, please refer to [Setup Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ Running the following code will load the [Boogu/Boogu-Image-0.1-Base](https://modelscope.cn/models/Boogu/Boogu-Image-0.1-Base) model for inference. VRAM management is enabled, the framework automatically controls parameter loading based on available VRAM, requiring a minimum of 8GB VRAM.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.boogu_image import BooguImagePipeline, ModelConfig
23
+ import torch
24
+
25
+
26
+ vram_config = {
27
+ "offload_dtype": torch.float8_e4m3fn,
28
+ "offload_device": "cpu",
29
+ "onload_dtype": torch.float8_e4m3fn,
30
+ "onload_device": "cpu",
31
+ "preparing_dtype": torch.float8_e4m3fn,
32
+ "preparing_device": "cuda",
33
+ "computation_dtype": torch.bfloat16,
34
+ "computation_device": "cuda",
35
+ }
36
+
37
+ pipe = BooguImagePipeline.from_pretrained(
38
+ torch_dtype=torch.bfloat16,
39
+ device="cuda",
40
+ model_configs=[
41
+ ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="transformer/*.safetensors", **vram_config),
42
+ ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="mllm/*.safetensors", **vram_config),
43
+ ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="vae/*.safetensors", **vram_config),
44
+ ],
45
+ processor_config=ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="mllm/"),
46
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 0.5,
47
+ )
48
+
49
+ output = pipe(
50
+ prompt="a cat",
51
+ negative_prompt="",
52
+ height=1024,
53
+ width=1024,
54
+ seed=42,
55
+ num_inference_steps=50,
56
+ cfg_scale=4.0,
57
+ )
58
+ output.save("image_Boogu-Image-0.1-Base.jpg")
59
+ ```
60
+
61
+ ## Model Overview
62
+
63
+ |Model ID|Inference|Low VRAM Inference|Full Training|Full Training Validation|LoRA Training|LoRA Training Validation|
64
+ |-|-|-|-|-|-|-|
65
+ |[Boogu/Boogu-Image-0.1-Base](https://modelscope.cn/models/Boogu/Boogu-Image-0.1-Base)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_inference/Boogu-Image-0.1-Base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_inference_low_vram/Boogu-Image-0.1-Base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/full/Boogu-Image-0.1-Base.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/validate_full/Boogu-Image-0.1-Base.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/lora/Boogu-Image-0.1-Base.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/validate_lora/Boogu-Image-0.1-Base.py)|
66
+ |[Boogu/Boogu-Image-0.1-Turbo](https://modelscope.cn/models/Boogu/Boogu-Image-0.1-Turbo)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_inference/Boogu-Image-0.1-Turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_inference_low_vram/Boogu-Image-0.1-Turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/full/Boogu-Image-0.1-Turbo.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/validate_full/Boogu-Image-0.1-Turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/lora/Boogu-Image-0.1-Turbo.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/validate_lora/Boogu-Image-0.1-Turbo.py)|
67
+ |[Boogu/Boogu-Image-0.1-Edit](https://modelscope.cn/models/Boogu/Boogu-Image-0.1-Edit)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_inference/Boogu-Image-0.1-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_inference_low_vram/Boogu-Image-0.1-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/full/Boogu-Image-0.1-Edit.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/validate_full/Boogu-Image-0.1-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/lora/Boogu-Image-0.1-Edit.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/boogu_image/model_training/validate_lora/Boogu-Image-0.1-Edit.py)|
68
+
69
+ ## Model Inference
70
+
71
+ The model is loaded via `BooguImagePipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models) for details.
72
+
73
+ The input parameters for `BooguImagePipeline` inference include:
74
+
75
+ * `prompt`: Text prompt describing the desired content or editing instruction.
76
+ * `negative_prompt`: Negative prompt specifying what should not appear in the result, defaults to empty string.
77
+ * `cfg_scale`: Classifier-free guidance scale factor, defaults to 4.0. Higher values make the output more closely follow the prompt.
78
+ * `input_image`: Input image for image-to-image (img2img). When provided, the input image is noised and denoised according to `denoising_strength`.
79
+ * `edit_image`: Image to be edited for instruction-guided editing. When provided, the model modifies the image according to the `prompt` instruction.
80
+ * `height`: Height of the output image, defaults to 1024. Must be divisible by 16.
81
+ * `width`: Width of the output image, defaults to 1024. Must be divisible by 16.
82
+ * `seed`: Random seed for reproducibility. Set to `None` for random seed.
83
+ * `denoising_strength`: Denoising strength controlling how much the input image is repainted, defaults to 1.0. Only effective when `input_image` is provided.
84
+ * `sigmas`: Custom sigma scheduling sequence to override the default scheduling strategy. Required for Turbo models.
85
+ * `num_inference_steps`: Number of inference steps, defaults to 20. More steps typically yield better quality.
86
+ * `max_sequence_length`: Maximum sequence length for the text encoder, defaults to 1280.
87
+ * `max_input_image_pixels`: Maximum pixel area for input images, defaults to 4194304. Images larger than this will be scaled down.
88
+ * `max_input_image_side_length`: Maximum side length for input images, defaults to 4096.
89
+ * `max_vlm_input_pil_pixels`: Maximum pixel area for VLM input images, defaults to 147456. Only effective in image editing mode.
90
+ * `max_vlm_input_pil_side_length`: Maximum side length for VLM input images, defaults to 768. Only effective in image editing mode.
91
+ * `rand_device`: Device for generating initial noise, defaults to "cpu".
92
+ * `progress_bar_cmd`: Progress bar display mode, defaults to tqdm.
93
+
94
+ When running low on VRAM, please refer to [VRAM Management](../Pipeline_Usage/VRAM_management.md) to enable VRAM management features.
95
+
96
+ ## Model Training
97
+
98
+ Models in the boogu_image series are trained uniformly via `examples/boogu_image/model_training/train.py`. The script parameters include:
99
+
100
+ * General Training Parameters
101
+ * Dataset Configuration
102
+ * `--dataset_base_path`: Root directory of the dataset.
103
+ * `--dataset_metadata_path`: Path to the dataset metadata file.
104
+ * `--dataset_repeat`: Number of dataset repeats per epoch.
105
+ * `--dataset_num_workers`: Number of processes per DataLoader.
106
+ * `--data_file_keys`: Field names to load from metadata, typically paths to image or video files, separated by `,`.
107
+ * Model Loading Configuration
108
+ * `--model_paths`: Paths to load models from, in JSON format.
109
+ * `--model_id_with_origin_paths`: Model IDs with original paths, separated by commas.
110
+ * `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
111
+ * `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
112
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
113
+ * Basic Training Configuration
114
+ * `--learning_rate`: Learning rate.
115
+ * `--num_epochs`: Number of epochs.
116
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
117
+ * `--find_unused_parameters`: Whether unused parameters exist in DDP training.
118
+ * `--weight_decay`: Weight decay magnitude.
119
+ * `--task`: Training task, defaults to `sft`.
120
+ * Output Configuration
121
+ * `--output_path`: Path to save the model.
122
+ * `--remove_prefix_in_ckpt`: Remove prefix in the model's state dict.
123
+ * `--save_steps`: Interval in training steps to save the model.
124
+ * LoRA Configuration
125
+ * `--lora_base_model`: Which model to add LoRA to.
126
+ * `--lora_target_modules`: Which layers to add LoRA to.
127
+ * `--lora_rank`: Rank of LoRA.
128
+ * `--lora_checkpoint`: Path to LoRA checkpoint.
129
+ * `--preset_lora_path`: Path to preset LoRA checkpoint for LoRA differential training.
130
+ * `--preset_lora_model`: Which model to integrate preset LoRA into, e.g., `dit`.
131
+ * Gradient Configuration
132
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
133
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to CPU memory.
134
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
135
+ * Resolution Configuration
136
+ * `--height`: Height of the image/video. Leave empty to enable dynamic resolution.
137
+ * `--width`: Width of the image/video. Leave empty to enable dynamic resolution.
138
+ * `--max_pixels`: Maximum pixel area, images larger than this will be scaled down during dynamic resolution.
139
+ * `--num_frames`: Number of frames for video (video generation models only).
140
+ * Boogu-Image Specific Parameters
141
+ * `--processor_path`: Path to the processor for processing text and image encoder inputs.
142
+ * `--initialize_model_on_cpu`: Whether to initialize models on CPU. By default, models are initialized on the accelerator device.
143
+
144
+ ```shell
145
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
146
+ ```
147
+
148
+ We provide recommended training scripts for each model, please refer to the table in "Model Overview" above. For guidance on writing model training scripts, see [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, see [Training Framework Overview](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).
docs/en/Model_Details/ERNIE-Image.md ADDED
@@ -0,0 +1,135 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ERNIE-Image
2
+
3
+ ERNIE-Image is a powerful image generation model with 8B parameters developed by Baidu, featuring a compact and efficient architecture with strong instruction-following capability. Based on an 8B DiT backbone, it delivers performance comparable to larger (20B+) models in certain scenarios while maintaining parameter efficiency. It offers reliable performance in instruction understanding and execution, text generation (English/Chinese/Japanese), and overall stability.
4
+
5
+ ## Installation
6
+
7
+ Before performing model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more information on installation, please refer to [Setup Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ Running the following code will load the [PaddlePaddle/ERNIE-Image](https://www.modelscope.cn/models/PaddlePaddle/ERNIE-Image) model for inference. VRAM management is enabled, the framework automatically controls parameter loading based on available VRAM, requiring a minimum of 3G VRAM.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.ernie_image import ErnieImagePipeline, ModelConfig
23
+ import torch
24
+
25
+ vram_config = {
26
+ "offload_dtype": torch.bfloat16,
27
+ "offload_device": "cpu",
28
+ "onload_dtype": torch.bfloat16,
29
+ "onload_device": "cpu",
30
+ "preparing_dtype": torch.bfloat16,
31
+ "preparing_device": "cuda",
32
+ "computation_dtype": torch.bfloat16,
33
+ "computation_device": "cuda",
34
+ }
35
+ pipe = ErnieImagePipeline.from_pretrained(
36
+ torch_dtype=torch.bfloat16,
37
+ device='cuda',
38
+ model_configs=[
39
+ ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="transformer/diffusion_pytorch_model*.safetensors", **vram_config),
40
+ ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="text_encoder/model.safetensors", **vram_config),
41
+ ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors", **vram_config),
42
+ ],
43
+ tokenizer_config=ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="tokenizer/"),
44
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 0.5,
45
+ )
46
+
47
+ image = pipe(
48
+ prompt="一只黑白相间的中华田园犬",
49
+ negative_prompt="",
50
+ height=1024,
51
+ width=1024,
52
+ seed=42,
53
+ num_inference_steps=50,
54
+ cfg_scale=4.0,
55
+ )
56
+ image.save("output.jpg")
57
+ ```
58
+
59
+ ## Model Overview
60
+
61
+ |Model ID|Inference|Low VRAM Inference|Full Training|Full Training Validation|LoRA Training|LoRA Training Validation|
62
+ |-|-|-|-|-|-|-|
63
+ |[PaddlePaddle/ERNIE-Image](https://www.modelscope.cn/models/PaddlePaddle/ERNIE-Image)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_inference/ERNIE-Image.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_inference_low_vram/ERNIE-Image.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_training/full/ERNIE-Image.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_training/validate_full/ERNIE-Image.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_training/lora/ERNIE-Image.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_training/validate_lora/ERNIE-Image.py)|
64
+ |[PaddlePaddle/ERNIE-Image-Turbo](https://www.modelscope.cn/models/PaddlePaddle/ERNIE-Image-Turbo)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_inference/ERNIE-Image-Turbo.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_inference_low_vram/ERNIE-Image-Turbo.py)|—|—|—|—|
65
+
66
+ ## Model Inference
67
+
68
+ The model is loaded via `ErnieImagePipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models) for details.
69
+
70
+ The input parameters for `ErnieImagePipeline` inference include:
71
+
72
+ * `prompt`: The prompt describing the content to appear in the image.
73
+ * `negative_prompt`: The negative prompt describing what should not appear in the image, default value is `""`.
74
+ * `cfg_scale`: Classifier-free guidance parameter, default value is 4.0.
75
+ * `height`: Image height, must be a multiple of 16, default value is 1024.
76
+ * `width`: Image width, must be a multiple of 16, default value is 1024.
77
+ * `seed`: Random seed. Default is `None`, meaning completely random.
78
+ * `rand_device`: The computing device for generating random Gaussian noise matrices, default is `"cuda"`. When set to `cuda`, different GPUs will produce different results.
79
+ * `num_inference_steps`: Number of inference steps, default value is 50.
80
+
81
+ If VRAM is insufficient, please enable [VRAM Management](../Pipeline_Usage/VRAM_management.md). We provide recommended low-VRAM configurations for each model in the "Model Overview" table above.
82
+
83
+ ## Model Training
84
+
85
+ ERNIE-Image series models are trained uniformly via [`examples/ernie_image/model_training/train.py`](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ernie_image/model_training/train.py). The script parameters include:
86
+
87
+ * General Training Parameters
88
+ * Dataset Configuration
89
+ * `--dataset_base_path`: Root directory of the dataset.
90
+ * `--dataset_metadata_path`: Path to the dataset metadata file.
91
+ * `--dataset_repeat`: Number of dataset repeats per epoch.
92
+ * `--dataset_num_workers`: Number of processes per DataLoader.
93
+ * `--data_file_keys`: Field names to load from metadata, typically paths to image or video files, separated by `,`.
94
+ * Model Loading Configuration
95
+ * `--model_paths`: Paths to load models from, in JSON format.
96
+ * `--model_id_with_origin_paths`: Model IDs with original paths, e.g., `"PaddlePaddle/ERNIE-Image:transformer/diffusion_pytorch_model*.safetensors"`, separated by commas.
97
+ * `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
98
+ * `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
99
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
100
+ * Basic Training Configuration
101
+ * `--learning_rate`: Learning rate.
102
+ * `--num_epochs`: Number of epochs.
103
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
104
+ * `--find_unused_parameters`: Whether unused parameters exist in DDP training.
105
+ * `--weight_decay`: Weight decay magnitude.
106
+ * `--task`: Training task, defaults to `sft`.
107
+ * Output Configuration
108
+ * `--output_path`: Path to save the model.
109
+ * `--remove_prefix_in_ckpt`: Remove prefix in the model's state dict.
110
+ * `--save_steps`: Interval in training steps to save the model.
111
+ * LoRA Configuration
112
+ * `--lora_base_model`: Which model to add LoRA to.
113
+ * `--lora_target_modules`: Which layers to add LoRA to.
114
+ * `--lora_rank`: Rank of LoRA.
115
+ * `--lora_checkpoint`: Path to LoRA checkpoint.
116
+ * `--preset_lora_path`: Path to preset LoRA checkpoint for LoRA differential training.
117
+ * `--preset_lora_model`: Which model to integrate preset LoRA into, e.g., `dit`.
118
+ * Gradient Configuration
119
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
120
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to CPU memory.
121
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
122
+ * Resolution Configuration
123
+ * `--height`: Height of the image. Leave empty to enable dynamic resolution.
124
+ * `--width`: Width of the image. Leave empty to enable dynamic resolution.
125
+ * `--max_pixels`: Maximum pixel area, images larger than this will be scaled down during dynamic resolution.
126
+ * ERNIE-Image Specific Parameters
127
+ * `--tokenizer_path`: Path to the tokenizer, leave empty to auto-download from remote.
128
+
129
+ We provide an example image dataset for testing, which can be downloaded with the following command:
130
+
131
+ ```shell
132
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
133
+ ```
134
+
135
+ We provide recommended training scripts for each model, please refer to the table in "Model Overview" above. For guidance on writing model training scripts, see [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, see [Training Framework Overview](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).
docs/en/Model_Details/FLUX.md ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # FLUX
2
+
3
+ ![Image](https://github.com/user-attachments/assets/c01258e2-f251-441a-aa1e-ebb22f02594d)
4
+
5
+ FLUX is an image generation model series developed and open-sourced by Black Forest Labs.
6
+
7
+ ## Installation
8
+
9
+ Before using this project for model inference and training, please install DiffSynth-Studio first.
10
+
11
+ ```shell
12
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
13
+ cd DiffSynth-Studio
14
+ pip install -e .
15
+ ```
16
+
17
+ For more information about installation, please refer to [Install Dependencies](../Pipeline_Usage/Setup.md).
18
+
19
+ ## Quick Start
20
+
21
+ Run the following code to quickly load the [black-forest-labs/FLUX.1-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-dev) model and perform inference. VRAM management is enabled, and the framework will automatically control model parameter loading based on remaining VRAM. Minimum 8GB VRAM is required to run.
22
+
23
+ ```python
24
+ import torch
25
+ from diffsynth.pipelines.flux_image import FluxImagePipeline, ModelConfig
26
+
27
+ vram_config = {
28
+ "offload_dtype": torch.float8_e4m3fn,
29
+ "offload_device": "cpu",
30
+ "onload_dtype": torch.float8_e4m3fn,
31
+ "onload_device": "cpu",
32
+ "preparing_dtype": torch.float8_e4m3fn,
33
+ "preparing_device": "cuda",
34
+ "computation_dtype": torch.bfloat16,
35
+ "computation_device": "cuda",
36
+ }
37
+ pipe = FluxImagePipeline.from_pretrained(
38
+ torch_dtype=torch.bfloat16,
39
+ device="cuda",
40
+ model_configs=[
41
+ ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="flux1-dev.safetensors", **vram_config),
42
+ ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="text_encoder/model.safetensors", **vram_config),
43
+ ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="text_encoder_2/*.safetensors", **vram_config),
44
+ ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="ae.safetensors", **vram_config),
45
+ ],
46
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 1,
47
+ )
48
+ prompt = "CG, masterpiece, best quality, solo, long hair, wavy hair, silver hair, blue eyes, blue dress, medium breasts, dress, underwater, air bubble, floating hair, refraction, portrait. The girl's flowing silver hair shimmers with every color of the rainbow and cascades down, merging with the floating flora around her."
49
+ image = pipe(prompt=prompt, seed=0)
50
+ image.save("image.jpg")
51
+ ```
52
+
53
+ ## Model Overview
54
+
55
+ | Model ID | Extra Parameters | Inference | Low VRAM Inference | Full Training | Validation After Full Training | LoRA Training | Validation After LoRA Training |
56
+ | - | - | - | - | - | - | - | - |
57
+ | [black-forest-labs/FLUX.1-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-dev) | | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev.py) |
58
+ | [black-forest-labs/FLUX.1-Krea-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-Krea-dev) | | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-Krea-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-Krea-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-Krea-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-Krea-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-Krea-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-Krea-dev.py) |
59
+ | [black-forest-labs/FLUX.1-Kontext-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-Kontext-dev) | `kontext_images` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-Kontext-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-Kontext-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-Kontext-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-Kontext-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-Kontext-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-Kontext-dev.py) |
60
+ | [black-forest-labs/FLUX.1-Fill-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-Fill-dev) | `flux_fill_image`, `flux_fill_mask` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-Fill-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-Fill-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-Fill-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-Fill-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-Fill-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-Fill-dev.py) |
61
+ | [black-forest-labs/FLUX.1-Redux-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-Redux-dev) | `flux_redux_image` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-Redux-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-Redux-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-Redux-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-Redux-dev.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-Redux-dev.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-Redux-dev.py) |
62
+ | [HuanJue/Insert-Anything](https://www.modelscope.cn/models/HuanJue/Insert-Anything) | `insert_anything_source_image`, `insert_anything_source_mask`, `insert_anything_ref_image`, `insert_anything_ref_mask` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/Insert-Anything.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/Insert-Anything.py) | - | - | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/Insert-Anything.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/Insert-Anything.py) |
63
+ | [alimama-creative/FLUX.1-dev-Controlnet-Inpainting-Beta](https://www.modelscope.cn/models/alimama-creative/FLUX.1-dev-Controlnet-Inpainting-Beta) | `controlnet_inputs` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-Controlnet-Inpainting-Beta.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-Controlnet-Inpainting-Beta.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev-Controlnet-Inpainting-Beta.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev-Controlnet-Inpainting-Beta.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev-Controlnet-Inpainting-Beta.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev-Controlnet-Inpainting-Beta.py) |
64
+ | [InstantX/FLUX.1-dev-Controlnet-Union-alpha](https://www.modelscope.cn/models/InstantX/FLUX.1-dev-Controlnet-Union-alpha) | `controlnet_inputs` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-Controlnet-Union-alpha.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-Controlnet-Union-alpha.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev-Controlnet-Union-alpha.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev-Controlnet-Union-alpha.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev-Controlnet-Union-alpha.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev-Controlnet-Union-alpha.py) |
65
+ | [jasperai/Flux.1-dev-Controlnet-Upscaler](https://www.modelscope.cn/models/jasperai/Flux.1-dev-Controlnet-Upscaler) | `controlnet_inputs` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-Controlnet-Upscaler.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-Controlnet-Upscaler.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev-Controlnet-Upscaler.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev-Controlnet-Upscaler.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev-Controlnet-Upscaler.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev-Controlnet-Upscaler.py) |
66
+ | [InstantX/FLUX.1-dev-IP-Adapter](https://www.modelscope.cn/models/InstantX/FLUX.1-dev-IP-Adapter) | `ipadapter_images`, `ipadapter_scale` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-IP-Adapter.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-IP-Adapter.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev-IP-Adapter.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev-IP-Adapter.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev-IP-Adapter.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev-IP-Adapter.py) |
67
+ | [ByteDance/InfiniteYou](https://www.modelscope.cn/models/ByteDance/InfiniteYou) | `infinityou_id_image`, `infinityou_guidance`, `controlnet_inputs` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-InfiniteYou.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-InfiniteYou.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev-InfiniteYou.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev-InfiniteYou.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev-InfiniteYou.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev-InfiniteYou.py) |
68
+ | [DiffSynth-Studio/Eligen](https://www.modelscope.cn/models/DiffSynth-Studio/Eligen) | `eligen_entity_prompts`, `eligen_entity_masks`, `eligen_enable_on_negative`, `eligen_enable_inpaint` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-EliGen.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-EliGen.py) | - | - | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLUX.1-dev-EliGen.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLUX.1-dev-EliGen.py) |
69
+ | [DiffSynth-Studio/LoRA-Encoder-FLUX.1-Dev](https://www.modelscope.cn/models/DiffSynth-Studio/LoRA-Encoder-FLUX.1-Dev) | `lora_encoder_inputs`, `lora_encoder_scale` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-LoRA-Encoder.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLUX.1-dev-LoRA-Encoder.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLUX.1-dev-LoRA-Encoder.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLUX.1-dev-LoRA-Encoder.py) | - | - |
70
+ | [DiffSynth-Studio/LoRAFusion-preview-FLUX.1-dev](https://modelscope.cn/models/DiffSynth-Studio/LoRAFusion-preview-FLUX.1-dev) | | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLUX.1-dev-LoRA-Fusion.py) | - | - | - | - | - |
71
+ | [stepfun-ai/Step1X-Edit](https://www.modelscope.cn/models/stepfun-ai/Step1X-Edit) | `step1x_reference_image` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/Step1X-Edit.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/Step1X-Edit.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/Step1X-Edit.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/Step1X-Edit.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/Step1X-Edit.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/Step1X-Edit.py) |
72
+ | [ostris/Flex.2-preview](https://www.modelscope.cn/models/ostris/Flex.2-preview) | `flex_inpaint_image`, `flex_inpaint_mask`, `flex_control_image`, `flex_control_strength`, `flex_control_stop` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/FLEX.2-preview.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/FLEX.2-preview.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/FLEX.2-preview.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/FLEX.2-preview.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/FLEX.2-preview.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/FLEX.2-preview.py) |
73
+ | [DiffSynth-Studio/Nexus-GenV2](https://www.modelscope.cn/models/DiffSynth-Studio/Nexus-GenV2) | `nexus_gen_reference_image` | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference/Nexus-Gen-Editing.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_inference_low_vram/Nexus-Gen-Editing.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/full/Nexus-Gen.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_full/Nexus-Gen.py) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/lora/Nexus-Gen.sh) | [code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/validate_lora/Nexus-Gen.py) |
74
+
75
+ Special Training Scripts:
76
+
77
+ * Differential LoRA Training: [doc](../Training/Differential_LoRA.md)
78
+ * FP8 Precision Training: [doc](../Training/FP8_Precision.md)
79
+ * Two-stage Split Training: [doc](../Training/Split_Training.md)
80
+ * End-to-end Direct Distillation: [doc](../Training/Direct_Distill.md)
81
+
82
+ ## Model Inference
83
+
84
+ Models are loaded via `FluxImagePipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models).
85
+
86
+ Input parameters for `FluxImagePipeline` inference include:
87
+
88
+ * `prompt`: Prompt describing the content appearing in the image.
89
+ * `negative_prompt`: Negative prompt describing content that should not appear in the image, default value is `""`.
90
+ * `cfg_scale`: Classifier-free guidance parameter, default value is 1. When set to a value greater than 1, CFG is enabled.
91
+ * `height`: Image height, must be a multiple of 16.
92
+ * `width`: Image width, must be a multiple of 16.
93
+ * `seed`: Random seed. Default is `None`, meaning completely random.
94
+ * `rand_device`: Computing device for generating random Gaussian noise matrix, default is `"cpu"`. When set to `cuda`, different GPUs will produce different generation results.
95
+ * `num_inference_steps`: Number of inference steps, default value is 30.
96
+ * `embedded_guidance`: Embedded guidance parameter, default value is 3.5.
97
+ * `t5_sequence_length`: Sequence length of the T5 text encoder, default is 512.
98
+ * `tiled`: Whether to enable VAE tiling inference, default is `False`. Setting to `True` can significantly reduce VRAM usage during VAE encoding/decoding stages, producing slight errors and slightly longer inference time.
99
+ * `tile_size`: Tile size during VAE encoding/decoding stages, default is 128, only effective when `tiled=True`.
100
+ * `tile_stride`: Tile stride during VAE encoding/decoding stages, default is 64, only effective when `tiled=True`, must be less than or equal to `tile_size`.
101
+ * `progress_bar_cmd`: Progress bar, default is `tqdm.tqdm`. Can be disabled by setting to `lambda x:x`.
102
+ * `controlnet_inputs`: ControlNet model inputs, type is `ControlNetInput` list.
103
+ * `ipadapter_images`: IP-Adapter model input image list.
104
+ * `ipadapter_scale`: Guidance strength of the IP-Adapter model.
105
+ * `infinityou_id_image`: InfiniteYou model input image.
106
+ * `infinityou_guidance`: Guidance strength of the InfiniteYou model.
107
+ * `kontext_images`: Kontext model input images.
108
+ * `eligen_entity_prompts`: EliGen partition control prompt list.
109
+ * `eligen_entity_masks`: EliGen partition control region mask image list.
110
+ * `eligen_enable_on_negative`: Whether to enable EliGen partition control on the negative side of CFG.
111
+ * `eligen_enable_inpaint`: Whether to enable EliGen partition control inpainting function.
112
+ * `lora_encoder_inputs`: LoRA encoder input image list.
113
+ * `lora_encoder_scale`: Guidance strength of the LoRA encoder.
114
+ * `step1x_reference_image`: Step1X model reference image.
115
+ * `flex_inpaint_image`: Flex model image to be inpainted.
116
+ * `flex_inpaint_mask`: Flex model inpainting mask.
117
+ * `flex_control_image`: Flex model control image.
118
+ * `flex_control_strength`: Flex model control strength.
119
+ * `flex_control_stop`: Flex model control stop timestep.
120
+ * `nexus_gen_reference_image`: Nexus-Gen model reference image.
121
+ * `flux_fill_image`: FLUX.1-Fill model image to be inpainted.
122
+ * `flux_fill_mask`: FLUX.1-Fill model inpainting mask.
123
+ * `flux_redux_image`: FLUX.1-Redux model reference image.
124
+ * `insert_anything_source_image`: Insert-Anything model source image, i.e., the target image to be edited.
125
+ * `insert_anything_source_mask`: Insert-Anything model source image mask, specifying the region to be edited.
126
+ * `insert_anything_ref_image`: Insert-Anything model reference image, providing the content to be inserted.
127
+ * `insert_anything_ref_mask`: Insert-Anything model reference image mask, specifying the target object in the reference image.
128
+
129
+ If VRAM is insufficient, please enable [VRAM Management](../Pipeline_Usage/VRAM_management.md). We provide recommended low VRAM configurations for each model in the example code, see the table in the "Model Overview" section above.
130
+
131
+ ## Model Training
132
+
133
+ FLUX series models are uniformly trained through [`examples/flux/model_training/train.py`](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux/model_training/train.py), and the script parameters include:
134
+
135
+ * General Training Parameters
136
+ * Dataset Basic Configuration
137
+ * `--dataset_base_path`: Root directory of the dataset.
138
+ * `--dataset_metadata_path`: Metadata file path of the dataset.
139
+ * `--dataset_repeat`: Number of times the dataset is repeated in each epoch.
140
+ * `--dataset_num_workers`: Number of processes for each DataLoader.
141
+ * `--data_file_keys`: Field names to be loaded from metadata, usually image or video file paths, separated by `,`.
142
+ * Model Loading Configuration
143
+ * `--model_paths`: Paths of models to be loaded. JSON format.
144
+ * `--model_id_with_origin_paths`: Model IDs with original paths, e.g., `"black-forest-labs/FLUX.1-dev:flux1-dev.safetensors"`. Separated by commas.
145
+ * `--extra_inputs`: Extra input parameters required by the model Pipeline, e.g., `controlnet_inputs` when training ControlNet models, separated by `,`.
146
+ * `--fp8_models`: Models loaded in FP8 format, consistent with `--model_paths` or `--model_id_with_origin_paths` format. Currently only supports models whose parameters are not updated by gradients (no gradient backpropagation, or gradients only update their LoRA).
147
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
148
+ * Training Basic Configuration
149
+ * `--learning_rate`: Learning rate.
150
+ * `--num_epochs`: Number of epochs.
151
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
152
+ * `--find_unused_parameters`: Whether there are unused parameters in DDP training. Some models contain redundant parameters that do not participate in gradient calculation, and this setting needs to be enabled to avoid errors in multi-GPU training.
153
+ * `--weight_decay`: Weight decay size, see [torch.optim.AdamW](https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html).
154
+ * `--task`: Training task, default is `sft`. Some models support more training modes, please refer to the documentation of each specific model.
155
+ * Output Configuration
156
+ * `--output_path`: Model saving path.
157
+ * `--remove_prefix_in_ckpt`: Remove prefix in the state dict of the model file.
158
+ * `--save_steps`: Interval of training steps to save the model. If this parameter is left blank, the model is saved once per epoch.
159
+ * LoRA Configuration
160
+ * `--lora_base_model`: Which model to add LoRA to.
161
+ * `--lora_target_modules`: Which layers to add LoRA to.
162
+ * `--lora_rank`: Rank of LoRA.
163
+ * `--lora_checkpoint`: Path of the LoRA checkpoint. If this path is provided, LoRA will be loaded from this checkpoint.
164
+ * `--preset_lora_path`: Preset LoRA checkpoint path. If this path is provided, this LoRA will be loaded in the form of being merged into the base model. This parameter is used for LoRA differential training.
165
+ * `--preset_lora_model`: Model that the preset LoRA is merged into, e.g., `dit`.
166
+ * Gradient Configuration
167
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
168
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to memory.
169
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
170
+ * Image Width/Height Configuration (Applicable to Image Generation and Video Generation Models)
171
+ * `--height`: Height of image or video. Leave `height` and `width` blank to enable dynamic resolution.
172
+ * `--width`: Width of image or video. Leave `height` and `width` blank to enable dynamic resolution.
173
+ * `--max_pixels`: Maximum pixel area of image or video frames. When dynamic resolution is enabled, images with resolution larger than this value will be downscaled, and images with resolution smaller than this value will remain unchanged.
174
+ * FLUX Specific Parameters
175
+ * `--tokenizer_1_path`: Path of the CLIP tokenizer, leave blank to automatically download from remote.
176
+ * `--tokenizer_2_path`: Path of the T5 tokenizer, leave blank to automatically download from remote.
177
+ * `--align_to_opensource_format`: Whether to align LoRA format to open-source format, only applicable to DiT's LoRA.
178
+
179
+ We have built a sample image dataset for your testing. You can download this dataset with the following command:
180
+
181
+ ```shell
182
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
183
+ ```
184
+
185
+ We have written recommended training scripts for each model, please refer to the table in the "Model Overview" section above. For how to write model training scripts, please refer to [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, please refer to [Training Framework Detailed Explanation](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).
docs/en/Model_Details/FLUX2.md ADDED
@@ -0,0 +1,155 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # FLUX.2
2
+
3
+ FLUX.2 is an image generation model trained and open-sourced by Black Forest Labs.
4
+
5
+ ## Installation
6
+
7
+ Before using this project for model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more information about installation, please refer to [Install Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ Run the following code to quickly load the [black-forest-labs/FLUX.2-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-dev) model and perform inference. VRAM management is enabled, and the framework will automatically control model parameter loading based on remaining VRAM. Minimum 10GB VRAM is required to run.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.flux2_image import Flux2ImagePipeline, ModelConfig
23
+ import torch
24
+
25
+ vram_config = {
26
+ "offload_dtype": "disk",
27
+ "offload_device": "disk",
28
+ "onload_dtype": torch.float8_e4m3fn,
29
+ "onload_device": "cpu",
30
+ "preparing_dtype": torch.float8_e4m3fn,
31
+ "preparing_device": "cuda",
32
+ "computation_dtype": torch.bfloat16,
33
+ "computation_device": "cuda",
34
+ }
35
+ pipe = Flux2ImagePipeline.from_pretrained(
36
+ torch_dtype=torch.bfloat16,
37
+ device="cuda",
38
+ model_configs=[
39
+ ModelConfig(model_id="black-forest-labs/FLUX.2-dev", origin_file_pattern="text_encoder/*.safetensors", **vram_config),
40
+ ModelConfig(model_id="black-forest-labs/FLUX.2-dev", origin_file_pattern="transformer/*.safetensors", **vram_config),
41
+ ModelConfig(model_id="black-forest-labs/FLUX.2-dev", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"),
42
+ ],
43
+ tokenizer_config=ModelConfig(model_id="black-forest-labs/FLUX.2-dev", origin_file_pattern="tokenizer/"),
44
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 0.5,
45
+ )
46
+ prompt = "High resolution. A dreamy underwater portrait of a serene young woman in a flowing blue dress. Her hair floats softly around her face, strands delicately suspended in the water. Clear, shimmering light filters through, casting gentle highlights, while tiny bubbles rise around her. Her expression is calm, her features finely detailed—creating a tranquil, ethereal scene."
47
+ image = pipe(prompt, seed=42, rand_device="cuda", num_inference_steps=50)
48
+ image.save("image.jpg")
49
+ ```
50
+
51
+ ## Model Overview
52
+
53
+ | Model ID | Inference | Low VRAM Inference | Full Training | Validation After Full Training | LoRA Training | Validation After LoRA Training |
54
+ | - | - | - | - | - | - | - |
55
+ |[black-forest-labs/FLUX.2-dev](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-dev)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/FLUX.2-dev.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/FLUX.2-dev.py)|-|-|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/lora/FLUX.2-dev.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_lora/FLUX.2-dev.py)|
56
+ |[black-forest-labs/FLUX.2-klein-4B](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-klein-4B)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/FLUX.2-klein-4B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/FLUX.2-klein-4B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/FLUX.2-klein-4B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/FLUX.2-klein-4B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/lora/FLUX.2-klein-4B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_lora/FLUX.2-klein-4B.py)|
57
+ |[black-forest-labs/FLUX.2-klein-9B](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-klein-9B)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/FLUX.2-klein-9B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/FLUX.2-klein-9B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/FLUX.2-klein-9B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/FLUX.2-klein-9B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/lora/FLUX.2-klein-9B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_lora/FLUX.2-klein-9B.py)|
58
+ |[black-forest-labs/FLUX.2-klein-base-4B](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-klein-base-4B)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/FLUX.2-klein-base-4B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/FLUX.2-klein-base-4B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/FLUX.2-klein-base-4B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/FLUX.2-klein-base-4B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/lora/FLUX.2-klein-base-4B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_lora/FLUX.2-klein-base-4B.py)|
59
+ |[black-forest-labs/FLUX.2-klein-base-9B](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-klein-base-9B)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/FLUX.2-klein-base-9B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/FLUX.2-klein-base-9B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/FLUX.2-klein-base-9B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/FLUX.2-klein-base-9B.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/lora/FLUX.2-klein-base-9B.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_lora/FLUX.2-klein-base-9B.py)|
60
+ |[DiffSynth-Studio/Template-KleinBase4B-Aesthetic](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Aesthetic)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Aesthetic.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Aesthetic.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Aesthetic.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Aesthetic.py)|-|-|
61
+ |[DiffSynth-Studio/Template-KleinBase4B-Brightness](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Brightness)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Brightness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Brightness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Brightness.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Brightness.py)|-|-|
62
+ |[DiffSynth-Studio/Template-KleinBase4B-Age](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Age)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Age.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Age.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Age.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Age.py)|-|-|
63
+ |[DiffSynth-Studio/Template-KleinBase4B-ControlNet](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-ControlNet)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-ControlNet.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-ControlNet.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-ControlNet.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-ControlNet.py)|-|-|
64
+ |[DiffSynth-Studio/Template-KleinBase4B-Edit](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Edit)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Edit.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Edit.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Edit.py)|-|-|
65
+ |[DiffSynth-Studio/Template-KleinBase4B-Inpaint](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Inpaint)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Inpaint.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Inpaint.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Inpaint.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Inpaint.py)|-|-|
66
+ |[DiffSynth-Studio/Template-KleinBase4B-PandaMeme](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-PandaMeme)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-PandaMeme.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-PandaMeme.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-PandaMeme.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-PandaMeme.py)|-|-|
67
+ |[DiffSynth-Studio/Template-KleinBase4B-Sharpness](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Sharpness)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Sharpness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Sharpness.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Sharpness.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Sharpness.py)|-|-|
68
+ |[DiffSynth-Studio/Template-KleinBase4B-SoftRGB](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-SoftRGB)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-SoftRGB.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-SoftRGB.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-SoftRGB.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-SoftRGB.py)|-|-|
69
+ |[DiffSynth-Studio/Template-KleinBase4B-Upscaler](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-Upscaler)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-Upscaler.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-Upscaler.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-Upscaler.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-Upscaler.py)|-|-|
70
+ |[DiffSynth-Studio/Template-KleinBase4B-ContentRef](https://www.modelscope.cn/models/DiffSynth-Studio/Template-KleinBase4B-ContentRef)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/Template-KleinBase4B-ContentRef.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/Template-KleinBase4B-ContentRef.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/Template-KleinBase4B-ContentRef.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/Template-KleinBase4B-ContentRef.py)|-|-|
71
+ |[DiffSynth-Studio/KleinBase4B-i2L-v2](https://www.modelscope.cn/models/DiffSynth-Studio/KleinBase4B-i2L-v2)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference/KleinBase4B-i2L-v2.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_inference_low_vram/KleinBase4B-i2L-v2.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/full/KleinBase4B-i2L-v2.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/validate_full/KleinBase4B-i2L-v2.py)|-|-|
72
+
73
+ Special Training Scripts:
74
+
75
+ * Differential LoRA Training: [doc](../Training/Differential_LoRA.md)
76
+ * FP8 Precision Training: [doc](../Training/FP8_Precision.md)
77
+ * Two-stage Split Training: [doc](../Training/Split_Training.md)
78
+ * End-to-end Direct Distillation: [doc](../Training/Direct_Distill.md)
79
+
80
+ ## Model Inference
81
+
82
+ Models are loaded via `Flux2ImagePipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models).
83
+
84
+ Input parameters for `Flux2ImagePipeline` inference include:
85
+
86
+ * `prompt`: Prompt describing the content appearing in the image.
87
+ * `negative_prompt`: Negative prompt describing content that should not appear in the image, default value is `""`.
88
+ * `cfg_scale`: Classifier-free guidance parameter, default value is 1. When set to a value greater than 1, CFG is enabled.
89
+ * `height`: Image height, must be a multiple of 16.
90
+ * `width`: Image width, must be a multiple of 16.
91
+ * `seed`: Random seed. Default is `None`, meaning completely random.
92
+ * `rand_device`: Computing device for generating random Gaussian noise matrix, default is `"cpu"`. When set to `cuda`, different GPUs will produce different generation results.
93
+ * `num_inference_steps`: Number of inference steps, default value is 30.
94
+ * `embedded_guidance`: Embedded guidance parameter, default value is 3.5.
95
+ * `t5_sequence_length`: Sequence length of the T5 text encoder, default is 512.
96
+ * `tiled`: Whether to enable VAE tiling inference, default is `False`. Setting to `True` can significantly reduce VRAM usage during VAE encoding/decoding stages, producing slight errors and slightly longer inference time.
97
+ * `tile_size`: Tile size during VAE encoding/decoding stages, default is 128, only effective when `tiled=True`.
98
+ * `tile_stride`: Tile stride during VAE encoding/decoding stages, default is 64, only effective when `tiled=True`, must be less than or equal to `tile_size`.
99
+ * `progress_bar_cmd`: Progress bar, default is `tqdm.tqdm`. Can be disabled by setting to `lambda x:x`.
100
+
101
+ If VRAM is insufficient, please enable [VRAM Management](../Pipeline_Usage/VRAM_management.md). We provide recommended low VRAM configurations for each model in the example code, see the table in the "Model Overview" section above.
102
+
103
+ ## Model Training
104
+
105
+ FLUX.2 series models are uniformly trained through [`examples/flux2/model_training/train.py`](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/flux2/model_training/train.py), and the script parameters include:
106
+
107
+ * General Training Parameters
108
+ * Dataset Basic Configuration
109
+ * `--dataset_base_path`: Root directory of the dataset.
110
+ * `--dataset_metadata_path`: Metadata file path of the dataset.
111
+ * `--dataset_repeat`: Number of times the dataset is repeated in each epoch.
112
+ * `--dataset_num_workers`: Number of processes for each DataLoader.
113
+ * `--data_file_keys`: Field names to be loaded from metadata, usually image or video file paths, separated by `,`.
114
+ * Model Loading Configuration
115
+ * `--model_paths`: Paths of models to be loaded. JSON format.
116
+ * `--model_id_with_origin_paths`: Model IDs with original paths, e.g., `"black-forest-labs/FLUX.2-dev:text_encoder/*.safetensors"`. Separated by commas.
117
+ * `--extra_inputs`: Extra input parameters required by the model Pipeline, e.g., `controlnet_inputs` when training ControlNet models, separated by `,`.
118
+ * `--fp8_models`: Models loaded in FP8 format, consistent with `--model_paths` or `--model_id_with_origin_paths` format. Currently only supports models whose parameters are not updated by gradients (no gradient backpropagation, or gradients only update their LoRA).
119
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
120
+ * Training Basic Configuration
121
+ * `--learning_rate`: Learning rate.
122
+ * `--num_epochs`: Number of epochs.
123
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
124
+ * `--find_unused_parameters`: Whether there are unused parameters in DDP training. Some models contain redundant parameters that do not participate in gradient calculation, and this setting needs to be enabled to avoid errors in multi-GPU training.
125
+ * `--weight_decay`: Weight decay size, see [torch.optim.AdamW](https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html).
126
+ * `--task`: Training task, default is `sft`. Some models support more training modes, please refer to the documentation of each specific model.
127
+ * Output Configuration
128
+ * `--output_path`: Model saving path.
129
+ * `--remove_prefix_in_ckpt`: Remove prefix in the state dict of the model file.
130
+ * `--save_steps`: Interval of training steps to save the model. If this parameter is left blank, the model is saved once per epoch.
131
+ * LoRA Configuration
132
+ * `--lora_base_model`: Which model to add LoRA to.
133
+ * `--lora_target_modules`: Which layers to add LoRA to.
134
+ * `--lora_rank`: Rank of LoRA.
135
+ * `--lora_checkpoint`: Path of the LoRA checkpoint. If this path is provided, LoRA will be loaded from this checkpoint.
136
+ * `--preset_lora_path`: Preset LoRA checkpoint path. If this path is provided, this LoRA will be loaded in the form of being merged into the base model. This parameter is used for LoRA differential training.
137
+ * `--preset_lora_model`: Model that the preset LoRA is merged into, e.g., `dit`.
138
+ * Gradient Configuration
139
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
140
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to memory.
141
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
142
+ * Image Width/Height Configuration (Applicable to Image Generation and Video Generation Models)
143
+ * `--height`: Height of image or video. Leave `height` and `width` blank to enable dynamic resolution.
144
+ * `--width`: Width of image or video. Leave `height` and `width` blank to enable dynamic resolution.
145
+ * `--max_pixels`: Maximum pixel area of image or video frames. When dynamic resolution is enabled, images with resolution larger than this value will be downscaled, and images with resolution smaller than this value will remain unchanged.
146
+ * FLUX.2 Specific Parameters
147
+ * `--tokenizer_path`: Path of the tokenizer, applicable to text-to-image models, leave blank to automatically download from remote.
148
+
149
+ We have built a sample image dataset for your testing. You can download this dataset with the following command:
150
+
151
+ ```shell
152
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
153
+ ```
154
+
155
+ We have written recommended training scripts for each model, please refer to the table in the "Model Overview" section above. For how to write model training scripts, please refer to [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, please refer to [Training Framework Detailed Explanation](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).
docs/en/Model_Details/HiDream-O1-Image.md ADDED
@@ -0,0 +1,143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # HiDream-O1-Image
2
+
3
+ HiDream-O1-Image is an image generation model open-sourced by HiDream.ai, based on the Pixel-Level Unified Transformer (UiT) architecture. This model unifies VAE, DiT, and TextEncoder within a single Qwen3VLModel, performing diffusion denoising directly in pixel patch space without requiring a separate VAE component.
4
+
5
+ ## Installation
6
+
7
+ Before performing model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more information on installation, please refer to [Setup Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ Running the following code will quickly load the [HiDream-ai/HiDream-O1-Image](https://modelscope.cn/HiDream-ai/HiDream-O1-Image) model for inference. VRAM management is enabled, the framework automatically controls parameter loading based on available VRAM, requiring a minimum of 3GB VRAM.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.hidream_o1_image import HiDreamO1ImagePipeline
23
+ from diffsynth.core.loader.config import ModelConfig
24
+ import torch
25
+
26
+
27
+ vram_config = {
28
+ "offload_dtype": torch.bfloat16,
29
+ "offload_device": "cpu",
30
+ "onload_dtype": torch.bfloat16,
31
+ "onload_device": "cpu",
32
+ "preparing_dtype": torch.bfloat16,
33
+ "preparing_device": "cuda",
34
+ "computation_dtype": torch.bfloat16,
35
+ "computation_device": "cuda",
36
+ }
37
+
38
+
39
+ pipe = HiDreamO1ImagePipeline.from_pretrained(
40
+ torch_dtype=torch.bfloat16,
41
+ device="cuda",
42
+ model_configs=[
43
+ ModelConfig(model_id="HiDream-ai/HiDream-O1-Image", origin_file_pattern="model-*.safetensors", **vram_config),
44
+ ],
45
+ processor_config=ModelConfig(model_id="HiDream-ai/HiDream-O1-Image", origin_file_pattern="./"),
46
+ vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 0.5,
47
+ )
48
+ image = pipe(
49
+ prompt="medium shot, eye-level, front view. A woman is seated in an ornate bedroom, illuminated by candlelight, with a calm and composed expression. The subject is a young woman with fair skin, light brown hair styled in an updo with loose tendrils framing her face, and blue eyes. She wears a cream-colored satin robe with delicate floral embroidery and lace trim along the neckline. Her ears are adorned with pearl drop earrings. She is seated on a bed with a dark, intricately carved wooden headboard. To her left, a wooden nightstand holds three lit white candles and a candelabra with multiple lit candles in the background. The bed is covered with patterned pillows and a dark, textured blanket. The walls are paneled with dark wood and feature a large, ornate tapestry with muted earth tones. The lighting creates soft highlights on her face and robe, with warm shadows cast across the room.",
50
+ negative_prompt=" ",
51
+ cfg_scale=4.0,
52
+ height=2048,
53
+ width=2048,
54
+ seed=42,
55
+ num_inference_steps=50,
56
+ )
57
+ image.save("image.jpg")
58
+ ```
59
+
60
+ ## Model Overview
61
+
62
+ |Model ID|Inference|Low VRAM Inference|Full Training|Full Training Validation|LoRA Training|LoRA Training Validation|
63
+ |-|-|-|-|-|-|-|
64
+ |[HiDream-ai/HiDream-O1-Image](https://modelscope.cn/HiDream-ai/HiDream-O1-Image)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_inference/HiDream-O1-Image.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_inference_low_vram/HiDream-O1-Image.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/full/HiDream-O1-Image.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/validate_full/HiDream-O1-Image.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/lora/HiDream-O1-Image.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/validate_lora/HiDream-O1-Image.py)|
65
+ |[HiDream-ai/HiDream-O1-Image-Dev](https://modelscope.cn/HiDream-ai/HiDream-O1-Image-Dev)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_inference/HiDream-O1-Image-Dev.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_inference_low_vram/HiDream-O1-Image-Dev.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/full/HiDream-O1-Image-Dev.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/validate_full/HiDream-O1-Image-Dev.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/lora/HiDream-O1-Image-Dev.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/validate_lora/HiDream-O1-Image-Dev.py)|
66
+ |[DiffSynth-Studio/HidreamO1-i2L-v2](https://www.modelscope.cn/models/DiffSynth-Studio/HidreamO1-i2L-v2)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_inference/HidreamO1-i2L-v2.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_inference_low_vram/HidreamO1-i2L-v2.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/full/HidreamO1-i2L-v2.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/hidream_o1_image/model_training/validate_full/HidreamO1-i2L-v2.py)|-|-|
67
+
68
+ ## Model Inference
69
+
70
+ The model is loaded via `HiDreamO1ImagePipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models) for details.
71
+
72
+ The input parameters for `HiDreamO1ImagePipeline` inference include:
73
+
74
+ * `prompt`: Text prompt.
75
+ * `negative_prompt`: Negative prompt, defaults to `" "`.
76
+ * `cfg_scale`: Classifier-Free Guidance scale, defaults to 4.0. For the Dev model, it is recommended to set to 1.0.
77
+ * `height`: Output image height, defaults to 2048.
78
+ * `width`: Output image width, defaults to 2048.
79
+ * `seed`: Random seed, defaults to random.
80
+ * `rand_device`: Noise generation device, defaults to `"cpu"`.
81
+ * `num_inference_steps`: Number of inference steps, defaults to 50 for Full model and 28 for Dev model.
82
+ * `model_type`: Model type, `"full"` for Full model, `"dev"` for distilled Dev model.
83
+ * `shift`: Timestep shift parameter affecting sigma computation, defaults to 3.0.
84
+ * `noise_scale`: Noise scaling factor, defaults to 8.0. For the Dev model, it is recommended to set to 7.5.
85
+ * `edit_image`: List of reference images for image editing. Defaults to None (text-to-image mode).
86
+ * `keep_original_aspect`: Whether to preserve the original aspect ratio of reference images, defaults to True.
87
+
88
+ > **VRAM Note**: HiDream-O1-Image has a large parameter count (~8B). When generating 2048x2048 images, it is recommended to enable VRAM management (vram_config) or use the low VRAM inference scripts.
89
+
90
+ ## Model Training
91
+
92
+ Models in the hidream_o1_image series are trained uniformly via `examples/hidream_o1_image/model_training/train.py`. The script parameters include:
93
+
94
+ * General Training Parameters
95
+ * Dataset Configuration
96
+ * `--dataset_base_path`: Root directory of the dataset.
97
+ * `--dataset_metadata_path`: Path to the dataset metadata file.
98
+ * `--dataset_repeat`: Number of dataset repeats per epoch.
99
+ * `--dataset_num_workers`: Number of processes per DataLoader.
100
+ * `--data_file_keys`: Field names to load from metadata, typically paths to image or video files, separated by `,`.
101
+ * Model Loading Configuration
102
+ * `--model_paths`: Paths to load models from, in JSON format.
103
+ * `--model_id_with_origin_paths`: Model IDs with original paths, separated by commas.
104
+ * `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
105
+ * `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
106
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
107
+ * Basic Training Configuration
108
+ * `--learning_rate`: Learning rate.
109
+ * `--num_epochs`: Number of epochs.
110
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
111
+ * `--find_unused_parameters`: Whether unused parameters exist in DDP training.
112
+ * `--weight_decay`: Weight decay magnitude.
113
+ * `--task`: Training task, defaults to `sft`.
114
+ * Output Configuration
115
+ * `--output_path`: Path to save the model.
116
+ * `--remove_prefix_in_ckpt`: Remove prefix in the model's state dict.
117
+ * `--save_steps`: Interval in training steps to save the model.
118
+ * LoRA Configuration
119
+ * `--lora_base_model`: Which model to add LoRA to.
120
+ * `--lora_target_modules`: Which layers to add LoRA to.
121
+ * `--lora_rank`: Rank of LoRA.
122
+ * `--lora_checkpoint`: Path to LoRA checkpoint.
123
+ * `--preset_lora_path`: Path to preset LoRA checkpoint for LoRA differential training.
124
+ * `--preset_lora_model`: Which model to integrate preset LoRA into, e.g., `dit`.
125
+ * Gradient Configuration
126
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
127
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to CPU memory.
128
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
129
+ * Resolution Configuration
130
+ * `--height`: Height of the image/video. Leave empty to enable dynamic resolution.
131
+ * `--width`: Width of the image/video. Leave empty to enable dynamic resolution.
132
+ * `--max_pixels`: Maximum pixel area, images larger than this will be scaled down during dynamic resolution.
133
+ * `--num_frames`: Number of frames for video (video generation models only).
134
+ * HiDream-O1-Image Specific Parameters
135
+ * `--processor_config`: Path to the processor configuration file, used for loading AutoProcessor for text tokenization.
136
+ * `--noise_scale`: Noise scaling factor, defaults to 8.0.
137
+ * `--initialize_model_on_cpu`: Whether to initialize the model on CPU, which can help reduce peak GPU VRAM usage.
138
+
139
+ ```shell
140
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
141
+ ```
142
+
143
+ We provide recommended training scripts for each model, please refer to the table in "Model Overview" above. For guidance on writing model training scripts, see [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, see [Training Framework Overview](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).
docs/en/Model_Details/Ideogram-4.md ADDED
@@ -0,0 +1,151 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Ideogram 4
2
+
3
+ Ideogram 4 is an image generation model open-sourced by Ideogram. DiffSynth-Studio supports inference, low VRAM inference, full training, and LoRA training for both the FP8 quantized version and the BF16 repackaged version.
4
+
5
+ ## Installation
6
+
7
+ Before performing model inference and training, please install DiffSynth-Studio first.
8
+
9
+ ```shell
10
+ git clone https://github.com/modelscope/DiffSynth-Studio.git
11
+ cd DiffSynth-Studio
12
+ pip install -e .
13
+ ```
14
+
15
+ For more information about installation, please refer to [Install Dependencies](../Pipeline_Usage/Setup.md).
16
+
17
+ ## Quick Start
18
+
19
+ Running the following code will load the [ideogram-ai/ideogram-4-fp8](https://www.modelscope.cn/models/ideogram-ai/ideogram-4-fp8) model for inference. A minimum of 24GB VRAM is required to run.
20
+
21
+ ```python
22
+ from diffsynth.pipelines.ideogram4 import Ideogram4Pipeline
23
+ from diffsynth.core import ModelConfig
24
+ import torch
25
+
26
+
27
+ pipe = Ideogram4Pipeline.from_pretrained(
28
+ torch_dtype=torch.bfloat16,
29
+ device="cuda",
30
+ model_configs=[
31
+ ModelConfig(model_id="ideogram-ai/ideogram-4-fp8", origin_file_pattern="transformer/diffusion_pytorch_model.safetensors"),
32
+ # unconditional_transformer is optional. You can delete this line to reduce VRAM required.
33
+ ModelConfig(model_id="ideogram-ai/ideogram-4-fp8", origin_file_pattern="unconditional_transformer/diffusion_pytorch_model.safetensors"),
34
+ ModelConfig(model_id="ideogram-ai/ideogram-4-fp8", origin_file_pattern="text_encoder/model.safetensors"),
35
+ ModelConfig(model_id="ideogram-ai/ideogram-4-fp8", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"),
36
+ ],
37
+ tokenizer_config=ModelConfig(model_id="ideogram-ai/ideogram-4-fp8", origin_file_pattern="tokenizer/"),
38
+ )
39
+ prompt = r"""
40
+ {
41
+ "high_level_description": "A medium-shot photograph of Formula 1 driver Max Verstappen wearing his Red Bull Racing racing suit and cap, smiling as he holds his racing helmet and talks to a man in a white shirt and black vest at a race track.",
42
+ "style_description": {
43
+ "aesthetics": "saturated primary colors, rule of thirds, joyful and triumphant",
44
+ "lighting": "overcast daylight, diffused, soft subtle shadows",
45
+ "photo": "shallow depth of field, sharp focus, eye-level, telephoto",
46
+ "medium": "photograph"
47
+ },
48
+ "compositional_deconstruction": {
49
+ "background": "The background is an out-of-focus racing paddock or track environment. Several blurred figures are visible, including one in an orange shirt. A purple and white structure with a red 'F1' logo stands on the left. The scene is outdoors with daylight, though the sky is not visible.",
50
+ "elements": [
51
+ {"type": "obj", "bbox": [55, 642, 1000, 937], "desc": "An older man standing in profile, facing left toward Max Verstappen. He has grey hair and fair skin. He is wearing a white long-sleeved button-down shirt with a navy blue quilted vest over it. He has a slight smile."},
52
+ {"type": "obj", "bbox": [34, 137, 1000, 617], "desc": "Max Verstappen, a fair-skinned male Formula 1 driver, positioned in the center. He is facing forward with a joyful expression and a slight smile. He wears a navy blue Red Bull Racing team uniform with numerous sponsor logos and a matching baseball cap with the number '1'. He is holding a white and red racing helmet in his hands. He has a silver watch on his left wrist."},
53
+ {"type": "obj", "bbox": [422, 212, 792, 452], "desc": "Max Verstappen's racing helmet, held in front of his chest. It features a white, red, and yellow design with the Red Bull logo and the 'Player 0.0' branding. The visor is clear and open."},
54
+ {"type": "text", "bbox": [657, 0, 755, 142], "text": "F1", "desc": "Large, stylized red logo on a black and purple background in the lower left."},
55
+ {"type": "text", "bbox": [768, 0, 818, 147], "text": "Formula 1\nWorld Championship™", "desc": "Small white sans-serif text below the F1 logo on the left side."},
56
+ {"type": "text", "bbox": [78, 447, 117, 510], "text": "ORACLE\nRed Bull\nRacing", "desc": "Very small white and orange logo on the front of the navy blue cap."},
57
+ {"type": "text", "bbox": [78, 417, 120, 440], "text": "1", "desc": "Bold red numeral '1' on the front left side of the navy blue cap."},
58
+ {"type": "text", "bbox": [332, 442, 363, 483], "text": "Red Bull", "desc": "Small yellow and red text logo on the collar of the uniform."},
59
+ {"type": "text", "bbox": [373, 490, 423, 532], "text": "RAUCH", "desc": "Small yellow and blue logo on the right chest of the uniform."},
60
+ {"type": "text", "bbox": [422, 473, 500, 532], "text": "BYBIT\nHONDA", "desc": "Medium-sized white sans-serif text on the right chest of the uniform."},
61
+ {"type": "text", "bbox": [410, 203, 442, 257], "text": "RAUCH", "desc": "Small yellow logo on the left upper arm of the uniform."},
62
+ {"type": "text", "bbox": [530, 448, 627, 510], "text": "Red Bull", "desc": "Medium red text logo on the right side of the torso, part of the Red Bull graphic."},
63
+ {"type": "text", "bbox": [680, 417, 768, 523], "text": "Red Bull", "desc": "Large red text logo across the lower torso of the uniform."},
64
+ {"type": "text", "bbox": [797, 475, 815, 518], "text": "MAX", "desc": "Small white text next to a Dutch flag on the belt area of the uniform."},
65
+ {"type": "text", "bbox": [558, 317, 715, 355], "text": "Player 0.0", "desc": "Black sans-serif text on a white band on the racing helmet."},
66
+ {"type": "text", "bbox": [560, 800, 582, 835], "text": "IA.COM", "desc": "Small blue sans-serif text on the right sleeve of the white shirt."},
67
+ {"type": "text", "bbox": [968, 8, 997, 332], "text": "© Anadolu Agency via Getty Images", "desc": "Small white watermark text in the bottom left corner."}
68
+ ]
69
+ }
70
+ }
71
+ """
72
+ image = pipe(prompt=prompt, height=1024, width=1024, num_inference_steps=48, cfg_scale=7.0, seed=42)
73
+ image.save("image_ideogram-4-fp8.jpg")
74
+ ```
75
+
76
+ ## Model Overview
77
+
78
+ |Model ID|Inference|Low VRAM Inference|Full Training|Full Training Validation|LoRA Training|LoRA Training Validation|
79
+ |-|-|-|-|-|-|-|
80
+ |[ideogram-ai/ideogram-4-fp8](https://www.modelscope.cn/models/ideogram-ai/ideogram-4-fp8)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ideogram4/model_inference/ideogram-4-fp8.py)|-|-|-|-|-|
81
+ |[DiffSynth-Studio/ideogram-4-bf16-repackage](https://www.modelscope.cn/models/DiffSynth-Studio/ideogram-4-bf16-repackage)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ideogram4/model_inference/ideogram-4-bf16-repackage.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ideogram4/model_inference_low_vram/ideogram-4-bf16-repackage.py)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ideogram4/model_training/full/Ideogram-4-bf16-repackage.sh)|-|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ideogram4/model_training/lora/Ideogram-4-bf16-repackage.sh)|[code](https://github.com/modelscope/DiffSynth-Studio/blob/main/examples/ideogram4/model_training/validate_lora/Ideogram-4-bf16-repackage.py)|
82
+
83
+ ## Model Inference
84
+
85
+ The model is loaded via `Ideogram4Pipeline.from_pretrained`, see [Loading Models](../Pipeline_Usage/Model_Inference.md#loading-models) for details.
86
+
87
+ The input parameters for `Ideogram4Pipeline` inference include:
88
+
89
+ * `prompt`: Prompt describing the content appearing in the image. Ideogram 4 supports structured JSON format prompts, including high-level description, style description, and compositional deconstruction.
90
+ * `negative_prompt`: Negative prompt describing content that should not appear in the image, default value is `""`.
91
+ * `cfg_scale`: Classifier-free guidance parameter, default value is 7.0.
92
+ * `input_image`: Input image for image-to-image generation, used in conjunction with `denoising_strength`.
93
+ * `denoising_strength`: Denoising strength, range is 0~1, default value is 1. When the value approaches 0, the generated image is similar to the input image; when the value approaches 1, the generated image differs more from the input image. When `input_image` parameter is not provided, do not set this to a non-1 value.
94
+ * `height`: Image height, must be a multiple of 16, default value is 1024.
95
+ * `width`: Image width, must be a multiple of 16, default value is 1024.
96
+ * `seed`: Random seed. Default is `None`, meaning completely random.
97
+ * `rand_device`: Computing device for generating random Gaussian noise matrix, default is `"cpu"`.
98
+ * `num_inference_steps`: Number of inference steps, default value is 50.
99
+
100
+ ## Model Training
101
+
102
+ Models in the ideogram4 series are trained uniformly via `examples/ideogram4/model_training/train.py`. The script parameters include:
103
+
104
+ * General Training Parameters
105
+ * Dataset Configuration
106
+ * `--dataset_base_path`: Root directory of the dataset.
107
+ * `--dataset_metadata_path`: Path to the dataset metadata file.
108
+ * `--dataset_repeat`: Number of dataset repeats per epoch.
109
+ * `--dataset_num_workers`: Number of processes per DataLoader.
110
+ * `--data_file_keys`: Field names to load from metadata, typically paths to image or video files, separated by `,`.
111
+ * Model Loading Configuration
112
+ * `--model_paths`: Paths to load models from, in JSON format.
113
+ * `--model_id_with_origin_paths`: Model IDs with original paths, separated by commas.
114
+ * `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
115
+ * `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
116
+ * `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
117
+ * Basic Training Configuration
118
+ * `--learning_rate`: Learning rate.
119
+ * `--num_epochs`: Number of epochs.
120
+ * `--trainable_models`: Trainable models, e.g., `dit`, `vae`, `text_encoder`.
121
+ * `--find_unused_parameters`: Whether unused parameters exist in DDP training.
122
+ * `--weight_decay`: Weight decay magnitude.
123
+ * `--task`: Training task, defaults to `sft`.
124
+ * Output Configuration
125
+ * `--output_path`: Path to save the model.
126
+ * `--remove_prefix_in_ckpt`: Remove prefix in the model's state dict.
127
+ * `--save_steps`: Interval in training steps to save the model.
128
+ * LoRA Configuration
129
+ * `--lora_base_model`: Which model to add LoRA to.
130
+ * `--lora_target_modules`: Which layers to add LoRA to.
131
+ * `--lora_rank`: Rank of LoRA.
132
+ * `--lora_checkpoint`: Path to LoRA checkpoint.
133
+ * `--preset_lora_path`: Path to preset LoRA checkpoint for LoRA differential training.
134
+ * `--preset_lora_model`: Which model to integrate preset LoRA into, e.g., `dit`.
135
+ * Gradient Configuration
136
+ * `--use_gradient_checkpointing`: Whether to enable gradient checkpointing.
137
+ * `--use_gradient_checkpointing_offload`: Whether to offload gradient checkpointing to CPU memory.
138
+ * `--gradient_accumulation_steps`: Number of gradient accumulation steps.
139
+ * Resolution Configuration
140
+ * `--height`: Height of the image/video. Leave empty to enable dynamic resolution.
141
+ * `--width`: Width of the image/video. Leave empty to enable dynamic resolution.
142
+ * `--max_pixels`: Maximum pixel area, images larger than this will be scaled down during dynamic resolution.
143
+ * `--num_frames`: Number of frames for video (video generation models only).
144
+ * Ideogram-4 Specific Parameters
145
+ * `--tokenizer_path`: Path to tokenizer. Defaults to downloading from `ideogram-ai/ideogram-4-fp8`.
146
+
147
+ ```shell
148
+ modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset
149
+ ```
150
+
151
+ We provide recommended training scripts for each model, please refer to the table in "Model Overview" above. For guidance on writing model training scripts, see [Model Training](../Pipeline_Usage/Model_Training.md); for more advanced training algorithms, see [Training Framework Overview](https://github.com/modelscope/DiffSynth-Studio/tree/main/docs/en/Training/).