fffiloni commited on
Commit
ef931f5
·
verified ·
1 Parent(s): d545e51

Migrated files batch 31

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +2 -0
  2. diffsynth/pipelines/pipeline_runner.py +105 -0
  3. diffsynth/pipelines/qwen_image.py +364 -0
  4. diffsynth/pipelines/sd3_image.py +147 -0
  5. diffsynth/pipelines/sd_image.py +191 -0
  6. diffsynth/pipelines/sd_video.py +269 -0
  7. diffsynth/pipelines/sdxl_image.py +226 -0
  8. diffsynth/pipelines/sdxl_video.py +226 -0
  9. diffsynth/pipelines/step_video.py +209 -0
  10. diffsynth/pipelines/svd_video.py +300 -0
  11. diffsynth/pipelines/wan_video.py +626 -0
  12. diffsynth/pipelines/wan_video_new.py +1125 -0
  13. diffsynth/pipelines/wan_video_relit_live.py +525 -0
  14. diffsynth/processors/FastBlend.py +142 -0
  15. diffsynth/processors/PILEditor.py +28 -0
  16. diffsynth/processors/RIFE.py +77 -0
  17. diffsynth/processors/__init__.py +0 -0
  18. diffsynth/processors/base.py +6 -0
  19. diffsynth/processors/sequencial_processor.py +41 -0
  20. diffsynth/prompters/__init__.py +12 -0
  21. diffsynth/prompters/base_prompter.py +70 -0
  22. diffsynth/prompters/cog_prompter.py +46 -0
  23. diffsynth/prompters/flux_prompter.py +74 -0
  24. diffsynth/prompters/hunyuan_dit_prompter.py +69 -0
  25. diffsynth/prompters/hunyuan_video_prompter.py +275 -0
  26. diffsynth/prompters/kolors_prompter.py +354 -0
  27. diffsynth/prompters/omnigen_prompter.py +356 -0
  28. diffsynth/prompters/omost.py +323 -0
  29. diffsynth/prompters/prompt_refiners.py +130 -0
  30. diffsynth/prompters/sd3_prompter.py +93 -0
  31. diffsynth/prompters/sd_prompter.py +73 -0
  32. diffsynth/prompters/sdxl_prompter.py +61 -0
  33. diffsynth/prompters/stepvideo_prompter.py +56 -0
  34. diffsynth/prompters/wan_prompter.py +109 -0
  35. diffsynth/schedulers/__init__.py +4 -0
  36. diffsynth/schedulers/continuous_ode.py +59 -0
  37. diffsynth/schedulers/ddim.py +105 -0
  38. diffsynth/schedulers/flow_match.py +120 -0
  39. diffsynth/schedulers/flow_match_plf.py +150 -0
  40. diffsynth/tokenizer_configs/__init__.py +0 -0
  41. diffsynth/tokenizer_configs/cog/tokenizer/added_tokens.json +102 -0
  42. diffsynth/tokenizer_configs/cog/tokenizer/special_tokens_map.json +125 -0
  43. diffsynth/tokenizer_configs/cog/tokenizer/spiece.model +3 -0
  44. diffsynth/tokenizer_configs/cog/tokenizer/tokenizer_config.json +940 -0
  45. diffsynth/tokenizer_configs/flux/tokenizer_1/merges.txt +0 -0
  46. diffsynth/tokenizer_configs/flux/tokenizer_1/special_tokens_map.json +30 -0
  47. diffsynth/tokenizer_configs/flux/tokenizer_1/tokenizer_config.json +30 -0
  48. diffsynth/tokenizer_configs/flux/tokenizer_1/vocab.json +0 -0
  49. diffsynth/tokenizer_configs/flux/tokenizer_2/special_tokens_map.json +125 -0
  50. diffsynth/tokenizer_configs/flux/tokenizer_2/spiece.model +3 -0
.gitattributes CHANGED
@@ -627,3 +627,5 @@ datasets/gradio_data/assets/images_demo/demo6.jpg filter=lfs diff=lfs merge=lfs
627
  datasets/gradio_data/assets/images_demo/demo7.jpg filter=lfs diff=lfs merge=lfs -text
628
  datasets/gradio_data/assets/images_demo/demo8.jpg filter=lfs diff=lfs merge=lfs -text
629
  datasets/gradio_data/assets/images_demo/demo9.jpg filter=lfs diff=lfs merge=lfs -text
 
 
 
627
  datasets/gradio_data/assets/images_demo/demo7.jpg filter=lfs diff=lfs merge=lfs -text
628
  datasets/gradio_data/assets/images_demo/demo8.jpg filter=lfs diff=lfs merge=lfs -text
629
  datasets/gradio_data/assets/images_demo/demo9.jpg filter=lfs diff=lfs merge=lfs -text
630
+ diffsynth/tokenizer_configs/hunyuan_video/tokenizer_2/tokenizer.json filter=lfs diff=lfs merge=lfs -text
631
+ diffsynth/tokenizer_configs/kolors/tokenizer/vocab.txt filter=lfs diff=lfs merge=lfs -text
diffsynth/pipelines/pipeline_runner.py ADDED
@@ -0,0 +1,105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os, torch, json
2
+ from .sd_video import ModelManager, SDVideoPipeline, ControlNetConfigUnit
3
+ from ..processors.sequencial_processor import SequencialProcessor
4
+ from ..data import VideoData, save_frames, save_video
5
+
6
+
7
+
8
+ class SDVideoPipelineRunner:
9
+ def __init__(self, in_streamlit=False):
10
+ self.in_streamlit = in_streamlit
11
+
12
+
13
+ def load_pipeline(self, model_list, textual_inversion_folder, device, lora_alphas, controlnet_units):
14
+ # Load models
15
+ model_manager = ModelManager(torch_dtype=torch.float16, device=device)
16
+ model_manager.load_models(model_list)
17
+ pipe = SDVideoPipeline.from_model_manager(
18
+ model_manager,
19
+ [
20
+ ControlNetConfigUnit(
21
+ processor_id=unit["processor_id"],
22
+ model_path=unit["model_path"],
23
+ scale=unit["scale"]
24
+ ) for unit in controlnet_units
25
+ ]
26
+ )
27
+ textual_inversion_paths = []
28
+ for file_name in os.listdir(textual_inversion_folder):
29
+ if file_name.endswith(".pt") or file_name.endswith(".bin") or file_name.endswith(".pth") or file_name.endswith(".safetensors"):
30
+ textual_inversion_paths.append(os.path.join(textual_inversion_folder, file_name))
31
+ pipe.prompter.load_textual_inversions(textual_inversion_paths)
32
+ return model_manager, pipe
33
+
34
+
35
+ def load_smoother(self, model_manager, smoother_configs):
36
+ smoother = SequencialProcessor.from_model_manager(model_manager, smoother_configs)
37
+ return smoother
38
+
39
+
40
+ def synthesize_video(self, model_manager, pipe, seed, smoother, **pipeline_inputs):
41
+ torch.manual_seed(seed)
42
+ if self.in_streamlit:
43
+ import streamlit as st
44
+ progress_bar_st = st.progress(0.0)
45
+ output_video = pipe(**pipeline_inputs, smoother=smoother, progress_bar_st=progress_bar_st)
46
+ progress_bar_st.progress(1.0)
47
+ else:
48
+ output_video = pipe(**pipeline_inputs, smoother=smoother)
49
+ model_manager.to("cpu")
50
+ return output_video
51
+
52
+
53
+ def load_video(self, video_file, image_folder, height, width, start_frame_id, end_frame_id):
54
+ video = VideoData(video_file=video_file, image_folder=image_folder, height=height, width=width)
55
+ if start_frame_id is None:
56
+ start_frame_id = 0
57
+ if end_frame_id is None:
58
+ end_frame_id = len(video)
59
+ frames = [video[i] for i in range(start_frame_id, end_frame_id)]
60
+ return frames
61
+
62
+
63
+ def add_data_to_pipeline_inputs(self, data, pipeline_inputs):
64
+ pipeline_inputs["input_frames"] = self.load_video(**data["input_frames"])
65
+ pipeline_inputs["num_frames"] = len(pipeline_inputs["input_frames"])
66
+ pipeline_inputs["width"], pipeline_inputs["height"] = pipeline_inputs["input_frames"][0].size
67
+ if len(data["controlnet_frames"]) > 0:
68
+ pipeline_inputs["controlnet_frames"] = [self.load_video(**unit) for unit in data["controlnet_frames"]]
69
+ return pipeline_inputs
70
+
71
+
72
+ def save_output(self, video, output_folder, fps, config):
73
+ os.makedirs(output_folder, exist_ok=True)
74
+ save_frames(video, os.path.join(output_folder, "frames"))
75
+ save_video(video, os.path.join(output_folder, "video.mp4"), fps=fps)
76
+ config["pipeline"]["pipeline_inputs"]["input_frames"] = []
77
+ config["pipeline"]["pipeline_inputs"]["controlnet_frames"] = []
78
+ with open(os.path.join(output_folder, "config.json"), 'w') as file:
79
+ json.dump(config, file, indent=4)
80
+
81
+
82
+ def run(self, config):
83
+ if self.in_streamlit:
84
+ import streamlit as st
85
+ if self.in_streamlit: st.markdown("Loading videos ...")
86
+ config["pipeline"]["pipeline_inputs"] = self.add_data_to_pipeline_inputs(config["data"], config["pipeline"]["pipeline_inputs"])
87
+ if self.in_streamlit: st.markdown("Loading videos ... done!")
88
+ if self.in_streamlit: st.markdown("Loading models ...")
89
+ model_manager, pipe = self.load_pipeline(**config["models"])
90
+ if self.in_streamlit: st.markdown("Loading models ... done!")
91
+ if "smoother_configs" in config:
92
+ if self.in_streamlit: st.markdown("Loading smoother ...")
93
+ smoother = self.load_smoother(model_manager, config["smoother_configs"])
94
+ if self.in_streamlit: st.markdown("Loading smoother ... done!")
95
+ else:
96
+ smoother = None
97
+ if self.in_streamlit: st.markdown("Synthesizing videos ...")
98
+ output_video = self.synthesize_video(model_manager, pipe, config["pipeline"]["seed"], smoother, **config["pipeline"]["pipeline_inputs"])
99
+ if self.in_streamlit: st.markdown("Synthesizing videos ... done!")
100
+ if self.in_streamlit: st.markdown("Saving videos ...")
101
+ self.save_output(output_video, config["data"]["output_folder"], config["data"]["fps"], config)
102
+ if self.in_streamlit: st.markdown("Saving videos ... done!")
103
+ if self.in_streamlit: st.markdown("Finished!")
104
+ video_file = open(os.path.join(os.path.join(config["data"]["output_folder"], "video.mp4")), 'rb')
105
+ if self.in_streamlit: st.video(video_file.read())
diffsynth/pipelines/qwen_image.py ADDED
@@ -0,0 +1,364 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from PIL import Image
3
+ from typing import Union
4
+ from PIL import Image
5
+ from tqdm import tqdm
6
+ from einops import rearrange
7
+
8
+ from ..models import ModelManager, load_state_dict
9
+ from ..models.qwen_image_dit import QwenImageDiT
10
+ from ..models.qwen_image_text_encoder import QwenImageTextEncoder
11
+ from ..models.qwen_image_vae import QwenImageVAE
12
+ from ..schedulers import FlowMatchScheduler
13
+ from ..utils import BasePipeline, ModelConfig, PipelineUnitRunner, PipelineUnit
14
+ from ..lora import GeneralLoRALoader
15
+
16
+ from ..vram_management import gradient_checkpoint_forward, enable_vram_management, AutoWrappedModule, AutoWrappedLinear
17
+
18
+
19
+
20
+ class QwenImagePipeline(BasePipeline):
21
+
22
+ def __init__(self, device="cuda", torch_dtype=torch.bfloat16):
23
+ super().__init__(
24
+ device=device, torch_dtype=torch_dtype,
25
+ height_division_factor=16, width_division_factor=16,
26
+ )
27
+ from transformers import Qwen2Tokenizer
28
+
29
+ self.scheduler = FlowMatchScheduler(sigma_min=0, sigma_max=1, extra_one_step=True, exponential_shift=True, exponential_shift_mu=0.8, shift_terminal=0.02)
30
+ self.text_encoder: QwenImageTextEncoder = None
31
+ self.dit: QwenImageDiT = None
32
+ self.vae: QwenImageVAE = None
33
+ self.tokenizer: Qwen2Tokenizer = None
34
+ self.unit_runner = PipelineUnitRunner()
35
+ self.in_iteration_models = ("dit",)
36
+ self.units = [
37
+ QwenImageUnit_ShapeChecker(),
38
+ QwenImageUnit_NoiseInitializer(),
39
+ QwenImageUnit_InputImageEmbedder(),
40
+ QwenImageUnit_PromptEmbedder(),
41
+ ]
42
+ self.model_fn = model_fn_qwen_image
43
+
44
+
45
+ def load_lora(self, module, path, alpha=1):
46
+ loader = GeneralLoRALoader(torch_dtype=self.torch_dtype, device=self.device)
47
+ lora = load_state_dict(path, torch_dtype=self.torch_dtype, device=self.device)
48
+ loader.load(module, lora, alpha=alpha)
49
+
50
+
51
+ def training_loss(self, **inputs):
52
+ timestep_id = torch.randint(0, self.scheduler.num_train_timesteps, (1,))
53
+ timestep = self.scheduler.timesteps[timestep_id].to(dtype=self.torch_dtype, device=self.device)
54
+
55
+ inputs["latents"] = self.scheduler.add_noise(inputs["input_latents"], inputs["noise"], timestep)
56
+ training_target = self.scheduler.training_target(inputs["input_latents"], inputs["noise"], timestep)
57
+
58
+ noise_pred = self.model_fn(**inputs, timestep=timestep)
59
+
60
+ loss = torch.nn.functional.mse_loss(noise_pred.float(), training_target.float())
61
+ loss = loss * self.scheduler.training_weight(timestep)
62
+ return loss
63
+
64
+
65
+ def enable_vram_management(self, num_persistent_param_in_dit=None, vram_limit=None, vram_buffer=0.5):
66
+ self.vram_management_enabled = True
67
+ if num_persistent_param_in_dit is not None:
68
+ vram_limit = None
69
+ else:
70
+ if vram_limit is None:
71
+ vram_limit = self.get_vram()
72
+ vram_limit = vram_limit - vram_buffer
73
+ if self.text_encoder is not None:
74
+ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLRotaryEmbedding, Qwen2RMSNorm
75
+ dtype = next(iter(self.text_encoder.parameters())).dtype
76
+ enable_vram_management(
77
+ self.text_encoder,
78
+ module_map = {
79
+ torch.nn.Linear: AutoWrappedLinear,
80
+ torch.nn.Embedding: AutoWrappedModule,
81
+ Qwen2_5_VLRotaryEmbedding: AutoWrappedModule,
82
+ Qwen2RMSNorm: AutoWrappedModule,
83
+ },
84
+ module_config = dict(
85
+ offload_dtype=dtype,
86
+ offload_device="cpu",
87
+ onload_dtype=dtype,
88
+ onload_device="cpu",
89
+ computation_dtype=self.torch_dtype,
90
+ computation_device=self.device,
91
+ ),
92
+ vram_limit=vram_limit,
93
+ )
94
+ if self.dit is not None:
95
+ from ..models.qwen_image_dit import RMSNorm
96
+ dtype = next(iter(self.dit.parameters())).dtype
97
+ device = "cpu" if vram_limit is not None else self.device
98
+ enable_vram_management(
99
+ self.dit,
100
+ module_map = {
101
+ RMSNorm: AutoWrappedModule,
102
+ torch.nn.Linear: AutoWrappedLinear,
103
+ },
104
+ module_config = dict(
105
+ offload_dtype=dtype,
106
+ offload_device="cpu",
107
+ onload_dtype=dtype,
108
+ onload_device=device,
109
+ computation_dtype=self.torch_dtype,
110
+ computation_device=self.device,
111
+ ),
112
+ max_num_param=num_persistent_param_in_dit,
113
+ overflow_module_config = dict(
114
+ offload_dtype=dtype,
115
+ offload_device="cpu",
116
+ onload_dtype=dtype,
117
+ onload_device="cpu",
118
+ computation_dtype=self.torch_dtype,
119
+ computation_device=self.device,
120
+ ),
121
+ vram_limit=vram_limit,
122
+ )
123
+ if self.vae is not None:
124
+ from ..models.qwen_image_vae import QwenImageRMS_norm
125
+ dtype = next(iter(self.vae.parameters())).dtype
126
+ enable_vram_management(
127
+ self.vae,
128
+ module_map = {
129
+ torch.nn.Linear: AutoWrappedLinear,
130
+ torch.nn.Conv3d: AutoWrappedModule,
131
+ torch.nn.Conv2d: AutoWrappedModule,
132
+ QwenImageRMS_norm: AutoWrappedModule,
133
+ },
134
+ module_config = dict(
135
+ offload_dtype=dtype,
136
+ offload_device="cpu",
137
+ onload_dtype=dtype,
138
+ onload_device="cpu",
139
+ computation_dtype=self.torch_dtype,
140
+ computation_device=self.device,
141
+ ),
142
+ vram_limit=vram_limit,
143
+ )
144
+
145
+
146
+ @staticmethod
147
+ def from_pretrained(
148
+ torch_dtype: torch.dtype = torch.bfloat16,
149
+ device: Union[str, torch.device] = "cuda",
150
+ model_configs: list[ModelConfig] = [],
151
+ tokenizer_config: ModelConfig = ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="tokenizer/"),
152
+ ):
153
+ # Download and load models
154
+ model_manager = ModelManager()
155
+ for model_config in model_configs:
156
+ model_config.download_if_necessary()
157
+ model_manager.load_model(
158
+ model_config.path,
159
+ device=model_config.offload_device or device,
160
+ torch_dtype=model_config.offload_dtype or torch_dtype
161
+ )
162
+
163
+ # Initialize pipeline
164
+ pipe = QwenImagePipeline(device=device, torch_dtype=torch_dtype)
165
+ pipe.text_encoder = model_manager.fetch_model("qwen_image_text_encoder")
166
+ pipe.dit = model_manager.fetch_model("qwen_image_dit")
167
+ pipe.vae = model_manager.fetch_model("qwen_image_vae")
168
+ if tokenizer_config is not None and pipe.text_encoder is not None:
169
+ tokenizer_config.download_if_necessary()
170
+ from transformers import Qwen2Tokenizer
171
+ pipe.tokenizer = Qwen2Tokenizer.from_pretrained(tokenizer_config.path)
172
+ return pipe
173
+
174
+
175
+ @torch.no_grad()
176
+ def __call__(
177
+ self,
178
+ # Prompt
179
+ prompt: str,
180
+ negative_prompt: str = "",
181
+ cfg_scale: float = 4.0,
182
+ # Image
183
+ input_image: Image.Image = None,
184
+ denoising_strength: float = 1.0,
185
+ # Shape
186
+ height: int = 1328,
187
+ width: int = 1328,
188
+ # Randomness
189
+ seed: int = None,
190
+ rand_device: str = "cpu",
191
+ # Steps
192
+ num_inference_steps: int = 30,
193
+ # Tile
194
+ tiled: bool = False,
195
+ tile_size: int = 128,
196
+ tile_stride: int = 64,
197
+ # Progress bar
198
+ progress_bar_cmd = tqdm,
199
+ ):
200
+ # Scheduler
201
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, dynamic_shift_len=(height // 16) * (width // 16))
202
+
203
+ # Parameters
204
+ inputs_posi = {
205
+ "prompt": prompt,
206
+ }
207
+ inputs_nega = {
208
+ "negative_prompt": negative_prompt,
209
+ }
210
+ inputs_shared = {
211
+ "cfg_scale": cfg_scale,
212
+ "input_image": input_image, "denoising_strength": denoising_strength,
213
+ "height": height, "width": width,
214
+ "seed": seed, "rand_device": rand_device,
215
+ "tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride,
216
+ }
217
+ for unit in self.units:
218
+ inputs_shared, inputs_posi, inputs_nega = self.unit_runner(unit, self, inputs_shared, inputs_posi, inputs_nega)
219
+
220
+ # Denoise
221
+ self.load_models_to_device(self.in_iteration_models)
222
+ models = {name: getattr(self, name) for name in self.in_iteration_models}
223
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
224
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
225
+
226
+ # Inference
227
+ noise_pred_posi = self.model_fn(**models, **inputs_shared, **inputs_posi, timestep=timestep, progress_id=progress_id)
228
+ if cfg_scale != 1.0:
229
+ noise_pred_nega = self.model_fn(**models, **inputs_shared, **inputs_nega, timestep=timestep, progress_id=progress_id)
230
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
231
+ else:
232
+ noise_pred = noise_pred_posi
233
+
234
+ # Scheduler
235
+ inputs_shared["latents"] = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], inputs_shared["latents"])
236
+
237
+ # Decode
238
+ self.load_models_to_device(['vae'])
239
+ image = self.vae.decode(inputs_shared["latents"], device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
240
+ image = self.vae_output_to_image(image)
241
+ self.load_models_to_device([])
242
+
243
+ return image
244
+
245
+
246
+
247
+ class QwenImageUnit_ShapeChecker(PipelineUnit):
248
+ def __init__(self):
249
+ super().__init__(input_params=("height", "width"))
250
+
251
+ def process(self, pipe: QwenImagePipeline, height, width):
252
+ height, width = pipe.check_resize_height_width(height, width)
253
+ return {"height": height, "width": width}
254
+
255
+
256
+
257
+ class QwenImageUnit_NoiseInitializer(PipelineUnit):
258
+ def __init__(self):
259
+ super().__init__(input_params=("height", "width", "seed", "rand_device"))
260
+
261
+ def process(self, pipe: QwenImagePipeline, height, width, seed, rand_device):
262
+ noise = pipe.generate_noise((1, 16, height//8, width//8), seed=seed, rand_device=rand_device, rand_torch_dtype=pipe.torch_dtype)
263
+ return {"noise": noise}
264
+
265
+
266
+
267
+ class QwenImageUnit_InputImageEmbedder(PipelineUnit):
268
+ def __init__(self):
269
+ super().__init__(
270
+ input_params=("input_image", "noise", "tiled", "tile_size", "tile_stride"),
271
+ onload_model_names=("vae",)
272
+ )
273
+
274
+ def process(self, pipe: QwenImagePipeline, input_image, noise, tiled, tile_size, tile_stride):
275
+ if input_image is None:
276
+ return {"latents": noise, "input_latents": None}
277
+ pipe.load_models_to_device(['vae'])
278
+ image = pipe.preprocess_image(input_image).to(device=pipe.device, dtype=pipe.torch_dtype)
279
+ input_latents = pipe.vae.encode(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
280
+ if pipe.scheduler.training:
281
+ return {"latents": noise, "input_latents": input_latents}
282
+ else:
283
+ latents = pipe.scheduler.add_noise(input_latents, noise, timestep=pipe.scheduler.timesteps[0])
284
+ return {"latents": latents, "input_latents": None}
285
+
286
+
287
+
288
+ class QwenImageUnit_PromptEmbedder(PipelineUnit):
289
+ def __init__(self):
290
+ super().__init__(
291
+ seperate_cfg=True,
292
+ input_params_posi={"prompt": "prompt"},
293
+ input_params_nega={"prompt": "negative_prompt"},
294
+ onload_model_names=("text_encoder",)
295
+ )
296
+
297
+ def extract_masked_hidden(self, hidden_states: torch.Tensor, mask: torch.Tensor):
298
+ bool_mask = mask.bool()
299
+ valid_lengths = bool_mask.sum(dim=1)
300
+ selected = hidden_states[bool_mask]
301
+ split_result = torch.split(selected, valid_lengths.tolist(), dim=0)
302
+ return split_result
303
+
304
+ def process(self, pipe: QwenImagePipeline, prompt) -> dict:
305
+ if pipe.text_encoder is not None:
306
+ prompt = [prompt]
307
+ template = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
308
+ drop_idx = 34
309
+ txt = [template.format(e) for e in prompt]
310
+ txt_tokens = pipe.tokenizer(txt, max_length=1024+drop_idx, padding=True, truncation=True, return_tensors="pt").to(pipe.device)
311
+ hidden_states = pipe.text_encoder(input_ids=txt_tokens.input_ids, attention_mask=txt_tokens.attention_mask, output_hidden_states=True,)[-1]
312
+
313
+ split_hidden_states = self.extract_masked_hidden(hidden_states, txt_tokens.attention_mask)
314
+ split_hidden_states = [e[drop_idx:] for e in split_hidden_states]
315
+ attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states]
316
+ max_seq_len = max([e.size(0) for e in split_hidden_states])
317
+ prompt_embeds = torch.stack([torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states])
318
+ encoder_attention_mask = torch.stack([torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list])
319
+ prompt_embeds = prompt_embeds.to(dtype=pipe.torch_dtype, device=pipe.device)
320
+ return {"prompt_emb": prompt_embeds, "prompt_emb_mask": encoder_attention_mask}
321
+ else:
322
+ return {}
323
+
324
+
325
+
326
+ def model_fn_qwen_image(
327
+ dit: QwenImageDiT = None,
328
+ latents=None,
329
+ timestep=None,
330
+ prompt_emb=None,
331
+ prompt_emb_mask=None,
332
+ height=None,
333
+ width=None,
334
+ use_gradient_checkpointing=False,
335
+ use_gradient_checkpointing_offload=False,
336
+ **kwargs
337
+ ):
338
+ img_shapes = [(latents.shape[0], latents.shape[2]//2, latents.shape[3]//2)]
339
+ txt_seq_lens = prompt_emb_mask.sum(dim=1).tolist()
340
+ timestep = timestep / 1000
341
+
342
+ image = rearrange(latents, "B C (H P) (W Q) -> B (H W) (C P Q)", H=height//16, W=width//16, P=2, Q=2)
343
+
344
+ image = dit.img_in(image)
345
+ text = dit.txt_in(dit.txt_norm(prompt_emb))
346
+ conditioning = dit.time_text_embed(timestep, image.dtype)
347
+ image_rotary_emb = dit.pos_embed(img_shapes, txt_seq_lens, device=latents.device)
348
+
349
+ for block in dit.transformer_blocks:
350
+ text, image = gradient_checkpoint_forward(
351
+ block,
352
+ use_gradient_checkpointing,
353
+ use_gradient_checkpointing_offload,
354
+ image=image,
355
+ text=text,
356
+ temb=conditioning,
357
+ image_rotary_emb=image_rotary_emb,
358
+ )
359
+
360
+ image = dit.norm_out(image, conditioning)
361
+ image = dit.proj_out(image)
362
+
363
+ latents = rearrange(image, "B (H W) (C P Q) -> B C (H P) (W Q)", H=height//16, W=width//16, P=2, Q=2)
364
+ return latents
diffsynth/pipelines/sd3_image.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import ModelManager, SD3TextEncoder1, SD3TextEncoder2, SD3TextEncoder3, SD3DiT, SD3VAEDecoder, SD3VAEEncoder
2
+ from ..prompters import SD3Prompter
3
+ from ..schedulers import FlowMatchScheduler
4
+ from .base import BasePipeline
5
+ import torch
6
+ from tqdm import tqdm
7
+
8
+
9
+
10
+ class SD3ImagePipeline(BasePipeline):
11
+
12
+ def __init__(self, device="cuda", torch_dtype=torch.float16):
13
+ super().__init__(device=device, torch_dtype=torch_dtype, height_division_factor=16, width_division_factor=16)
14
+ self.scheduler = FlowMatchScheduler()
15
+ self.prompter = SD3Prompter()
16
+ # models
17
+ self.text_encoder_1: SD3TextEncoder1 = None
18
+ self.text_encoder_2: SD3TextEncoder2 = None
19
+ self.text_encoder_3: SD3TextEncoder3 = None
20
+ self.dit: SD3DiT = None
21
+ self.vae_decoder: SD3VAEDecoder = None
22
+ self.vae_encoder: SD3VAEEncoder = None
23
+ self.model_names = ['text_encoder_1', 'text_encoder_2', 'text_encoder_3', 'dit', 'vae_decoder', 'vae_encoder']
24
+
25
+
26
+ def denoising_model(self):
27
+ return self.dit
28
+
29
+
30
+ def fetch_models(self, model_manager: ModelManager, prompt_refiner_classes=[]):
31
+ self.text_encoder_1 = model_manager.fetch_model("sd3_text_encoder_1")
32
+ self.text_encoder_2 = model_manager.fetch_model("sd3_text_encoder_2")
33
+ self.text_encoder_3 = model_manager.fetch_model("sd3_text_encoder_3")
34
+ self.dit = model_manager.fetch_model("sd3_dit")
35
+ self.vae_decoder = model_manager.fetch_model("sd3_vae_decoder")
36
+ self.vae_encoder = model_manager.fetch_model("sd3_vae_encoder")
37
+ self.prompter.fetch_models(self.text_encoder_1, self.text_encoder_2, self.text_encoder_3)
38
+ self.prompter.load_prompt_refiners(model_manager, prompt_refiner_classes)
39
+
40
+
41
+ @staticmethod
42
+ def from_model_manager(model_manager: ModelManager, prompt_refiner_classes=[], device=None):
43
+ pipe = SD3ImagePipeline(
44
+ device=model_manager.device if device is None else device,
45
+ torch_dtype=model_manager.torch_dtype,
46
+ )
47
+ pipe.fetch_models(model_manager, prompt_refiner_classes)
48
+ return pipe
49
+
50
+
51
+ def encode_image(self, image, tiled=False, tile_size=64, tile_stride=32):
52
+ latents = self.vae_encoder(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
53
+ return latents
54
+
55
+
56
+ def decode_image(self, latent, tiled=False, tile_size=64, tile_stride=32):
57
+ image = self.vae_decoder(latent.to(self.device), tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
58
+ image = self.vae_output_to_image(image)
59
+ return image
60
+
61
+
62
+ def encode_prompt(self, prompt, positive=True, t5_sequence_length=77):
63
+ prompt_emb, pooled_prompt_emb = self.prompter.encode_prompt(
64
+ prompt, device=self.device, positive=positive, t5_sequence_length=t5_sequence_length
65
+ )
66
+ return {"prompt_emb": prompt_emb, "pooled_prompt_emb": pooled_prompt_emb}
67
+
68
+
69
+ def prepare_extra_input(self, latents=None):
70
+ return {}
71
+
72
+
73
+ @torch.no_grad()
74
+ def __call__(
75
+ self,
76
+ prompt,
77
+ local_prompts=[],
78
+ masks=[],
79
+ mask_scales=[],
80
+ negative_prompt="",
81
+ cfg_scale=7.5,
82
+ input_image=None,
83
+ denoising_strength=1.0,
84
+ height=1024,
85
+ width=1024,
86
+ num_inference_steps=20,
87
+ t5_sequence_length=77,
88
+ tiled=False,
89
+ tile_size=128,
90
+ tile_stride=64,
91
+ seed=None,
92
+ progress_bar_cmd=tqdm,
93
+ progress_bar_st=None,
94
+ ):
95
+ height, width = self.check_resize_height_width(height, width)
96
+
97
+ # Tiler parameters
98
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
99
+
100
+ # Prepare scheduler
101
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength)
102
+
103
+ # Prepare latent tensors
104
+ if input_image is not None:
105
+ self.load_models_to_device(['vae_encoder'])
106
+ image = self.preprocess_image(input_image).to(device=self.device, dtype=self.torch_dtype)
107
+ latents = self.encode_image(image, **tiler_kwargs)
108
+ noise = self.generate_noise((1, 16, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
109
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
110
+ else:
111
+ latents = self.generate_noise((1, 16, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
112
+
113
+ # Encode prompts
114
+ self.load_models_to_device(['text_encoder_1', 'text_encoder_2', 'text_encoder_3'])
115
+ prompt_emb_posi = self.encode_prompt(prompt, positive=True, t5_sequence_length=t5_sequence_length)
116
+ prompt_emb_nega = self.encode_prompt(negative_prompt, positive=False, t5_sequence_length=t5_sequence_length)
117
+ prompt_emb_locals = [self.encode_prompt(prompt_local, t5_sequence_length=t5_sequence_length) for prompt_local in local_prompts]
118
+
119
+ # Denoise
120
+ self.load_models_to_device(['dit'])
121
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
122
+ timestep = timestep.unsqueeze(0).to(self.device)
123
+
124
+ # Classifier-free guidance
125
+ inference_callback = lambda prompt_emb_posi: self.dit(
126
+ latents, timestep=timestep, **prompt_emb_posi, **tiler_kwargs,
127
+ )
128
+ noise_pred_posi = self.control_noise_via_local_prompts(prompt_emb_posi, prompt_emb_locals, masks, mask_scales, inference_callback)
129
+ noise_pred_nega = self.dit(
130
+ latents, timestep=timestep, **prompt_emb_nega, **tiler_kwargs,
131
+ )
132
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
133
+
134
+ # DDIM
135
+ latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents)
136
+
137
+ # UI
138
+ if progress_bar_st is not None:
139
+ progress_bar_st.progress(progress_id / len(self.scheduler.timesteps))
140
+
141
+ # Decode image
142
+ self.load_models_to_device(['vae_decoder'])
143
+ image = self.decode_image(latents, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
144
+
145
+ # offload all models
146
+ self.load_models_to_device([])
147
+ return image
diffsynth/pipelines/sd_image.py ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import SDTextEncoder, SDUNet, SDVAEDecoder, SDVAEEncoder, SDIpAdapter, IpAdapterCLIPImageEmbedder
2
+ from ..models.model_manager import ModelManager
3
+ from ..controlnets import MultiControlNetManager, ControlNetUnit, ControlNetConfigUnit, Annotator
4
+ from ..prompters import SDPrompter
5
+ from ..schedulers import EnhancedDDIMScheduler
6
+ from .base import BasePipeline
7
+ from .dancer import lets_dance
8
+ from typing import List
9
+ import torch
10
+ from tqdm import tqdm
11
+
12
+
13
+
14
+ class SDImagePipeline(BasePipeline):
15
+
16
+ def __init__(self, device="cuda", torch_dtype=torch.float16):
17
+ super().__init__(device=device, torch_dtype=torch_dtype)
18
+ self.scheduler = EnhancedDDIMScheduler()
19
+ self.prompter = SDPrompter()
20
+ # models
21
+ self.text_encoder: SDTextEncoder = None
22
+ self.unet: SDUNet = None
23
+ self.vae_decoder: SDVAEDecoder = None
24
+ self.vae_encoder: SDVAEEncoder = None
25
+ self.controlnet: MultiControlNetManager = None
26
+ self.ipadapter_image_encoder: IpAdapterCLIPImageEmbedder = None
27
+ self.ipadapter: SDIpAdapter = None
28
+ self.model_names = ['text_encoder', 'unet', 'vae_decoder', 'vae_encoder', 'controlnet', 'ipadapter_image_encoder', 'ipadapter']
29
+
30
+
31
+ def denoising_model(self):
32
+ return self.unet
33
+
34
+
35
+ def fetch_models(self, model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
36
+ # Main models
37
+ self.text_encoder = model_manager.fetch_model("sd_text_encoder")
38
+ self.unet = model_manager.fetch_model("sd_unet")
39
+ self.vae_decoder = model_manager.fetch_model("sd_vae_decoder")
40
+ self.vae_encoder = model_manager.fetch_model("sd_vae_encoder")
41
+ self.prompter.fetch_models(self.text_encoder)
42
+ self.prompter.load_prompt_refiners(model_manager, prompt_refiner_classes)
43
+
44
+ # ControlNets
45
+ controlnet_units = []
46
+ for config in controlnet_config_units:
47
+ controlnet_unit = ControlNetUnit(
48
+ Annotator(config.processor_id, device=self.device),
49
+ model_manager.fetch_model("sd_controlnet", config.model_path),
50
+ config.scale
51
+ )
52
+ controlnet_units.append(controlnet_unit)
53
+ self.controlnet = MultiControlNetManager(controlnet_units)
54
+
55
+ # IP-Adapters
56
+ self.ipadapter = model_manager.fetch_model("sd_ipadapter")
57
+ self.ipadapter_image_encoder = model_manager.fetch_model("sd_ipadapter_clip_image_encoder")
58
+
59
+
60
+ @staticmethod
61
+ def from_model_manager(model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[], device=None):
62
+ pipe = SDImagePipeline(
63
+ device=model_manager.device if device is None else device,
64
+ torch_dtype=model_manager.torch_dtype,
65
+ )
66
+ pipe.fetch_models(model_manager, controlnet_config_units, prompt_refiner_classes=[])
67
+ return pipe
68
+
69
+
70
+ def encode_image(self, image, tiled=False, tile_size=64, tile_stride=32):
71
+ latents = self.vae_encoder(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
72
+ return latents
73
+
74
+
75
+ def decode_image(self, latent, tiled=False, tile_size=64, tile_stride=32):
76
+ image = self.vae_decoder(latent.to(self.device), tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
77
+ image = self.vae_output_to_image(image)
78
+ return image
79
+
80
+
81
+ def encode_prompt(self, prompt, clip_skip=1, positive=True):
82
+ prompt_emb = self.prompter.encode_prompt(prompt, clip_skip=clip_skip, device=self.device, positive=positive)
83
+ return {"encoder_hidden_states": prompt_emb}
84
+
85
+
86
+ def prepare_extra_input(self, latents=None):
87
+ return {}
88
+
89
+
90
+ @torch.no_grad()
91
+ def __call__(
92
+ self,
93
+ prompt,
94
+ local_prompts=[],
95
+ masks=[],
96
+ mask_scales=[],
97
+ negative_prompt="",
98
+ cfg_scale=7.5,
99
+ clip_skip=1,
100
+ input_image=None,
101
+ ipadapter_images=None,
102
+ ipadapter_scale=1.0,
103
+ controlnet_image=None,
104
+ denoising_strength=1.0,
105
+ height=512,
106
+ width=512,
107
+ num_inference_steps=20,
108
+ tiled=False,
109
+ tile_size=64,
110
+ tile_stride=32,
111
+ seed=None,
112
+ progress_bar_cmd=tqdm,
113
+ progress_bar_st=None,
114
+ ):
115
+ height, width = self.check_resize_height_width(height, width)
116
+
117
+ # Tiler parameters
118
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
119
+
120
+ # Prepare scheduler
121
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength)
122
+
123
+ # Prepare latent tensors
124
+ if input_image is not None:
125
+ self.load_models_to_device(['vae_encoder'])
126
+ image = self.preprocess_image(input_image).to(device=self.device, dtype=self.torch_dtype)
127
+ latents = self.encode_image(image, **tiler_kwargs)
128
+ noise = self.generate_noise((1, 4, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
129
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
130
+ else:
131
+ latents = self.generate_noise((1, 4, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
132
+
133
+ # Encode prompts
134
+ self.load_models_to_device(['text_encoder'])
135
+ prompt_emb_posi = self.encode_prompt(prompt, clip_skip=clip_skip, positive=True)
136
+ prompt_emb_nega = self.encode_prompt(negative_prompt, clip_skip=clip_skip, positive=False)
137
+ prompt_emb_locals = [self.encode_prompt(prompt_local, clip_skip=clip_skip, positive=True) for prompt_local in local_prompts]
138
+
139
+ # IP-Adapter
140
+ if ipadapter_images is not None:
141
+ self.load_models_to_device(['ipadapter_image_encoder'])
142
+ ipadapter_image_encoding = self.ipadapter_image_encoder(ipadapter_images)
143
+ self.load_models_to_device(['ipadapter'])
144
+ ipadapter_kwargs_list_posi = {"ipadapter_kwargs_list": self.ipadapter(ipadapter_image_encoding, scale=ipadapter_scale)}
145
+ ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": self.ipadapter(torch.zeros_like(ipadapter_image_encoding))}
146
+ else:
147
+ ipadapter_kwargs_list_posi, ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": {}}, {"ipadapter_kwargs_list": {}}
148
+
149
+ # Prepare ControlNets
150
+ if controlnet_image is not None:
151
+ self.load_models_to_device(['controlnet'])
152
+ controlnet_image = self.controlnet.process_image(controlnet_image).to(device=self.device, dtype=self.torch_dtype)
153
+ controlnet_image = controlnet_image.unsqueeze(1)
154
+ controlnet_kwargs = {"controlnet_frames": controlnet_image}
155
+ else:
156
+ controlnet_kwargs = {"controlnet_frames": None}
157
+
158
+ # Denoise
159
+ self.load_models_to_device(['controlnet', 'unet'])
160
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
161
+ timestep = timestep.unsqueeze(0).to(self.device)
162
+
163
+ # Classifier-free guidance
164
+ inference_callback = lambda prompt_emb_posi: lets_dance(
165
+ self.unet, motion_modules=None, controlnet=self.controlnet,
166
+ sample=latents, timestep=timestep,
167
+ **prompt_emb_posi, **controlnet_kwargs, **tiler_kwargs, **ipadapter_kwargs_list_posi,
168
+ device=self.device,
169
+ )
170
+ noise_pred_posi = self.control_noise_via_local_prompts(prompt_emb_posi, prompt_emb_locals, masks, mask_scales, inference_callback)
171
+ noise_pred_nega = lets_dance(
172
+ self.unet, motion_modules=None, controlnet=self.controlnet,
173
+ sample=latents, timestep=timestep, **prompt_emb_nega, **controlnet_kwargs, **tiler_kwargs, **ipadapter_kwargs_list_nega,
174
+ device=self.device,
175
+ )
176
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
177
+
178
+ # DDIM
179
+ latents = self.scheduler.step(noise_pred, timestep, latents)
180
+
181
+ # UI
182
+ if progress_bar_st is not None:
183
+ progress_bar_st.progress(progress_id / len(self.scheduler.timesteps))
184
+
185
+ # Decode image
186
+ self.load_models_to_device(['vae_decoder'])
187
+ image = self.decode_image(latents, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
188
+
189
+ # offload all models
190
+ self.load_models_to_device([])
191
+ return image
diffsynth/pipelines/sd_video.py ADDED
@@ -0,0 +1,269 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import SDTextEncoder, SDUNet, SDVAEDecoder, SDVAEEncoder, SDIpAdapter, IpAdapterCLIPImageEmbedder, SDMotionModel
2
+ from ..models.model_manager import ModelManager
3
+ from ..controlnets import MultiControlNetManager, ControlNetUnit, ControlNetConfigUnit, Annotator
4
+ from ..prompters import SDPrompter
5
+ from ..schedulers import EnhancedDDIMScheduler
6
+ from .sd_image import SDImagePipeline
7
+ from .dancer import lets_dance
8
+ from typing import List
9
+ import torch
10
+ from tqdm import tqdm
11
+
12
+
13
+
14
+ def lets_dance_with_long_video(
15
+ unet: SDUNet,
16
+ motion_modules: SDMotionModel = None,
17
+ controlnet: MultiControlNetManager = None,
18
+ sample = None,
19
+ timestep = None,
20
+ encoder_hidden_states = None,
21
+ ipadapter_kwargs_list = {},
22
+ controlnet_frames = None,
23
+ unet_batch_size = 1,
24
+ controlnet_batch_size = 1,
25
+ cross_frame_attention = False,
26
+ tiled=False,
27
+ tile_size=64,
28
+ tile_stride=32,
29
+ device="cuda",
30
+ animatediff_batch_size=16,
31
+ animatediff_stride=8,
32
+ ):
33
+ num_frames = sample.shape[0]
34
+ hidden_states_output = [(torch.zeros(sample[0].shape, dtype=sample[0].dtype), 0) for i in range(num_frames)]
35
+
36
+ for batch_id in range(0, num_frames, animatediff_stride):
37
+ batch_id_ = min(batch_id + animatediff_batch_size, num_frames)
38
+
39
+ # process this batch
40
+ hidden_states_batch = lets_dance(
41
+ unet, motion_modules, controlnet,
42
+ sample[batch_id: batch_id_].to(device),
43
+ timestep,
44
+ encoder_hidden_states,
45
+ ipadapter_kwargs_list=ipadapter_kwargs_list,
46
+ controlnet_frames=controlnet_frames[:, batch_id: batch_id_].to(device) if controlnet_frames is not None else None,
47
+ unet_batch_size=unet_batch_size, controlnet_batch_size=controlnet_batch_size,
48
+ cross_frame_attention=cross_frame_attention,
49
+ tiled=tiled, tile_size=tile_size, tile_stride=tile_stride, device=device
50
+ ).cpu()
51
+
52
+ # update hidden_states
53
+ for i, hidden_states_updated in zip(range(batch_id, batch_id_), hidden_states_batch):
54
+ bias = max(1 - abs(i - (batch_id + batch_id_ - 1) / 2) / ((batch_id_ - batch_id - 1 + 1e-2) / 2), 1e-2)
55
+ hidden_states, num = hidden_states_output[i]
56
+ hidden_states = hidden_states * (num / (num + bias)) + hidden_states_updated * (bias / (num + bias))
57
+ hidden_states_output[i] = (hidden_states, num + bias)
58
+
59
+ if batch_id_ == num_frames:
60
+ break
61
+
62
+ # output
63
+ hidden_states = torch.stack([h for h, _ in hidden_states_output])
64
+ return hidden_states
65
+
66
+
67
+
68
+ class SDVideoPipeline(SDImagePipeline):
69
+
70
+ def __init__(self, device="cuda", torch_dtype=torch.float16, use_original_animatediff=True):
71
+ super().__init__(device=device, torch_dtype=torch_dtype)
72
+ self.scheduler = EnhancedDDIMScheduler(beta_schedule="linear" if use_original_animatediff else "scaled_linear")
73
+ self.prompter = SDPrompter()
74
+ # models
75
+ self.text_encoder: SDTextEncoder = None
76
+ self.unet: SDUNet = None
77
+ self.vae_decoder: SDVAEDecoder = None
78
+ self.vae_encoder: SDVAEEncoder = None
79
+ self.controlnet: MultiControlNetManager = None
80
+ self.ipadapter_image_encoder: IpAdapterCLIPImageEmbedder = None
81
+ self.ipadapter: SDIpAdapter = None
82
+ self.motion_modules: SDMotionModel = None
83
+
84
+
85
+ def fetch_models(self, model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
86
+ # Main models
87
+ self.text_encoder = model_manager.fetch_model("sd_text_encoder")
88
+ self.unet = model_manager.fetch_model("sd_unet")
89
+ self.vae_decoder = model_manager.fetch_model("sd_vae_decoder")
90
+ self.vae_encoder = model_manager.fetch_model("sd_vae_encoder")
91
+ self.prompter.fetch_models(self.text_encoder)
92
+ self.prompter.load_prompt_refiners(model_manager, prompt_refiner_classes)
93
+
94
+ # ControlNets
95
+ controlnet_units = []
96
+ for config in controlnet_config_units:
97
+ controlnet_unit = ControlNetUnit(
98
+ Annotator(config.processor_id, device=self.device),
99
+ model_manager.fetch_model("sd_controlnet", config.model_path),
100
+ config.scale
101
+ )
102
+ controlnet_units.append(controlnet_unit)
103
+ self.controlnet = MultiControlNetManager(controlnet_units)
104
+
105
+ # IP-Adapters
106
+ self.ipadapter = model_manager.fetch_model("sd_ipadapter")
107
+ self.ipadapter_image_encoder = model_manager.fetch_model("sd_ipadapter_clip_image_encoder")
108
+
109
+ # Motion Modules
110
+ self.motion_modules = model_manager.fetch_model("sd_motion_modules")
111
+ if self.motion_modules is None:
112
+ self.scheduler = EnhancedDDIMScheduler(beta_schedule="scaled_linear")
113
+
114
+
115
+ @staticmethod
116
+ def from_model_manager(model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
117
+ pipe = SDVideoPipeline(
118
+ device=model_manager.device,
119
+ torch_dtype=model_manager.torch_dtype,
120
+ )
121
+ pipe.fetch_models(model_manager, controlnet_config_units, prompt_refiner_classes)
122
+ return pipe
123
+
124
+
125
+ def decode_video(self, latents, tiled=False, tile_size=64, tile_stride=32):
126
+ images = [
127
+ self.decode_image(latents[frame_id: frame_id+1], tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
128
+ for frame_id in range(latents.shape[0])
129
+ ]
130
+ return images
131
+
132
+
133
+ def encode_video(self, processed_images, tiled=False, tile_size=64, tile_stride=32):
134
+ latents = []
135
+ for image in processed_images:
136
+ image = self.preprocess_image(image).to(device=self.device, dtype=self.torch_dtype)
137
+ latent = self.encode_image(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
138
+ latents.append(latent.cpu())
139
+ latents = torch.concat(latents, dim=0)
140
+ return latents
141
+
142
+
143
+ @torch.no_grad()
144
+ def __call__(
145
+ self,
146
+ prompt,
147
+ negative_prompt="",
148
+ cfg_scale=7.5,
149
+ clip_skip=1,
150
+ num_frames=None,
151
+ input_frames=None,
152
+ ipadapter_images=None,
153
+ ipadapter_scale=1.0,
154
+ controlnet_frames=None,
155
+ denoising_strength=1.0,
156
+ height=512,
157
+ width=512,
158
+ num_inference_steps=20,
159
+ animatediff_batch_size = 16,
160
+ animatediff_stride = 8,
161
+ unet_batch_size = 1,
162
+ controlnet_batch_size = 1,
163
+ cross_frame_attention = False,
164
+ smoother=None,
165
+ smoother_progress_ids=[],
166
+ tiled=False,
167
+ tile_size=64,
168
+ tile_stride=32,
169
+ seed=None,
170
+ progress_bar_cmd=tqdm,
171
+ progress_bar_st=None,
172
+ ):
173
+ height, width = self.check_resize_height_width(height, width)
174
+
175
+ # Tiler parameters, batch size ...
176
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
177
+ other_kwargs = {
178
+ "animatediff_batch_size": animatediff_batch_size, "animatediff_stride": animatediff_stride,
179
+ "unet_batch_size": unet_batch_size, "controlnet_batch_size": controlnet_batch_size,
180
+ "cross_frame_attention": cross_frame_attention,
181
+ }
182
+
183
+ # Prepare scheduler
184
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength)
185
+
186
+ # Prepare latent tensors
187
+ if self.motion_modules is None:
188
+ noise = self.generate_noise((1, 4, height//8, width//8), seed=seed, device="cpu", dtype=self.torch_dtype).repeat(num_frames, 1, 1, 1)
189
+ else:
190
+ noise = self.generate_noise((num_frames, 4, height//8, width//8), seed=seed, device="cpu", dtype=self.torch_dtype)
191
+ if input_frames is None or denoising_strength == 1.0:
192
+ latents = noise
193
+ else:
194
+ latents = self.encode_video(input_frames, **tiler_kwargs)
195
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
196
+
197
+ # Encode prompts
198
+ prompt_emb_posi = self.encode_prompt(prompt, clip_skip=clip_skip, positive=True)
199
+ prompt_emb_nega = self.encode_prompt(negative_prompt, clip_skip=clip_skip, positive=False)
200
+
201
+ # IP-Adapter
202
+ if ipadapter_images is not None:
203
+ ipadapter_image_encoding = self.ipadapter_image_encoder(ipadapter_images)
204
+ ipadapter_kwargs_list_posi = {"ipadapter_kwargs_list": self.ipadapter(ipadapter_image_encoding, scale=ipadapter_scale)}
205
+ ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": self.ipadapter(torch.zeros_like(ipadapter_image_encoding))}
206
+ else:
207
+ ipadapter_kwargs_list_posi, ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": {}}, {"ipadapter_kwargs_list": {}}
208
+
209
+ # Prepare ControlNets
210
+ if controlnet_frames is not None:
211
+ if isinstance(controlnet_frames[0], list):
212
+ controlnet_frames_ = []
213
+ for processor_id in range(len(controlnet_frames)):
214
+ controlnet_frames_.append(
215
+ torch.stack([
216
+ self.controlnet.process_image(controlnet_frame, processor_id=processor_id).to(self.torch_dtype)
217
+ for controlnet_frame in progress_bar_cmd(controlnet_frames[processor_id])
218
+ ], dim=1)
219
+ )
220
+ controlnet_frames = torch.concat(controlnet_frames_, dim=0)
221
+ else:
222
+ controlnet_frames = torch.stack([
223
+ self.controlnet.process_image(controlnet_frame).to(self.torch_dtype)
224
+ for controlnet_frame in progress_bar_cmd(controlnet_frames)
225
+ ], dim=1)
226
+ controlnet_kwargs = {"controlnet_frames": controlnet_frames}
227
+ else:
228
+ controlnet_kwargs = {"controlnet_frames": None}
229
+
230
+ # Denoise
231
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
232
+ timestep = timestep.unsqueeze(0).to(self.device)
233
+
234
+ # Classifier-free guidance
235
+ noise_pred_posi = lets_dance_with_long_video(
236
+ self.unet, motion_modules=self.motion_modules, controlnet=self.controlnet,
237
+ sample=latents, timestep=timestep,
238
+ **prompt_emb_posi, **controlnet_kwargs, **ipadapter_kwargs_list_posi, **other_kwargs, **tiler_kwargs,
239
+ device=self.device,
240
+ )
241
+ noise_pred_nega = lets_dance_with_long_video(
242
+ self.unet, motion_modules=self.motion_modules, controlnet=self.controlnet,
243
+ sample=latents, timestep=timestep,
244
+ **prompt_emb_nega, **controlnet_kwargs, **ipadapter_kwargs_list_nega, **other_kwargs, **tiler_kwargs,
245
+ device=self.device,
246
+ )
247
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
248
+
249
+ # DDIM and smoother
250
+ if smoother is not None and progress_id in smoother_progress_ids:
251
+ rendered_frames = self.scheduler.step(noise_pred, timestep, latents, to_final=True)
252
+ rendered_frames = self.decode_video(rendered_frames)
253
+ rendered_frames = smoother(rendered_frames, original_frames=input_frames)
254
+ target_latents = self.encode_video(rendered_frames)
255
+ noise_pred = self.scheduler.return_to_timestep(timestep, latents, target_latents)
256
+ latents = self.scheduler.step(noise_pred, timestep, latents)
257
+
258
+ # UI
259
+ if progress_bar_st is not None:
260
+ progress_bar_st.progress(progress_id / len(self.scheduler.timesteps))
261
+
262
+ # Decode image
263
+ output_frames = self.decode_video(latents, **tiler_kwargs)
264
+
265
+ # Post-process
266
+ if smoother is not None and (num_inference_steps in smoother_progress_ids or -1 in smoother_progress_ids):
267
+ output_frames = smoother(output_frames, original_frames=input_frames)
268
+
269
+ return output_frames
diffsynth/pipelines/sdxl_image.py ADDED
@@ -0,0 +1,226 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import SDXLTextEncoder, SDXLTextEncoder2, SDXLUNet, SDXLVAEDecoder, SDXLVAEEncoder, SDXLIpAdapter, IpAdapterXLCLIPImageEmbedder
2
+ from ..models.kolors_text_encoder import ChatGLMModel
3
+ from ..models.model_manager import ModelManager
4
+ from ..controlnets import MultiControlNetManager, ControlNetUnit, ControlNetConfigUnit, Annotator
5
+ from ..prompters import SDXLPrompter, KolorsPrompter
6
+ from ..schedulers import EnhancedDDIMScheduler
7
+ from .base import BasePipeline
8
+ from .dancer import lets_dance_xl
9
+ from typing import List
10
+ import torch
11
+ from tqdm import tqdm
12
+ from einops import repeat
13
+
14
+
15
+
16
+ class SDXLImagePipeline(BasePipeline):
17
+
18
+ def __init__(self, device="cuda", torch_dtype=torch.float16):
19
+ super().__init__(device=device, torch_dtype=torch_dtype)
20
+ self.scheduler = EnhancedDDIMScheduler()
21
+ self.prompter = SDXLPrompter()
22
+ # models
23
+ self.text_encoder: SDXLTextEncoder = None
24
+ self.text_encoder_2: SDXLTextEncoder2 = None
25
+ self.text_encoder_kolors: ChatGLMModel = None
26
+ self.unet: SDXLUNet = None
27
+ self.vae_decoder: SDXLVAEDecoder = None
28
+ self.vae_encoder: SDXLVAEEncoder = None
29
+ self.controlnet: MultiControlNetManager = None
30
+ self.ipadapter_image_encoder: IpAdapterXLCLIPImageEmbedder = None
31
+ self.ipadapter: SDXLIpAdapter = None
32
+ self.model_names = ['text_encoder', 'text_encoder_2', 'text_encoder_kolors', 'unet', 'vae_decoder', 'vae_encoder', 'controlnet', 'ipadapter_image_encoder', 'ipadapter']
33
+
34
+
35
+ def denoising_model(self):
36
+ return self.unet
37
+
38
+
39
+ def fetch_models(self, model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
40
+ # Main models
41
+ self.text_encoder = model_manager.fetch_model("sdxl_text_encoder")
42
+ self.text_encoder_2 = model_manager.fetch_model("sdxl_text_encoder_2")
43
+ self.text_encoder_kolors = model_manager.fetch_model("kolors_text_encoder")
44
+ self.unet = model_manager.fetch_model("sdxl_unet")
45
+ self.vae_decoder = model_manager.fetch_model("sdxl_vae_decoder")
46
+ self.vae_encoder = model_manager.fetch_model("sdxl_vae_encoder")
47
+
48
+ # ControlNets
49
+ controlnet_units = []
50
+ for config in controlnet_config_units:
51
+ controlnet_unit = ControlNetUnit(
52
+ Annotator(config.processor_id, device=self.device),
53
+ model_manager.fetch_model("sdxl_controlnet", config.model_path),
54
+ config.scale
55
+ )
56
+ controlnet_units.append(controlnet_unit)
57
+ self.controlnet = MultiControlNetManager(controlnet_units)
58
+
59
+ # IP-Adapters
60
+ self.ipadapter = model_manager.fetch_model("sdxl_ipadapter")
61
+ self.ipadapter_image_encoder = model_manager.fetch_model("sdxl_ipadapter_clip_image_encoder")
62
+
63
+ # Kolors
64
+ if self.text_encoder_kolors is not None:
65
+ print("Switch to Kolors. The prompter and scheduler will be replaced.")
66
+ self.prompter = KolorsPrompter()
67
+ self.prompter.fetch_models(self.text_encoder_kolors)
68
+ self.scheduler = EnhancedDDIMScheduler(beta_end=0.014, num_train_timesteps=1100)
69
+ else:
70
+ self.prompter.fetch_models(self.text_encoder, self.text_encoder_2)
71
+ self.prompter.load_prompt_refiners(model_manager, prompt_refiner_classes)
72
+
73
+
74
+ @staticmethod
75
+ def from_model_manager(model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[], device=None):
76
+ pipe = SDXLImagePipeline(
77
+ device=model_manager.device if device is None else device,
78
+ torch_dtype=model_manager.torch_dtype,
79
+ )
80
+ pipe.fetch_models(model_manager, controlnet_config_units, prompt_refiner_classes)
81
+ return pipe
82
+
83
+
84
+ def encode_image(self, image, tiled=False, tile_size=64, tile_stride=32):
85
+ latents = self.vae_encoder(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
86
+ return latents
87
+
88
+
89
+ def decode_image(self, latent, tiled=False, tile_size=64, tile_stride=32):
90
+ image = self.vae_decoder(latent.to(self.device), tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
91
+ image = self.vae_output_to_image(image)
92
+ return image
93
+
94
+
95
+ def encode_prompt(self, prompt, clip_skip=1, clip_skip_2=2, positive=True):
96
+ add_prompt_emb, prompt_emb = self.prompter.encode_prompt(
97
+ prompt,
98
+ clip_skip=clip_skip, clip_skip_2=clip_skip_2,
99
+ device=self.device,
100
+ positive=positive,
101
+ )
102
+ return {"encoder_hidden_states": prompt_emb, "add_text_embeds": add_prompt_emb}
103
+
104
+
105
+ def prepare_extra_input(self, latents=None):
106
+ height, width = latents.shape[2] * 8, latents.shape[3] * 8
107
+ add_time_id = torch.tensor([height, width, 0, 0, height, width], device=self.device).repeat(latents.shape[0])
108
+ return {"add_time_id": add_time_id}
109
+
110
+
111
+ @torch.no_grad()
112
+ def __call__(
113
+ self,
114
+ prompt,
115
+ local_prompts=[],
116
+ masks=[],
117
+ mask_scales=[],
118
+ negative_prompt="",
119
+ cfg_scale=7.5,
120
+ clip_skip=1,
121
+ clip_skip_2=2,
122
+ input_image=None,
123
+ ipadapter_images=None,
124
+ ipadapter_scale=1.0,
125
+ ipadapter_use_instant_style=False,
126
+ controlnet_image=None,
127
+ denoising_strength=1.0,
128
+ height=1024,
129
+ width=1024,
130
+ num_inference_steps=20,
131
+ tiled=False,
132
+ tile_size=64,
133
+ tile_stride=32,
134
+ seed=None,
135
+ progress_bar_cmd=tqdm,
136
+ progress_bar_st=None,
137
+ ):
138
+ height, width = self.check_resize_height_width(height, width)
139
+
140
+ # Tiler parameters
141
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
142
+
143
+ # Prepare scheduler
144
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength)
145
+
146
+ # Prepare latent tensors
147
+ if input_image is not None:
148
+ self.load_models_to_device(['vae_encoder'])
149
+ image = self.preprocess_image(input_image).to(device=self.device, dtype=self.torch_dtype)
150
+ latents = self.encode_image(image, **tiler_kwargs)
151
+ noise = self.generate_noise((1, 4, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
152
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
153
+ else:
154
+ latents = self.generate_noise((1, 4, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
155
+
156
+ # Encode prompts
157
+ self.load_models_to_device(['text_encoder', 'text_encoder_2', 'text_encoder_kolors'])
158
+ prompt_emb_posi = self.encode_prompt(prompt, clip_skip=clip_skip, clip_skip_2=clip_skip_2, positive=True)
159
+ prompt_emb_nega = self.encode_prompt(negative_prompt, clip_skip=clip_skip, clip_skip_2=clip_skip_2, positive=False)
160
+ prompt_emb_locals = [self.encode_prompt(prompt_local, clip_skip=clip_skip, clip_skip_2=clip_skip_2, positive=True) for prompt_local in local_prompts]
161
+
162
+ # IP-Adapter
163
+ if ipadapter_images is not None:
164
+ if ipadapter_use_instant_style:
165
+ self.ipadapter.set_less_adapter()
166
+ else:
167
+ self.ipadapter.set_full_adapter()
168
+ self.load_models_to_device(['ipadapter_image_encoder'])
169
+ ipadapter_image_encoding = self.ipadapter_image_encoder(ipadapter_images)
170
+ self.load_models_to_device(['ipadapter'])
171
+ ipadapter_kwargs_list_posi = {"ipadapter_kwargs_list": self.ipadapter(ipadapter_image_encoding, scale=ipadapter_scale)}
172
+ ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": self.ipadapter(torch.zeros_like(ipadapter_image_encoding))}
173
+ else:
174
+ ipadapter_kwargs_list_posi, ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": {}}, {"ipadapter_kwargs_list": {}}
175
+
176
+ # Prepare ControlNets
177
+ if controlnet_image is not None:
178
+ self.load_models_to_device(['controlnet'])
179
+ controlnet_image = self.controlnet.process_image(controlnet_image).to(device=self.device, dtype=self.torch_dtype)
180
+ controlnet_image = controlnet_image.unsqueeze(1)
181
+ controlnet_kwargs = {"controlnet_frames": controlnet_image}
182
+ else:
183
+ controlnet_kwargs = {"controlnet_frames": None}
184
+
185
+ # Prepare extra input
186
+ extra_input = self.prepare_extra_input(latents)
187
+
188
+ # Denoise
189
+ self.load_models_to_device(['controlnet', 'unet'])
190
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
191
+ timestep = timestep.unsqueeze(0).to(self.device)
192
+
193
+ # Classifier-free guidance
194
+ inference_callback = lambda prompt_emb_posi: lets_dance_xl(
195
+ self.unet, motion_modules=None, controlnet=self.controlnet,
196
+ sample=latents, timestep=timestep, **extra_input,
197
+ **prompt_emb_posi, **controlnet_kwargs, **tiler_kwargs, **ipadapter_kwargs_list_posi,
198
+ device=self.device,
199
+ )
200
+ noise_pred_posi = self.control_noise_via_local_prompts(prompt_emb_posi, prompt_emb_locals, masks, mask_scales, inference_callback)
201
+
202
+ if cfg_scale != 1.0:
203
+ noise_pred_nega = lets_dance_xl(
204
+ self.unet, motion_modules=None, controlnet=self.controlnet,
205
+ sample=latents, timestep=timestep, **extra_input,
206
+ **prompt_emb_nega, **controlnet_kwargs, **tiler_kwargs, **ipadapter_kwargs_list_nega,
207
+ device=self.device,
208
+ )
209
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
210
+ else:
211
+ noise_pred = noise_pred_posi
212
+
213
+ # DDIM
214
+ latents = self.scheduler.step(noise_pred, timestep, latents)
215
+
216
+ # UI
217
+ if progress_bar_st is not None:
218
+ progress_bar_st.progress(progress_id / len(self.scheduler.timesteps))
219
+
220
+ # Decode image
221
+ self.load_models_to_device(['vae_decoder'])
222
+ image = self.decode_image(latents, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
223
+
224
+ # offload all models
225
+ self.load_models_to_device([])
226
+ return image
diffsynth/pipelines/sdxl_video.py ADDED
@@ -0,0 +1,226 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import SDXLTextEncoder, SDXLTextEncoder2, SDXLUNet, SDXLVAEDecoder, SDXLVAEEncoder, SDXLIpAdapter, IpAdapterXLCLIPImageEmbedder, SDXLMotionModel
2
+ from ..models.kolors_text_encoder import ChatGLMModel
3
+ from ..models.model_manager import ModelManager
4
+ from ..controlnets import MultiControlNetManager, ControlNetUnit, ControlNetConfigUnit, Annotator
5
+ from ..prompters import SDXLPrompter, KolorsPrompter
6
+ from ..schedulers import EnhancedDDIMScheduler
7
+ from .sdxl_image import SDXLImagePipeline
8
+ from .dancer import lets_dance_xl
9
+ from typing import List
10
+ import torch
11
+ from tqdm import tqdm
12
+
13
+
14
+
15
+ class SDXLVideoPipeline(SDXLImagePipeline):
16
+
17
+ def __init__(self, device="cuda", torch_dtype=torch.float16, use_original_animatediff=True):
18
+ super().__init__(device=device, torch_dtype=torch_dtype)
19
+ self.scheduler = EnhancedDDIMScheduler(beta_schedule="linear" if use_original_animatediff else "scaled_linear")
20
+ self.prompter = SDXLPrompter()
21
+ # models
22
+ self.text_encoder: SDXLTextEncoder = None
23
+ self.text_encoder_2: SDXLTextEncoder2 = None
24
+ self.text_encoder_kolors: ChatGLMModel = None
25
+ self.unet: SDXLUNet = None
26
+ self.vae_decoder: SDXLVAEDecoder = None
27
+ self.vae_encoder: SDXLVAEEncoder = None
28
+ # self.controlnet: MultiControlNetManager = None (TODO)
29
+ self.ipadapter_image_encoder: IpAdapterXLCLIPImageEmbedder = None
30
+ self.ipadapter: SDXLIpAdapter = None
31
+ self.motion_modules: SDXLMotionModel = None
32
+
33
+
34
+ def fetch_models(self, model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
35
+ # Main models
36
+ self.text_encoder = model_manager.fetch_model("sdxl_text_encoder")
37
+ self.text_encoder_2 = model_manager.fetch_model("sdxl_text_encoder_2")
38
+ self.text_encoder_kolors = model_manager.fetch_model("kolors_text_encoder")
39
+ self.unet = model_manager.fetch_model("sdxl_unet")
40
+ self.vae_decoder = model_manager.fetch_model("sdxl_vae_decoder")
41
+ self.vae_encoder = model_manager.fetch_model("sdxl_vae_encoder")
42
+ self.prompter.fetch_models(self.text_encoder)
43
+ self.prompter.load_prompt_refiners(model_manager, prompt_refiner_classes)
44
+
45
+ # ControlNets (TODO)
46
+
47
+ # IP-Adapters
48
+ self.ipadapter = model_manager.fetch_model("sdxl_ipadapter")
49
+ self.ipadapter_image_encoder = model_manager.fetch_model("sdxl_ipadapter_clip_image_encoder")
50
+
51
+ # Motion Modules
52
+ self.motion_modules = model_manager.fetch_model("sdxl_motion_modules")
53
+ if self.motion_modules is None:
54
+ self.scheduler = EnhancedDDIMScheduler(beta_schedule="scaled_linear")
55
+
56
+ # Kolors
57
+ if self.text_encoder_kolors is not None:
58
+ print("Switch to Kolors. The prompter will be replaced.")
59
+ self.prompter = KolorsPrompter()
60
+ self.prompter.fetch_models(self.text_encoder_kolors)
61
+ # The schedulers of AniamteDiff and Kolors are incompatible. We align it with AniamteDiff.
62
+ if self.motion_modules is None:
63
+ self.scheduler = EnhancedDDIMScheduler(beta_end=0.014, num_train_timesteps=1100)
64
+ else:
65
+ self.prompter.fetch_models(self.text_encoder, self.text_encoder_2)
66
+
67
+
68
+ @staticmethod
69
+ def from_model_manager(model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
70
+ pipe = SDXLVideoPipeline(
71
+ device=model_manager.device,
72
+ torch_dtype=model_manager.torch_dtype,
73
+ )
74
+ pipe.fetch_models(model_manager, controlnet_config_units, prompt_refiner_classes)
75
+ return pipe
76
+
77
+
78
+ def decode_video(self, latents, tiled=False, tile_size=64, tile_stride=32):
79
+ images = [
80
+ self.decode_image(latents[frame_id: frame_id+1], tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
81
+ for frame_id in range(latents.shape[0])
82
+ ]
83
+ return images
84
+
85
+
86
+ def encode_video(self, processed_images, tiled=False, tile_size=64, tile_stride=32):
87
+ latents = []
88
+ for image in processed_images:
89
+ image = self.preprocess_image(image).to(device=self.device, dtype=self.torch_dtype)
90
+ latent = self.encode_image(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
91
+ latents.append(latent.cpu())
92
+ latents = torch.concat(latents, dim=0)
93
+ return latents
94
+
95
+
96
+ @torch.no_grad()
97
+ def __call__(
98
+ self,
99
+ prompt,
100
+ negative_prompt="",
101
+ cfg_scale=7.5,
102
+ clip_skip=1,
103
+ num_frames=None,
104
+ input_frames=None,
105
+ ipadapter_images=None,
106
+ ipadapter_scale=1.0,
107
+ ipadapter_use_instant_style=False,
108
+ controlnet_frames=None,
109
+ denoising_strength=1.0,
110
+ height=512,
111
+ width=512,
112
+ num_inference_steps=20,
113
+ animatediff_batch_size = 16,
114
+ animatediff_stride = 8,
115
+ unet_batch_size = 1,
116
+ controlnet_batch_size = 1,
117
+ cross_frame_attention = False,
118
+ smoother=None,
119
+ smoother_progress_ids=[],
120
+ tiled=False,
121
+ tile_size=64,
122
+ tile_stride=32,
123
+ seed=None,
124
+ progress_bar_cmd=tqdm,
125
+ progress_bar_st=None,
126
+ ):
127
+ height, width = self.check_resize_height_width(height, width)
128
+
129
+ # Tiler parameters, batch size ...
130
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
131
+
132
+ # Prepare scheduler
133
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength)
134
+
135
+ # Prepare latent tensors
136
+ if self.motion_modules is None:
137
+ noise = self.generate_noise((1, 4, height//8, width//8), seed=seed, device="cpu", dtype=self.torch_dtype).repeat(num_frames, 1, 1, 1)
138
+ else:
139
+ noise = self.generate_noise((num_frames, 4, height//8, width//8), seed=seed, device="cpu", dtype=self.torch_dtype)
140
+ if input_frames is None or denoising_strength == 1.0:
141
+ latents = noise
142
+ else:
143
+ latents = self.encode_video(input_frames, **tiler_kwargs)
144
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
145
+ latents = latents.to(self.device) # will be deleted for supporting long videos
146
+
147
+ # Encode prompts
148
+ prompt_emb_posi = self.encode_prompt(prompt, clip_skip=clip_skip, positive=True)
149
+ prompt_emb_nega = self.encode_prompt(negative_prompt, clip_skip=clip_skip, positive=False)
150
+
151
+ # IP-Adapter
152
+ if ipadapter_images is not None:
153
+ if ipadapter_use_instant_style:
154
+ self.ipadapter.set_less_adapter()
155
+ else:
156
+ self.ipadapter.set_full_adapter()
157
+ ipadapter_image_encoding = self.ipadapter_image_encoder(ipadapter_images)
158
+ ipadapter_kwargs_list_posi = {"ipadapter_kwargs_list": self.ipadapter(ipadapter_image_encoding, scale=ipadapter_scale)}
159
+ ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": self.ipadapter(torch.zeros_like(ipadapter_image_encoding))}
160
+ else:
161
+ ipadapter_kwargs_list_posi, ipadapter_kwargs_list_nega = {"ipadapter_kwargs_list": {}}, {"ipadapter_kwargs_list": {}}
162
+
163
+ # Prepare ControlNets
164
+ if controlnet_frames is not None:
165
+ if isinstance(controlnet_frames[0], list):
166
+ controlnet_frames_ = []
167
+ for processor_id in range(len(controlnet_frames)):
168
+ controlnet_frames_.append(
169
+ torch.stack([
170
+ self.controlnet.process_image(controlnet_frame, processor_id=processor_id).to(self.torch_dtype)
171
+ for controlnet_frame in progress_bar_cmd(controlnet_frames[processor_id])
172
+ ], dim=1)
173
+ )
174
+ controlnet_frames = torch.concat(controlnet_frames_, dim=0)
175
+ else:
176
+ controlnet_frames = torch.stack([
177
+ self.controlnet.process_image(controlnet_frame).to(self.torch_dtype)
178
+ for controlnet_frame in progress_bar_cmd(controlnet_frames)
179
+ ], dim=1)
180
+ controlnet_kwargs = {"controlnet_frames": controlnet_frames}
181
+ else:
182
+ controlnet_kwargs = {"controlnet_frames": None}
183
+
184
+ # Prepare extra input
185
+ extra_input = self.prepare_extra_input(latents)
186
+
187
+ # Denoise
188
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
189
+ timestep = timestep.unsqueeze(0).to(self.device)
190
+
191
+ # Classifier-free guidance
192
+ noise_pred_posi = lets_dance_xl(
193
+ self.unet, motion_modules=self.motion_modules, controlnet=None,
194
+ sample=latents, timestep=timestep,
195
+ **prompt_emb_posi, **controlnet_kwargs, **ipadapter_kwargs_list_posi, **extra_input, **tiler_kwargs,
196
+ device=self.device,
197
+ )
198
+ noise_pred_nega = lets_dance_xl(
199
+ self.unet, motion_modules=self.motion_modules, controlnet=None,
200
+ sample=latents, timestep=timestep,
201
+ **prompt_emb_nega, **controlnet_kwargs, **ipadapter_kwargs_list_nega, **extra_input, **tiler_kwargs,
202
+ device=self.device,
203
+ )
204
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
205
+
206
+ # DDIM and smoother
207
+ if smoother is not None and progress_id in smoother_progress_ids:
208
+ rendered_frames = self.scheduler.step(noise_pred, timestep, latents, to_final=True)
209
+ rendered_frames = self.decode_video(rendered_frames)
210
+ rendered_frames = smoother(rendered_frames, original_frames=input_frames)
211
+ target_latents = self.encode_video(rendered_frames)
212
+ noise_pred = self.scheduler.return_to_timestep(timestep, latents, target_latents)
213
+ latents = self.scheduler.step(noise_pred, timestep, latents)
214
+
215
+ # UI
216
+ if progress_bar_st is not None:
217
+ progress_bar_st.progress(progress_id / len(self.scheduler.timesteps))
218
+
219
+ # Decode image
220
+ output_frames = self.decode_video(latents, **tiler_kwargs)
221
+
222
+ # Post-process
223
+ if smoother is not None and (num_inference_steps in smoother_progress_ids or -1 in smoother_progress_ids):
224
+ output_frames = smoother(output_frames, original_frames=input_frames)
225
+
226
+ return output_frames
diffsynth/pipelines/step_video.py ADDED
@@ -0,0 +1,209 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import ModelManager
2
+ from ..models.hunyuan_dit_text_encoder import HunyuanDiTCLIPTextEncoder
3
+ from ..models.stepvideo_text_encoder import STEP1TextEncoder
4
+ from ..models.stepvideo_dit import StepVideoModel
5
+ from ..models.stepvideo_vae import StepVideoVAE
6
+ from ..schedulers.flow_match import FlowMatchScheduler
7
+ from .base import BasePipeline
8
+ from ..prompters import StepVideoPrompter
9
+ import torch
10
+ from einops import rearrange
11
+ import numpy as np
12
+ from PIL import Image
13
+ from ..vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
14
+ from transformers.models.bert.modeling_bert import BertEmbeddings
15
+ from ..models.stepvideo_dit import RMSNorm
16
+ from ..models.stepvideo_vae import CausalConv, CausalConvAfterNorm, Upsample2D, BaseGroupNorm
17
+
18
+
19
+
20
+ class StepVideoPipeline(BasePipeline):
21
+
22
+ def __init__(self, device="cuda", torch_dtype=torch.float16):
23
+ super().__init__(device=device, torch_dtype=torch_dtype)
24
+ self.scheduler = FlowMatchScheduler(sigma_min=0.0, extra_one_step=True, shift=13.0, reverse_sigmas=True, num_train_timesteps=1)
25
+ self.prompter = StepVideoPrompter()
26
+ self.text_encoder_1: HunyuanDiTCLIPTextEncoder = None
27
+ self.text_encoder_2: STEP1TextEncoder = None
28
+ self.dit: StepVideoModel = None
29
+ self.vae: StepVideoVAE = None
30
+ self.model_names = ['text_encoder_1', 'text_encoder_2', 'dit', 'vae']
31
+
32
+
33
+ def enable_vram_management(self, num_persistent_param_in_dit=None):
34
+ dtype = next(iter(self.text_encoder_1.parameters())).dtype
35
+ enable_vram_management(
36
+ self.text_encoder_1,
37
+ module_map = {
38
+ torch.nn.Linear: AutoWrappedLinear,
39
+ BertEmbeddings: AutoWrappedModule,
40
+ torch.nn.LayerNorm: AutoWrappedModule,
41
+ },
42
+ module_config = dict(
43
+ offload_dtype=dtype,
44
+ offload_device="cpu",
45
+ onload_dtype=dtype,
46
+ onload_device="cpu",
47
+ computation_dtype=torch.float32,
48
+ computation_device=self.device,
49
+ ),
50
+ )
51
+ dtype = next(iter(self.text_encoder_2.parameters())).dtype
52
+ enable_vram_management(
53
+ self.text_encoder_2,
54
+ module_map = {
55
+ torch.nn.Linear: AutoWrappedLinear,
56
+ RMSNorm: AutoWrappedModule,
57
+ torch.nn.Embedding: AutoWrappedModule,
58
+ },
59
+ module_config = dict(
60
+ offload_dtype=dtype,
61
+ offload_device="cpu",
62
+ onload_dtype=dtype,
63
+ onload_device="cpu",
64
+ computation_dtype=self.torch_dtype,
65
+ computation_device=self.device,
66
+ ),
67
+ )
68
+ dtype = next(iter(self.dit.parameters())).dtype
69
+ enable_vram_management(
70
+ self.dit,
71
+ module_map = {
72
+ torch.nn.Linear: AutoWrappedLinear,
73
+ torch.nn.Conv2d: AutoWrappedModule,
74
+ torch.nn.LayerNorm: AutoWrappedModule,
75
+ RMSNorm: AutoWrappedModule,
76
+ },
77
+ module_config = dict(
78
+ offload_dtype=dtype,
79
+ offload_device="cpu",
80
+ onload_dtype=dtype,
81
+ onload_device=self.device,
82
+ computation_dtype=self.torch_dtype,
83
+ computation_device=self.device,
84
+ ),
85
+ max_num_param=num_persistent_param_in_dit,
86
+ overflow_module_config = dict(
87
+ offload_dtype=dtype,
88
+ offload_device="cpu",
89
+ onload_dtype=dtype,
90
+ onload_device="cpu",
91
+ computation_dtype=self.torch_dtype,
92
+ computation_device=self.device,
93
+ ),
94
+ )
95
+ dtype = next(iter(self.vae.parameters())).dtype
96
+ enable_vram_management(
97
+ self.vae,
98
+ module_map = {
99
+ torch.nn.Linear: AutoWrappedLinear,
100
+ torch.nn.Conv3d: AutoWrappedModule,
101
+ CausalConv: AutoWrappedModule,
102
+ CausalConvAfterNorm: AutoWrappedModule,
103
+ Upsample2D: AutoWrappedModule,
104
+ BaseGroupNorm: AutoWrappedModule,
105
+ },
106
+ module_config = dict(
107
+ offload_dtype=dtype,
108
+ offload_device="cpu",
109
+ onload_dtype=dtype,
110
+ onload_device="cpu",
111
+ computation_dtype=self.torch_dtype,
112
+ computation_device=self.device,
113
+ ),
114
+ )
115
+ self.enable_cpu_offload()
116
+
117
+
118
+ def fetch_models(self, model_manager: ModelManager):
119
+ self.text_encoder_1 = model_manager.fetch_model("hunyuan_dit_clip_text_encoder")
120
+ self.text_encoder_2 = model_manager.fetch_model("stepvideo_text_encoder_2")
121
+ self.dit = model_manager.fetch_model("stepvideo_dit")
122
+ self.vae = model_manager.fetch_model("stepvideo_vae")
123
+ self.prompter.fetch_models(self.text_encoder_1, self.text_encoder_2)
124
+
125
+
126
+ @staticmethod
127
+ def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None):
128
+ if device is None: device = model_manager.device
129
+ if torch_dtype is None: torch_dtype = model_manager.torch_dtype
130
+ pipe = StepVideoPipeline(device=device, torch_dtype=torch_dtype)
131
+ pipe.fetch_models(model_manager)
132
+ return pipe
133
+
134
+
135
+ def encode_prompt(self, prompt, positive=True):
136
+ clip_embeds, llm_embeds, llm_mask = self.prompter.encode_prompt(prompt, device=self.device, positive=positive)
137
+ clip_embeds = clip_embeds.to(dtype=self.torch_dtype, device=self.device)
138
+ llm_embeds = llm_embeds.to(dtype=self.torch_dtype, device=self.device)
139
+ llm_mask = llm_mask.to(dtype=self.torch_dtype, device=self.device)
140
+ return {"encoder_hidden_states_2": clip_embeds, "encoder_hidden_states": llm_embeds, "encoder_attention_mask": llm_mask}
141
+
142
+
143
+ def tensor2video(self, frames):
144
+ frames = rearrange(frames, "C T H W -> T H W C")
145
+ frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
146
+ frames = [Image.fromarray(frame) for frame in frames]
147
+ return frames
148
+
149
+
150
+ @torch.no_grad()
151
+ def __call__(
152
+ self,
153
+ prompt,
154
+ negative_prompt="",
155
+ input_video=None,
156
+ denoising_strength=1.0,
157
+ seed=None,
158
+ rand_device="cpu",
159
+ height=544,
160
+ width=992,
161
+ num_frames=204,
162
+ cfg_scale=9.0,
163
+ num_inference_steps=30,
164
+ tiled=True,
165
+ tile_size=(34, 34),
166
+ tile_stride=(16, 16),
167
+ smooth_scale=0.6,
168
+ progress_bar_cmd=lambda x: x,
169
+ progress_bar_st=None,
170
+ ):
171
+ # Tiler parameters
172
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
173
+
174
+ # Scheduler
175
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength)
176
+
177
+ # Initialize noise
178
+ latents = self.generate_noise((1, max(num_frames//17*3, 1), 64, height//16, width//16), seed=seed, device=rand_device, dtype=self.torch_dtype).to(self.device)
179
+
180
+ # Encode prompts
181
+ self.load_models_to_device(["text_encoder_1", "text_encoder_2"])
182
+ prompt_emb_posi = self.encode_prompt(prompt, positive=True)
183
+ if cfg_scale != 1.0:
184
+ prompt_emb_nega = self.encode_prompt(negative_prompt, positive=False)
185
+
186
+ # Denoise
187
+ self.load_models_to_device(["dit"])
188
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
189
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
190
+ print(f"Step {progress_id + 1} / {len(self.scheduler.timesteps)}")
191
+
192
+ # Inference
193
+ noise_pred_posi = self.dit(latents, timestep=timestep, **prompt_emb_posi)
194
+ if cfg_scale != 1.0:
195
+ noise_pred_nega = self.dit(latents, timestep=timestep, **prompt_emb_nega)
196
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
197
+ else:
198
+ noise_pred = noise_pred_posi
199
+
200
+ # Scheduler
201
+ latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents)
202
+
203
+ # Decode
204
+ self.load_models_to_device(['vae'])
205
+ frames = self.vae.decode(latents, device=self.device, smooth_scale=smooth_scale, **tiler_kwargs)
206
+ self.load_models_to_device([])
207
+ frames = self.tensor2video(frames[0])
208
+
209
+ return frames
diffsynth/pipelines/svd_video.py ADDED
@@ -0,0 +1,300 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import ModelManager, SVDImageEncoder, SVDUNet, SVDVAEEncoder, SVDVAEDecoder
2
+ from ..schedulers import ContinuousODEScheduler
3
+ from .base import BasePipeline
4
+ import torch
5
+ from tqdm import tqdm
6
+ from PIL import Image
7
+ import numpy as np
8
+ from einops import rearrange, repeat
9
+
10
+
11
+
12
+ class SVDVideoPipeline(BasePipeline):
13
+
14
+ def __init__(self, device="cuda", torch_dtype=torch.float16):
15
+ super().__init__(device=device, torch_dtype=torch_dtype)
16
+ self.scheduler = ContinuousODEScheduler()
17
+ # models
18
+ self.image_encoder: SVDImageEncoder = None
19
+ self.unet: SVDUNet = None
20
+ self.vae_encoder: SVDVAEEncoder = None
21
+ self.vae_decoder: SVDVAEDecoder = None
22
+
23
+
24
+ def fetch_models(self, model_manager: ModelManager):
25
+ self.image_encoder = model_manager.fetch_model("svd_image_encoder")
26
+ self.unet = model_manager.fetch_model("svd_unet")
27
+ self.vae_encoder = model_manager.fetch_model("svd_vae_encoder")
28
+ self.vae_decoder = model_manager.fetch_model("svd_vae_decoder")
29
+
30
+
31
+ @staticmethod
32
+ def from_model_manager(model_manager: ModelManager, **kwargs):
33
+ pipe = SVDVideoPipeline(
34
+ device=model_manager.device,
35
+ torch_dtype=model_manager.torch_dtype
36
+ )
37
+ pipe.fetch_models(model_manager)
38
+ return pipe
39
+
40
+
41
+ def encode_image_with_clip(self, image):
42
+ image = self.preprocess_image(image).to(device=self.device, dtype=self.torch_dtype)
43
+ image = SVDCLIPImageProcessor().resize_with_antialiasing(image, (224, 224))
44
+ image = (image + 1.0) / 2.0
45
+ mean = torch.tensor([0.48145466, 0.4578275, 0.40821073]).reshape(1, 3, 1, 1).to(device=self.device, dtype=self.torch_dtype)
46
+ std = torch.tensor([0.26862954, 0.26130258, 0.27577711]).reshape(1, 3, 1, 1).to(device=self.device, dtype=self.torch_dtype)
47
+ image = (image - mean) / std
48
+ image_emb = self.image_encoder(image)
49
+ return image_emb
50
+
51
+
52
+ def encode_image_with_vae(self, image, noise_aug_strength, seed=None):
53
+ image = self.preprocess_image(image).to(device=self.device, dtype=self.torch_dtype)
54
+ noise = self.generate_noise(image.shape, seed=seed, device=self.device, dtype=self.torch_dtype)
55
+ image = image + noise_aug_strength * noise
56
+ image_emb = self.vae_encoder(image) / self.vae_encoder.scaling_factor
57
+ return image_emb
58
+
59
+
60
+ def encode_video_with_vae(self, video):
61
+ video = torch.concat([self.preprocess_image(frame) for frame in video], dim=0)
62
+ video = rearrange(video, "T C H W -> 1 C T H W")
63
+ video = video.to(device=self.device, dtype=self.torch_dtype)
64
+ latents = self.vae_encoder.encode_video(video)
65
+ latents = rearrange(latents[0], "C T H W -> T C H W")
66
+ return latents
67
+
68
+
69
+ def tensor2video(self, frames):
70
+ frames = rearrange(frames, "C T H W -> T H W C")
71
+ frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
72
+ frames = [Image.fromarray(frame) for frame in frames]
73
+ return frames
74
+
75
+
76
+ def calculate_noise_pred(
77
+ self,
78
+ latents,
79
+ timestep,
80
+ add_time_id,
81
+ cfg_scales,
82
+ image_emb_vae_posi, image_emb_clip_posi,
83
+ image_emb_vae_nega, image_emb_clip_nega
84
+ ):
85
+ # Positive side
86
+ noise_pred_posi = self.unet(
87
+ torch.cat([latents, image_emb_vae_posi], dim=1),
88
+ timestep, image_emb_clip_posi, add_time_id
89
+ )
90
+ # Negative side
91
+ noise_pred_nega = self.unet(
92
+ torch.cat([latents, image_emb_vae_nega], dim=1),
93
+ timestep, image_emb_clip_nega, add_time_id
94
+ )
95
+
96
+ # Classifier-free guidance
97
+ noise_pred = noise_pred_nega + cfg_scales * (noise_pred_posi - noise_pred_nega)
98
+
99
+ return noise_pred
100
+
101
+
102
+ def post_process_latents(self, latents, post_normalize=True, contrast_enhance_scale=1.0):
103
+ if post_normalize:
104
+ mean, std = latents.mean(), latents.std()
105
+ latents = (latents - latents.mean(dim=[1, 2, 3], keepdim=True)) / latents.std(dim=[1, 2, 3], keepdim=True) * std + mean
106
+ latents = latents * contrast_enhance_scale
107
+ return latents
108
+
109
+
110
+ @torch.no_grad()
111
+ def __call__(
112
+ self,
113
+ input_image=None,
114
+ input_video=None,
115
+ mask_frames=[],
116
+ mask_frame_ids=[],
117
+ min_cfg_scale=1.0,
118
+ max_cfg_scale=3.0,
119
+ denoising_strength=1.0,
120
+ num_frames=25,
121
+ height=576,
122
+ width=1024,
123
+ fps=7,
124
+ motion_bucket_id=127,
125
+ noise_aug_strength=0.02,
126
+ num_inference_steps=20,
127
+ post_normalize=True,
128
+ contrast_enhance_scale=1.2,
129
+ seed=None,
130
+ progress_bar_cmd=tqdm,
131
+ progress_bar_st=None,
132
+ ):
133
+ height, width = self.check_resize_height_width(height, width)
134
+
135
+ # Prepare scheduler
136
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength)
137
+
138
+ # Prepare latent tensors
139
+ noise = self.generate_noise((num_frames, 4, height//8, width//8), seed=seed, device=self.device, dtype=self.torch_dtype)
140
+ if denoising_strength == 1.0:
141
+ latents = noise.clone()
142
+ else:
143
+ latents = self.encode_video_with_vae(input_video)
144
+ latents = self.scheduler.add_noise(latents, noise, self.scheduler.timesteps[0])
145
+
146
+ # Prepare mask frames
147
+ if len(mask_frames) > 0:
148
+ mask_latents = self.encode_video_with_vae(mask_frames)
149
+
150
+ # Encode image
151
+ image_emb_clip_posi = self.encode_image_with_clip(input_image)
152
+ image_emb_clip_nega = torch.zeros_like(image_emb_clip_posi)
153
+ image_emb_vae_posi = repeat(self.encode_image_with_vae(input_image, noise_aug_strength, seed=seed), "B C H W -> (B T) C H W", T=num_frames)
154
+ image_emb_vae_nega = torch.zeros_like(image_emb_vae_posi)
155
+
156
+ # Prepare classifier-free guidance
157
+ cfg_scales = torch.linspace(min_cfg_scale, max_cfg_scale, num_frames)
158
+ cfg_scales = cfg_scales.reshape(num_frames, 1, 1, 1).to(device=self.device, dtype=self.torch_dtype)
159
+
160
+ # Prepare positional id
161
+ add_time_id = torch.tensor([[fps-1, motion_bucket_id, noise_aug_strength]], device=self.device)
162
+
163
+ # Denoise
164
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
165
+
166
+ # Mask frames
167
+ for frame_id, mask_frame_id in enumerate(mask_frame_ids):
168
+ latents[mask_frame_id] = self.scheduler.add_noise(mask_latents[frame_id], noise[mask_frame_id], timestep)
169
+
170
+ # Fetch model output
171
+ noise_pred = self.calculate_noise_pred(
172
+ latents, timestep, add_time_id, cfg_scales,
173
+ image_emb_vae_posi, image_emb_clip_posi, image_emb_vae_nega, image_emb_clip_nega
174
+ )
175
+
176
+ # Forward Euler
177
+ latents = self.scheduler.step(noise_pred, timestep, latents)
178
+
179
+ # Update progress bar
180
+ if progress_bar_st is not None:
181
+ progress_bar_st.progress(progress_id / len(self.scheduler.timesteps))
182
+
183
+ # Decode image
184
+ latents = self.post_process_latents(latents, post_normalize=post_normalize, contrast_enhance_scale=contrast_enhance_scale)
185
+ video = self.vae_decoder.decode_video(latents, progress_bar=progress_bar_cmd)
186
+ video = self.tensor2video(video)
187
+
188
+ return video
189
+
190
+
191
+
192
+ class SVDCLIPImageProcessor:
193
+ def __init__(self):
194
+ pass
195
+
196
+ def resize_with_antialiasing(self, input, size, interpolation="bicubic", align_corners=True):
197
+ h, w = input.shape[-2:]
198
+ factors = (h / size[0], w / size[1])
199
+
200
+ # First, we have to determine sigma
201
+ # Taken from skimage: https://github.com/scikit-image/scikit-image/blob/v0.19.2/skimage/transform/_warps.py#L171
202
+ sigmas = (
203
+ max((factors[0] - 1.0) / 2.0, 0.001),
204
+ max((factors[1] - 1.0) / 2.0, 0.001),
205
+ )
206
+
207
+ # Now kernel size. Good results are for 3 sigma, but that is kind of slow. Pillow uses 1 sigma
208
+ # https://github.com/python-pillow/Pillow/blob/master/src/libImaging/Resample.c#L206
209
+ # But they do it in the 2 passes, which gives better results. Let's try 2 sigmas for now
210
+ ks = int(max(2.0 * 2 * sigmas[0], 3)), int(max(2.0 * 2 * sigmas[1], 3))
211
+
212
+ # Make sure it is odd
213
+ if (ks[0] % 2) == 0:
214
+ ks = ks[0] + 1, ks[1]
215
+
216
+ if (ks[1] % 2) == 0:
217
+ ks = ks[0], ks[1] + 1
218
+
219
+ input = self._gaussian_blur2d(input, ks, sigmas)
220
+
221
+ output = torch.nn.functional.interpolate(input, size=size, mode=interpolation, align_corners=align_corners)
222
+ return output
223
+
224
+
225
+ def _compute_padding(self, kernel_size):
226
+ """Compute padding tuple."""
227
+ # 4 or 6 ints: (padding_left, padding_right,padding_top,padding_bottom)
228
+ # https://pytorch.org/docs/stable/nn.html#torch.nn.functional.pad
229
+ if len(kernel_size) < 2:
230
+ raise AssertionError(kernel_size)
231
+ computed = [k - 1 for k in kernel_size]
232
+
233
+ # for even kernels we need to do asymmetric padding :(
234
+ out_padding = 2 * len(kernel_size) * [0]
235
+
236
+ for i in range(len(kernel_size)):
237
+ computed_tmp = computed[-(i + 1)]
238
+
239
+ pad_front = computed_tmp // 2
240
+ pad_rear = computed_tmp - pad_front
241
+
242
+ out_padding[2 * i + 0] = pad_front
243
+ out_padding[2 * i + 1] = pad_rear
244
+
245
+ return out_padding
246
+
247
+
248
+ def _filter2d(self, input, kernel):
249
+ # prepare kernel
250
+ b, c, h, w = input.shape
251
+ tmp_kernel = kernel[:, None, ...].to(device=input.device, dtype=input.dtype)
252
+
253
+ tmp_kernel = tmp_kernel.expand(-1, c, -1, -1)
254
+
255
+ height, width = tmp_kernel.shape[-2:]
256
+
257
+ padding_shape: list[int] = self._compute_padding([height, width])
258
+ input = torch.nn.functional.pad(input, padding_shape, mode="reflect")
259
+
260
+ # kernel and input tensor reshape to align element-wise or batch-wise params
261
+ tmp_kernel = tmp_kernel.reshape(-1, 1, height, width)
262
+ input = input.view(-1, tmp_kernel.size(0), input.size(-2), input.size(-1))
263
+
264
+ # convolve the tensor with the kernel.
265
+ output = torch.nn.functional.conv2d(input, tmp_kernel, groups=tmp_kernel.size(0), padding=0, stride=1)
266
+
267
+ out = output.view(b, c, h, w)
268
+ return out
269
+
270
+
271
+ def _gaussian(self, window_size: int, sigma):
272
+ if isinstance(sigma, float):
273
+ sigma = torch.tensor([[sigma]])
274
+
275
+ batch_size = sigma.shape[0]
276
+
277
+ x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) - window_size // 2).expand(batch_size, -1)
278
+
279
+ if window_size % 2 == 0:
280
+ x = x + 0.5
281
+
282
+ gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0)))
283
+
284
+ return gauss / gauss.sum(-1, keepdim=True)
285
+
286
+
287
+ def _gaussian_blur2d(self, input, kernel_size, sigma):
288
+ if isinstance(sigma, tuple):
289
+ sigma = torch.tensor([sigma], dtype=input.dtype)
290
+ else:
291
+ sigma = sigma.to(dtype=input.dtype)
292
+
293
+ ky, kx = int(kernel_size[0]), int(kernel_size[1])
294
+ bs = sigma.shape[0]
295
+ kernel_x = self._gaussian(kx, sigma[:, 1].view(bs, 1))
296
+ kernel_y = self._gaussian(ky, sigma[:, 0].view(bs, 1))
297
+ out_x = self._filter2d(input, kernel_x[..., None, :])
298
+ out = self._filter2d(out_x, kernel_y[..., None])
299
+
300
+ return out
diffsynth/pipelines/wan_video.py ADDED
@@ -0,0 +1,626 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import types
2
+ from ..models import ModelManager
3
+ from ..models.wan_video_dit import WanModel
4
+ from ..models.wan_video_text_encoder import WanTextEncoder
5
+ from ..models.wan_video_vae import WanVideoVAE
6
+ from ..models.wan_video_image_encoder import WanImageEncoder
7
+ from ..models.wan_video_vace import VaceWanModel
8
+ from ..schedulers.flow_match import FlowMatchScheduler
9
+ from .base import BasePipeline
10
+ from ..prompters import WanPrompter
11
+ import torch, os
12
+ from einops import rearrange
13
+ import numpy as np
14
+ from PIL import Image
15
+ from tqdm import tqdm
16
+ from typing import Optional
17
+
18
+ from ..vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
19
+ from ..models.wan_video_text_encoder import T5RelativeEmbedding, T5LayerNorm
20
+ from ..models.wan_video_dit import RMSNorm, sinusoidal_embedding_1d
21
+ from ..models.wan_video_vae import RMS_norm, CausalConv3d, Upsample
22
+ from ..models.wan_video_motion_controller import WanMotionControllerModel
23
+
24
+
25
+
26
+ class WanVideoPipeline(BasePipeline):
27
+
28
+ def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None):
29
+ super().__init__(device=device, torch_dtype=torch_dtype)
30
+ self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
31
+ self.prompter = WanPrompter(tokenizer_path=tokenizer_path)
32
+ self.text_encoder: WanTextEncoder = None
33
+ self.image_encoder: WanImageEncoder = None
34
+ self.dit: WanModel = None
35
+ self.vae: WanVideoVAE = None
36
+ self.motion_controller: WanMotionControllerModel = None
37
+ self.vace: VaceWanModel = None
38
+ self.model_names = ['text_encoder', 'dit', 'vae', 'image_encoder', 'motion_controller', 'vace']
39
+ self.height_division_factor = 16
40
+ self.width_division_factor = 16
41
+ self.use_unified_sequence_parallel = False
42
+
43
+
44
+ def enable_vram_management(self, num_persistent_param_in_dit=None):
45
+ dtype = next(iter(self.text_encoder.parameters())).dtype
46
+ enable_vram_management(
47
+ self.text_encoder,
48
+ module_map = {
49
+ torch.nn.Linear: AutoWrappedLinear,
50
+ torch.nn.Embedding: AutoWrappedModule,
51
+ T5RelativeEmbedding: AutoWrappedModule,
52
+ T5LayerNorm: AutoWrappedModule,
53
+ },
54
+ module_config = dict(
55
+ offload_dtype=dtype,
56
+ offload_device="cpu",
57
+ onload_dtype=dtype,
58
+ onload_device="cpu",
59
+ computation_dtype=self.torch_dtype,
60
+ computation_device=self.device,
61
+ ),
62
+ )
63
+ dtype = next(iter(self.dit.parameters())).dtype
64
+ enable_vram_management(
65
+ self.dit,
66
+ module_map = {
67
+ torch.nn.Linear: AutoWrappedLinear,
68
+ torch.nn.Conv3d: AutoWrappedModule,
69
+ torch.nn.LayerNorm: AutoWrappedModule,
70
+ RMSNorm: AutoWrappedModule,
71
+ torch.nn.Conv2d: AutoWrappedModule,
72
+ },
73
+ module_config = dict(
74
+ offload_dtype=dtype,
75
+ offload_device="cpu",
76
+ onload_dtype=dtype,
77
+ onload_device=self.device,
78
+ computation_dtype=self.torch_dtype,
79
+ computation_device=self.device,
80
+ ),
81
+ max_num_param=num_persistent_param_in_dit,
82
+ overflow_module_config = dict(
83
+ offload_dtype=dtype,
84
+ offload_device="cpu",
85
+ onload_dtype=dtype,
86
+ onload_device="cpu",
87
+ computation_dtype=self.torch_dtype,
88
+ computation_device=self.device,
89
+ ),
90
+ )
91
+ dtype = next(iter(self.vae.parameters())).dtype
92
+ enable_vram_management(
93
+ self.vae,
94
+ module_map = {
95
+ torch.nn.Linear: AutoWrappedLinear,
96
+ torch.nn.Conv2d: AutoWrappedModule,
97
+ RMS_norm: AutoWrappedModule,
98
+ CausalConv3d: AutoWrappedModule,
99
+ Upsample: AutoWrappedModule,
100
+ torch.nn.SiLU: AutoWrappedModule,
101
+ torch.nn.Dropout: AutoWrappedModule,
102
+ },
103
+ module_config = dict(
104
+ offload_dtype=dtype,
105
+ offload_device="cpu",
106
+ onload_dtype=dtype,
107
+ onload_device=self.device,
108
+ computation_dtype=self.torch_dtype,
109
+ computation_device=self.device,
110
+ ),
111
+ )
112
+ if self.image_encoder is not None:
113
+ dtype = next(iter(self.image_encoder.parameters())).dtype
114
+ enable_vram_management(
115
+ self.image_encoder,
116
+ module_map = {
117
+ torch.nn.Linear: AutoWrappedLinear,
118
+ torch.nn.Conv2d: AutoWrappedModule,
119
+ torch.nn.LayerNorm: AutoWrappedModule,
120
+ },
121
+ module_config = dict(
122
+ offload_dtype=dtype,
123
+ offload_device="cpu",
124
+ onload_dtype=dtype,
125
+ onload_device="cpu",
126
+ computation_dtype=dtype,
127
+ computation_device=self.device,
128
+ ),
129
+ )
130
+ if self.motion_controller is not None:
131
+ dtype = next(iter(self.motion_controller.parameters())).dtype
132
+ enable_vram_management(
133
+ self.motion_controller,
134
+ module_map = {
135
+ torch.nn.Linear: AutoWrappedLinear,
136
+ },
137
+ module_config = dict(
138
+ offload_dtype=dtype,
139
+ offload_device="cpu",
140
+ onload_dtype=dtype,
141
+ onload_device="cpu",
142
+ computation_dtype=dtype,
143
+ computation_device=self.device,
144
+ ),
145
+ )
146
+ if self.vace is not None:
147
+ enable_vram_management(
148
+ self.vace,
149
+ module_map = {
150
+ torch.nn.Linear: AutoWrappedLinear,
151
+ torch.nn.Conv3d: AutoWrappedModule,
152
+ torch.nn.LayerNorm: AutoWrappedModule,
153
+ RMSNorm: AutoWrappedModule,
154
+ },
155
+ module_config = dict(
156
+ offload_dtype=dtype,
157
+ offload_device="cpu",
158
+ onload_dtype=dtype,
159
+ onload_device=self.device,
160
+ computation_dtype=self.torch_dtype,
161
+ computation_device=self.device,
162
+ ),
163
+ )
164
+ self.enable_cpu_offload()
165
+
166
+
167
+ def fetch_models(self, model_manager: ModelManager):
168
+ text_encoder_model_and_path = model_manager.fetch_model("wan_video_text_encoder", require_model_path=True)
169
+ if text_encoder_model_and_path is not None:
170
+ self.text_encoder, tokenizer_path = text_encoder_model_and_path
171
+ self.prompter.fetch_models(self.text_encoder)
172
+ self.prompter.fetch_tokenizer(os.path.join(os.path.dirname(tokenizer_path), "google/umt5-xxl"))
173
+ self.dit = model_manager.fetch_model("wan_video_dit")
174
+ self.vae = model_manager.fetch_model("wan_video_vae")
175
+ self.image_encoder = model_manager.fetch_model("wan_video_image_encoder")
176
+ self.motion_controller = model_manager.fetch_model("wan_video_motion_controller")
177
+ self.vace = model_manager.fetch_model("wan_video_vace")
178
+
179
+
180
+ @staticmethod
181
+ def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None, use_usp=False):
182
+ if device is None: device = model_manager.device
183
+ if torch_dtype is None: torch_dtype = model_manager.torch_dtype
184
+ pipe = WanVideoPipeline(device=device, torch_dtype=torch_dtype)
185
+ pipe.fetch_models(model_manager)
186
+ if use_usp:
187
+ from xfuser.core.distributed import get_sequence_parallel_world_size
188
+ from ..distributed.xdit_context_parallel import usp_attn_forward, usp_dit_forward
189
+
190
+ for block in pipe.dit.blocks:
191
+ block.self_attn.forward = types.MethodType(usp_attn_forward, block.self_attn)
192
+ pipe.dit.forward = types.MethodType(usp_dit_forward, pipe.dit)
193
+ pipe.sp_size = get_sequence_parallel_world_size()
194
+ pipe.use_unified_sequence_parallel = True
195
+ return pipe
196
+
197
+
198
+ def denoising_model(self):
199
+ return self.dit
200
+
201
+
202
+ def encode_prompt(self, prompt, positive=True):
203
+ prompt_emb = self.prompter.encode_prompt(prompt, positive=positive, device=self.device)
204
+ return {"context": prompt_emb}
205
+
206
+
207
+ def encode_image(self, image, end_image, num_frames, height, width, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)):
208
+ image = self.preprocess_image(image.resize((width, height))).to(self.device)
209
+ clip_context = self.image_encoder.encode_image([image])
210
+ msk = torch.ones(1, num_frames, height//8, width//8, device=self.device)
211
+ msk[:, 1:] = 0
212
+ if end_image is not None:
213
+ end_image = self.preprocess_image(end_image.resize((width, height))).to(self.device)
214
+ vae_input = torch.concat([image.transpose(0,1), torch.zeros(3, num_frames-2, height, width).to(image.device), end_image.transpose(0,1)],dim=1)
215
+ if self.dit.has_image_pos_emb:
216
+ clip_context = torch.concat([clip_context, self.image_encoder.encode_image([end_image])], dim=1)
217
+ msk[:, -1:] = 1
218
+ else:
219
+ vae_input = torch.concat([image.transpose(0, 1), torch.zeros(3, num_frames-1, height, width).to(image.device)], dim=1)
220
+
221
+ msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
222
+ msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
223
+ msk = msk.transpose(1, 2)[0]
224
+
225
+ y = self.vae.encode([vae_input.to(dtype=self.torch_dtype, device=self.device)], device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)[0]
226
+ y = y.to(dtype=self.torch_dtype, device=self.device)
227
+ y = torch.concat([msk, y])
228
+ y = y.unsqueeze(0)
229
+ clip_context = clip_context.to(dtype=self.torch_dtype, device=self.device)
230
+ y = y.to(dtype=self.torch_dtype, device=self.device)
231
+ return {"clip_feature": clip_context, "y": y}
232
+
233
+
234
+ def encode_control_video(self, control_video, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
235
+ control_video = self.preprocess_images(control_video)
236
+ control_video = torch.stack(control_video, dim=2).to(dtype=self.torch_dtype, device=self.device)
237
+ latents = self.encode_video(control_video, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=self.torch_dtype, device=self.device)
238
+ return latents
239
+
240
+
241
+ def prepare_reference_image(self, reference_image, height, width):
242
+ if reference_image is not None:
243
+ self.load_models_to_device(["vae"])
244
+ reference_image = reference_image.resize((width, height))
245
+ reference_image = self.preprocess_images([reference_image])
246
+ reference_image = torch.stack(reference_image, dim=2).to(dtype=self.torch_dtype, device=self.device)
247
+ reference_latents = self.vae.encode(reference_image, device=self.device)
248
+ return {"reference_latents": reference_latents}
249
+ else:
250
+ return {}
251
+
252
+
253
+ def prepare_controlnet_kwargs(self, control_video, num_frames, height, width, clip_feature=None, y=None, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
254
+ if control_video is not None:
255
+ control_latents = self.encode_control_video(control_video, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
256
+ if clip_feature is None or y is None:
257
+ clip_feature = torch.zeros((1, 257, 1280), dtype=self.torch_dtype, device=self.device)
258
+ y = torch.zeros((1, 16, (num_frames - 1) // 4 + 1, height//8, width//8), dtype=self.torch_dtype, device=self.device)
259
+ else:
260
+ y = y[:, -16:]
261
+ y = torch.concat([control_latents, y], dim=1)
262
+ return {"clip_feature": clip_feature, "y": y}
263
+
264
+
265
+ def tensor2video(self, frames):
266
+ frames = rearrange(frames, "C T H W -> T H W C")
267
+ frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
268
+ frames = [Image.fromarray(frame) for frame in frames]
269
+ return frames
270
+
271
+
272
+ def prepare_extra_input(self, latents=None):
273
+ return {}
274
+
275
+
276
+ def encode_video(self, input_video, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
277
+ latents = self.vae.encode(input_video, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
278
+ return latents
279
+
280
+
281
+ def decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
282
+ frames = self.vae.decode(latents, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
283
+ return frames
284
+
285
+
286
+ def prepare_unified_sequence_parallel(self):
287
+ return {"use_unified_sequence_parallel": self.use_unified_sequence_parallel}
288
+
289
+
290
+ def prepare_motion_bucket_id(self, motion_bucket_id):
291
+ motion_bucket_id = torch.Tensor((motion_bucket_id,)).to(dtype=self.torch_dtype, device=self.device)
292
+ return {"motion_bucket_id": motion_bucket_id}
293
+
294
+
295
+ def prepare_vace_kwargs(
296
+ self,
297
+ latents,
298
+ vace_video=None, vace_mask=None, vace_reference_image=None, vace_scale=1.0,
299
+ height=480, width=832, num_frames=81,
300
+ seed=None, rand_device="cpu",
301
+ tiled=True, tile_size=(34, 34), tile_stride=(18, 16)
302
+ ):
303
+ if vace_video is not None or vace_mask is not None or vace_reference_image is not None:
304
+ self.load_models_to_device(["vae"])
305
+ if vace_video is None:
306
+ vace_video = torch.zeros((1, 3, num_frames, height, width), dtype=self.torch_dtype, device=self.device)
307
+ else:
308
+ vace_video = self.preprocess_images(vace_video)
309
+ vace_video = torch.stack(vace_video, dim=2).to(dtype=self.torch_dtype, device=self.device)
310
+
311
+ if vace_mask is None:
312
+ vace_mask = torch.ones_like(vace_video)
313
+ else:
314
+ vace_mask = self.preprocess_images(vace_mask)
315
+ vace_mask = torch.stack(vace_mask, dim=2).to(dtype=self.torch_dtype, device=self.device)
316
+
317
+ inactive = vace_video * (1 - vace_mask) + 0 * vace_mask
318
+ reactive = vace_video * vace_mask + 0 * (1 - vace_mask)
319
+ inactive = self.encode_video(inactive, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=self.torch_dtype, device=self.device)
320
+ reactive = self.encode_video(reactive, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=self.torch_dtype, device=self.device)
321
+ vace_video_latents = torch.concat((inactive, reactive), dim=1)
322
+
323
+ vace_mask_latents = rearrange(vace_mask[0,0], "T (H P) (W Q) -> 1 (P Q) T H W", P=8, Q=8)
324
+ vace_mask_latents = torch.nn.functional.interpolate(vace_mask_latents, size=((vace_mask_latents.shape[2] + 3) // 4, vace_mask_latents.shape[3], vace_mask_latents.shape[4]), mode='nearest-exact')
325
+
326
+ if vace_reference_image is None:
327
+ pass
328
+ else:
329
+ vace_reference_image = self.preprocess_images([vace_reference_image])
330
+ vace_reference_image = torch.stack(vace_reference_image, dim=2).to(dtype=self.torch_dtype, device=self.device)
331
+ vace_reference_latents = self.encode_video(vace_reference_image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=self.torch_dtype, device=self.device)
332
+ vace_reference_latents = torch.concat((vace_reference_latents, torch.zeros_like(vace_reference_latents)), dim=1)
333
+ vace_video_latents = torch.concat((vace_reference_latents, vace_video_latents), dim=2)
334
+ vace_mask_latents = torch.concat((torch.zeros_like(vace_mask_latents[:, :, :1]), vace_mask_latents), dim=2)
335
+
336
+ noise = self.generate_noise((1, 16, 1, latents.shape[3], latents.shape[4]), seed=seed, device=rand_device, dtype=torch.float32)
337
+ noise = noise.to(dtype=self.torch_dtype, device=self.device)
338
+ latents = torch.concat((noise, latents), dim=2)
339
+
340
+ vace_context = torch.concat((vace_video_latents, vace_mask_latents), dim=1)
341
+ return latents, {"vace_context": vace_context, "vace_scale": vace_scale}
342
+ else:
343
+ return latents, {"vace_context": None, "vace_scale": vace_scale}
344
+
345
+
346
+ @torch.no_grad()
347
+ def __call__(
348
+ self,
349
+ prompt,
350
+ negative_prompt="",
351
+ input_image=None,
352
+ end_image=None,
353
+ input_video=None,
354
+ control_video=None,
355
+ reference_image=None,
356
+ vace_video=None,
357
+ vace_video_mask=None,
358
+ vace_reference_image=None,
359
+ vace_scale=1.0,
360
+ denoising_strength=1.0,
361
+ seed=None,
362
+ rand_device="cpu",
363
+ height=480,
364
+ width=832,
365
+ num_frames=81,
366
+ cfg_scale=5.0,
367
+ num_inference_steps=50,
368
+ sigma_shift=5.0,
369
+ motion_bucket_id=None,
370
+ tiled=True,
371
+ tile_size=(30, 52),
372
+ tile_stride=(15, 26),
373
+ tea_cache_l1_thresh=None,
374
+ tea_cache_model_id="",
375
+ progress_bar_cmd=tqdm,
376
+ progress_bar_st=None,
377
+ ):
378
+ # Parameter check
379
+ height, width = self.check_resize_height_width(height, width)
380
+ if num_frames % 4 != 1:
381
+ num_frames = (num_frames + 2) // 4 * 4 + 1
382
+ print(f"Only `num_frames % 4 == 1` is acceptable. We round it up to {num_frames}.")
383
+
384
+ # Tiler parameters
385
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
386
+
387
+ # Scheduler
388
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, shift=sigma_shift)
389
+
390
+ # Initialize noise
391
+ noise = self.generate_noise((1, 16, (num_frames - 1) // 4 + 1, height//8, width//8), seed=seed, device=rand_device, dtype=torch.float32)
392
+ noise = noise.to(dtype=self.torch_dtype, device=self.device)
393
+ if input_video is not None:
394
+ self.load_models_to_device(['vae'])
395
+ input_video = self.preprocess_images(input_video)
396
+ input_video = torch.stack(input_video, dim=2).to(dtype=self.torch_dtype, device=self.device)
397
+ latents = self.encode_video(input_video, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
398
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
399
+ else:
400
+ latents = noise
401
+
402
+ # Encode prompts
403
+ self.load_models_to_device(["text_encoder"])
404
+ prompt_emb_posi = self.encode_prompt(prompt, positive=True)
405
+ if cfg_scale != 1.0:
406
+ prompt_emb_nega = self.encode_prompt(negative_prompt, positive=False)
407
+
408
+ # Encode image
409
+ if input_image is not None and self.image_encoder is not None:
410
+ self.load_models_to_device(["image_encoder", "vae"])
411
+ image_emb = self.encode_image(input_image, end_image, num_frames, height, width, **tiler_kwargs)
412
+ else:
413
+ image_emb = {}
414
+
415
+ # Reference image
416
+ reference_image_kwargs = self.prepare_reference_image(reference_image, height, width)
417
+
418
+ # ControlNet
419
+ if control_video is not None:
420
+ self.load_models_to_device(["image_encoder", "vae"])
421
+ image_emb = self.prepare_controlnet_kwargs(control_video, num_frames, height, width, **image_emb, **tiler_kwargs)
422
+
423
+ # Motion Controller
424
+ if self.motion_controller is not None and motion_bucket_id is not None:
425
+ motion_kwargs = self.prepare_motion_bucket_id(motion_bucket_id)
426
+ else:
427
+ motion_kwargs = {}
428
+
429
+ # Extra input
430
+ extra_input = self.prepare_extra_input(latents)
431
+
432
+ # VACE
433
+ latents, vace_kwargs = self.prepare_vace_kwargs(
434
+ latents, vace_video, vace_video_mask, vace_reference_image, vace_scale,
435
+ height=height, width=width, num_frames=num_frames, seed=seed, rand_device=rand_device, **tiler_kwargs
436
+ )
437
+
438
+ # TeaCache
439
+ tea_cache_posi = {"tea_cache": TeaCache(num_inference_steps, rel_l1_thresh=tea_cache_l1_thresh, model_id=tea_cache_model_id) if tea_cache_l1_thresh is not None else None}
440
+ tea_cache_nega = {"tea_cache": TeaCache(num_inference_steps, rel_l1_thresh=tea_cache_l1_thresh, model_id=tea_cache_model_id) if tea_cache_l1_thresh is not None else None}
441
+
442
+ # Unified Sequence Parallel
443
+ usp_kwargs = self.prepare_unified_sequence_parallel()
444
+
445
+ # Denoise
446
+ self.load_models_to_device(["dit", "motion_controller", "vace"])
447
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
448
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
449
+
450
+ # Inference
451
+ noise_pred_posi = model_fn_wan_video(
452
+ self.dit, motion_controller=self.motion_controller, vace=self.vace,
453
+ x=latents, timestep=timestep,
454
+ **prompt_emb_posi, **image_emb, **extra_input,
455
+ **tea_cache_posi, **usp_kwargs, **motion_kwargs, **vace_kwargs, **reference_image_kwargs,
456
+ )
457
+ if cfg_scale != 1.0:
458
+ noise_pred_nega = model_fn_wan_video(
459
+ self.dit, motion_controller=self.motion_controller, vace=self.vace,
460
+ x=latents, timestep=timestep,
461
+ **prompt_emb_nega, **image_emb, **extra_input,
462
+ **tea_cache_nega, **usp_kwargs, **motion_kwargs, **vace_kwargs, **reference_image_kwargs,
463
+ )
464
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
465
+ else:
466
+ noise_pred = noise_pred_posi
467
+
468
+ # Scheduler
469
+ latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents)
470
+
471
+ if vace_reference_image is not None:
472
+ latents = latents[:, :, 1:]
473
+
474
+ # Decode
475
+ self.load_models_to_device(['vae'])
476
+ frames = self.decode_video(latents, **tiler_kwargs)
477
+ self.load_models_to_device([])
478
+ frames = self.tensor2video(frames[0])
479
+
480
+ return frames
481
+
482
+
483
+
484
+ class TeaCache:
485
+ def __init__(self, num_inference_steps, rel_l1_thresh, model_id):
486
+ self.num_inference_steps = num_inference_steps
487
+ self.step = 0
488
+ self.accumulated_rel_l1_distance = 0
489
+ self.previous_modulated_input = None
490
+ self.rel_l1_thresh = rel_l1_thresh
491
+ self.previous_residual = None
492
+ self.previous_hidden_states = None
493
+
494
+ self.coefficients_dict = {
495
+ "Wan2.1-T2V-1.3B": [-5.21862437e+04, 9.23041404e+03, -5.28275948e+02, 1.36987616e+01, -4.99875664e-02],
496
+ "Wan2.1-T2V-14B": [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01],
497
+ "Wan2.1-I2V-14B-480P": [2.57151496e+05, -3.54229917e+04, 1.40286849e+03, -1.35890334e+01, 1.32517977e-01],
498
+ "Wan2.1-I2V-14B-720P": [ 8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02],
499
+ }
500
+ if model_id not in self.coefficients_dict:
501
+ supported_model_ids = ", ".join([i for i in self.coefficients_dict])
502
+ raise ValueError(f"{model_id} is not a supported TeaCache model id. Please choose a valid model id in ({supported_model_ids}).")
503
+ self.coefficients = self.coefficients_dict[model_id]
504
+
505
+ def check(self, dit: WanModel, x, t_mod):
506
+ modulated_inp = t_mod.clone()
507
+ if self.step == 0 or self.step == self.num_inference_steps - 1:
508
+ should_calc = True
509
+ self.accumulated_rel_l1_distance = 0
510
+ else:
511
+ coefficients = self.coefficients
512
+ rescale_func = np.poly1d(coefficients)
513
+ self.accumulated_rel_l1_distance += rescale_func(((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()).cpu().item())
514
+ if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
515
+ should_calc = False
516
+ else:
517
+ should_calc = True
518
+ self.accumulated_rel_l1_distance = 0
519
+ self.previous_modulated_input = modulated_inp
520
+ self.step += 1
521
+ if self.step == self.num_inference_steps:
522
+ self.step = 0
523
+ if should_calc:
524
+ self.previous_hidden_states = x.clone()
525
+ return not should_calc
526
+
527
+ def store(self, hidden_states):
528
+ self.previous_residual = hidden_states - self.previous_hidden_states
529
+ self.previous_hidden_states = None
530
+
531
+ def update(self, hidden_states):
532
+ hidden_states = hidden_states + self.previous_residual
533
+ return hidden_states
534
+
535
+
536
+
537
+ def model_fn_wan_video(
538
+ dit: WanModel,
539
+ motion_controller: WanMotionControllerModel = None,
540
+ vace: VaceWanModel = None,
541
+ x: torch.Tensor = None,
542
+ timestep: torch.Tensor = None,
543
+ context: torch.Tensor = None,
544
+ clip_feature: Optional[torch.Tensor] = None,
545
+ y: Optional[torch.Tensor] = None,
546
+ reference_latents = None,
547
+ vace_context = None,
548
+ vace_scale = 1.0,
549
+ tea_cache: TeaCache = None,
550
+ use_unified_sequence_parallel: bool = False,
551
+ motion_bucket_id: Optional[torch.Tensor] = None,
552
+ **kwargs,
553
+ ):
554
+ if use_unified_sequence_parallel:
555
+ import torch.distributed as dist
556
+ from xfuser.core.distributed import (get_sequence_parallel_rank,
557
+ get_sequence_parallel_world_size,
558
+ get_sp_group)
559
+
560
+ t = dit.time_embedding(sinusoidal_embedding_1d(dit.freq_dim, timestep))
561
+ t_mod = dit.time_projection(t).unflatten(1, (6, dit.dim))
562
+ if motion_bucket_id is not None and motion_controller is not None:
563
+ t_mod = t_mod + motion_controller(motion_bucket_id).unflatten(1, (6, dit.dim))
564
+ context = dit.text_embedding(context)
565
+
566
+ if dit.has_image_input:
567
+ x = torch.cat([x, y], dim=1) # (b, c_x + c_y, f, h, w)
568
+ clip_embdding = dit.img_emb(clip_feature)
569
+ context = torch.cat([clip_embdding, context], dim=1)
570
+
571
+ x, (f, h, w) = dit.patchify(x)
572
+
573
+ # Reference image
574
+ if reference_latents is not None:
575
+ reference_latents = dit.ref_conv(reference_latents[:, :, 0]).flatten(2).transpose(1, 2)
576
+ x = torch.concat([reference_latents, x], dim=1)
577
+ f += 1
578
+
579
+ freqs = torch.cat([
580
+ dit.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
581
+ dit.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
582
+ dit.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
583
+ ], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
584
+
585
+ # TeaCache
586
+ if tea_cache is not None:
587
+ tea_cache_update = tea_cache.check(dit, x, t_mod)
588
+ else:
589
+ tea_cache_update = False
590
+
591
+ if vace_context is not None:
592
+ vace_hints = vace(x, vace_context, context, t_mod, freqs)
593
+
594
+ # blocks
595
+ if use_unified_sequence_parallel:
596
+ if dist.is_initialized() and dist.get_world_size() > 1:
597
+ chunks = torch.chunk(x, get_sequence_parallel_world_size(), dim=1)
598
+ pad_shape = chunks[0].shape[1] - chunks[-1].shape[1]
599
+ chunks = [torch.nn.functional.pad(chunk, (0, 0, 0, chunks[0].shape[1]-chunk.shape[1]), value=0) for chunk in chunks]
600
+ x = chunks[get_sequence_parallel_rank()]
601
+
602
+ if tea_cache_update:
603
+ x = tea_cache.update(x)
604
+ else:
605
+ for block_id, block in enumerate(dit.blocks):
606
+ x = block(x, context, t_mod, freqs)
607
+ if vace_context is not None and block_id in vace.vace_layers_mapping:
608
+ current_vace_hint = vace_hints[vace.vace_layers_mapping[block_id]]
609
+ if use_unified_sequence_parallel and dist.is_initialized() and dist.get_world_size() > 1:
610
+ current_vace_hint = torch.chunk(current_vace_hint, get_sequence_parallel_world_size(), dim=1)[get_sequence_parallel_rank()]
611
+ current_vace_hint = torch.nn.functional.pad(current_vace_hint, (0, 0, 0, chunks[0].shape[1] - current_vace_hint.shape[1]), value=0)
612
+ x = x + current_vace_hint * vace_scale
613
+ if tea_cache is not None:
614
+ tea_cache.store(x)
615
+
616
+ x = dit.head(x, t)
617
+ if use_unified_sequence_parallel:
618
+ if dist.is_initialized() and dist.get_world_size() > 1:
619
+ x = get_sp_group().all_gather(x, dim=1)
620
+ x = x[:, :-pad_shape] if pad_shape > 0 else x
621
+ # Remove reference latents
622
+ if reference_latents is not None:
623
+ x = x[:, reference_latents.shape[1]:]
624
+ f -= 1
625
+ x = dit.unpatchify(x, (f, h, w))
626
+ return x
diffsynth/pipelines/wan_video_new.py ADDED
@@ -0,0 +1,1125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, warnings, glob, os, types
2
+ import numpy as np
3
+ from PIL import Image
4
+ from einops import repeat, reduce
5
+ from typing import Optional, Union
6
+ from dataclasses import dataclass
7
+ from modelscope import snapshot_download
8
+ from einops import rearrange
9
+ import numpy as np
10
+ from PIL import Image
11
+ from tqdm import tqdm
12
+ from typing import Optional
13
+ from typing_extensions import Literal
14
+
15
+ from ..utils import BasePipeline, ModelConfig, PipelineUnit, PipelineUnitRunner
16
+ from ..models import ModelManager, load_state_dict
17
+ from ..models.wan_video_dit import WanModel, RMSNorm, sinusoidal_embedding_1d
18
+ from ..models.wan_video_text_encoder import WanTextEncoder, T5RelativeEmbedding, T5LayerNorm
19
+ from ..models.wan_video_vae import WanVideoVAE, RMS_norm, CausalConv3d, Upsample
20
+ from ..models.wan_video_image_encoder import WanImageEncoder
21
+ from ..models.wan_video_vace import VaceWanModel
22
+ from ..models.wan_video_motion_controller import WanMotionControllerModel
23
+ from ..schedulers.flow_match import FlowMatchScheduler
24
+ from ..prompters import WanPrompter
25
+ from ..vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear, WanAutoCastLayerNorm
26
+ from ..lora import GeneralLoRALoader
27
+
28
+
29
+
30
+ class WanVideoPipeline(BasePipeline):
31
+
32
+ def __init__(self, device="cuda", torch_dtype=torch.bfloat16, tokenizer_path=None):
33
+ super().__init__(
34
+ device=device, torch_dtype=torch_dtype,
35
+ height_division_factor=16, width_division_factor=16, time_division_factor=4, time_division_remainder=1
36
+ )
37
+ self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
38
+ self.prompter = WanPrompter(tokenizer_path=tokenizer_path)
39
+ self.text_encoder: WanTextEncoder = None
40
+ self.image_encoder: WanImageEncoder = None
41
+ self.dit: WanModel = None
42
+ self.dit2: WanModel = None
43
+ self.vae: WanVideoVAE = None
44
+ self.motion_controller: WanMotionControllerModel = None
45
+ self.vace: VaceWanModel = None
46
+ self.in_iteration_models = ("dit", "motion_controller", "vace")
47
+ self.in_iteration_models_2 = ("dit2", "motion_controller", "vace")
48
+ self.unit_runner = PipelineUnitRunner()
49
+ self.units = [
50
+ WanVideoUnit_ShapeChecker(),
51
+ WanVideoUnit_NoiseInitializer(),
52
+ WanVideoUnit_InputVideoEmbedder(),
53
+ WanVideoUnit_PromptEmbedder(),
54
+ WanVideoUnit_ImageEmbedderVAE(),
55
+ WanVideoUnit_ImageEmbedderCLIP(),
56
+ WanVideoUnit_ImageEmbedderFused(),
57
+ WanVideoUnit_FunControl(),
58
+ WanVideoUnit_FunReference(),
59
+ WanVideoUnit_FunCameraControl(),
60
+ WanVideoUnit_SpeedControl(),
61
+ WanVideoUnit_VACE(),
62
+ WanVideoUnit_UnifiedSequenceParallel(),
63
+ WanVideoUnit_TeaCache(),
64
+ WanVideoUnit_CfgMerger(),
65
+ ]
66
+ self.model_fn = model_fn_wan_video
67
+
68
+
69
+ def load_lora(self, module, path, alpha=1):
70
+ loader = GeneralLoRALoader(torch_dtype=self.torch_dtype, device=self.device)
71
+ lora = load_state_dict(path, torch_dtype=self.torch_dtype, device=self.device)
72
+ loader.load(module, lora, alpha=alpha)
73
+
74
+
75
+ def training_loss(self, **inputs):
76
+ max_timestep_boundary = int(inputs.get("max_timestep_boundary", 1) * self.scheduler.num_train_timesteps)
77
+ min_timestep_boundary = int(inputs.get("min_timestep_boundary", 0) * self.scheduler.num_train_timesteps)
78
+ timestep_id = torch.randint(min_timestep_boundary, max_timestep_boundary, (1,))
79
+ timestep = self.scheduler.timesteps[timestep_id].to(dtype=self.torch_dtype, device=self.device)
80
+
81
+ inputs["latents"] = self.scheduler.add_noise(inputs["input_latents"], inputs["noise"], timestep)
82
+ training_target = self.scheduler.training_target(inputs["input_latents"], inputs["noise"], timestep)
83
+
84
+ noise_pred = self.model_fn(**inputs, timestep=timestep)
85
+
86
+ loss = torch.nn.functional.mse_loss(noise_pred.float(), training_target.float())
87
+ loss = loss * self.scheduler.training_weight(timestep)
88
+ return loss
89
+
90
+
91
+ def enable_vram_management(self, num_persistent_param_in_dit=None, vram_limit=None, vram_buffer=0.5):
92
+ self.vram_management_enabled = True
93
+ if num_persistent_param_in_dit is not None:
94
+ vram_limit = None
95
+ else:
96
+ if vram_limit is None:
97
+ vram_limit = self.get_vram()
98
+ vram_limit = vram_limit - vram_buffer
99
+ if self.text_encoder is not None:
100
+ dtype = next(iter(self.text_encoder.parameters())).dtype
101
+ enable_vram_management(
102
+ self.text_encoder,
103
+ module_map = {
104
+ torch.nn.Linear: AutoWrappedLinear,
105
+ torch.nn.Embedding: AutoWrappedModule,
106
+ T5RelativeEmbedding: AutoWrappedModule,
107
+ T5LayerNorm: AutoWrappedModule,
108
+ },
109
+ module_config = dict(
110
+ offload_dtype=dtype,
111
+ offload_device="cpu",
112
+ onload_dtype=dtype,
113
+ onload_device="cpu",
114
+ computation_dtype=self.torch_dtype,
115
+ computation_device=self.device,
116
+ ),
117
+ vram_limit=vram_limit,
118
+ )
119
+ if self.dit is not None:
120
+ dtype = next(iter(self.dit.parameters())).dtype
121
+ device = "cpu" if vram_limit is not None else self.device
122
+ enable_vram_management(
123
+ self.dit,
124
+ module_map = {
125
+ torch.nn.Linear: AutoWrappedLinear,
126
+ torch.nn.Conv3d: AutoWrappedModule,
127
+ torch.nn.LayerNorm: WanAutoCastLayerNorm,
128
+ RMSNorm: AutoWrappedModule,
129
+ torch.nn.Conv2d: AutoWrappedModule,
130
+ },
131
+ module_config = dict(
132
+ offload_dtype=dtype,
133
+ offload_device="cpu",
134
+ onload_dtype=dtype,
135
+ onload_device=device,
136
+ computation_dtype=self.torch_dtype,
137
+ computation_device=self.device,
138
+ ),
139
+ max_num_param=num_persistent_param_in_dit,
140
+ overflow_module_config = dict(
141
+ offload_dtype=dtype,
142
+ offload_device="cpu",
143
+ onload_dtype=dtype,
144
+ onload_device="cpu",
145
+ computation_dtype=self.torch_dtype,
146
+ computation_device=self.device,
147
+ ),
148
+ vram_limit=vram_limit,
149
+ )
150
+ if self.dit2 is not None:
151
+ dtype = next(iter(self.dit2.parameters())).dtype
152
+ device = "cpu" if vram_limit is not None else self.device
153
+ enable_vram_management(
154
+ self.dit2,
155
+ module_map = {
156
+ torch.nn.Linear: AutoWrappedLinear,
157
+ torch.nn.Conv3d: AutoWrappedModule,
158
+ torch.nn.LayerNorm: WanAutoCastLayerNorm,
159
+ RMSNorm: AutoWrappedModule,
160
+ torch.nn.Conv2d: AutoWrappedModule,
161
+ },
162
+ module_config = dict(
163
+ offload_dtype=dtype,
164
+ offload_device="cpu",
165
+ onload_dtype=dtype,
166
+ onload_device=device,
167
+ computation_dtype=self.torch_dtype,
168
+ computation_device=self.device,
169
+ ),
170
+ max_num_param=num_persistent_param_in_dit,
171
+ overflow_module_config = dict(
172
+ offload_dtype=dtype,
173
+ offload_device="cpu",
174
+ onload_dtype=dtype,
175
+ onload_device="cpu",
176
+ computation_dtype=self.torch_dtype,
177
+ computation_device=self.device,
178
+ ),
179
+ vram_limit=vram_limit,
180
+ )
181
+ if self.vae is not None:
182
+ dtype = next(iter(self.vae.parameters())).dtype
183
+ enable_vram_management(
184
+ self.vae,
185
+ module_map = {
186
+ torch.nn.Linear: AutoWrappedLinear,
187
+ torch.nn.Conv2d: AutoWrappedModule,
188
+ RMS_norm: AutoWrappedModule,
189
+ CausalConv3d: AutoWrappedModule,
190
+ Upsample: AutoWrappedModule,
191
+ torch.nn.SiLU: AutoWrappedModule,
192
+ torch.nn.Dropout: AutoWrappedModule,
193
+ },
194
+ module_config = dict(
195
+ offload_dtype=dtype,
196
+ offload_device="cpu",
197
+ onload_dtype=dtype,
198
+ onload_device=self.device,
199
+ computation_dtype=self.torch_dtype,
200
+ computation_device=self.device,
201
+ ),
202
+ )
203
+ if self.image_encoder is not None:
204
+ dtype = next(iter(self.image_encoder.parameters())).dtype
205
+ enable_vram_management(
206
+ self.image_encoder,
207
+ module_map = {
208
+ torch.nn.Linear: AutoWrappedLinear,
209
+ torch.nn.Conv2d: AutoWrappedModule,
210
+ torch.nn.LayerNorm: AutoWrappedModule,
211
+ },
212
+ module_config = dict(
213
+ offload_dtype=dtype,
214
+ offload_device="cpu",
215
+ onload_dtype=dtype,
216
+ onload_device="cpu",
217
+ computation_dtype=dtype,
218
+ computation_device=self.device,
219
+ ),
220
+ )
221
+ if self.motion_controller is not None:
222
+ dtype = next(iter(self.motion_controller.parameters())).dtype
223
+ enable_vram_management(
224
+ self.motion_controller,
225
+ module_map = {
226
+ torch.nn.Linear: AutoWrappedLinear,
227
+ },
228
+ module_config = dict(
229
+ offload_dtype=dtype,
230
+ offload_device="cpu",
231
+ onload_dtype=dtype,
232
+ onload_device="cpu",
233
+ computation_dtype=dtype,
234
+ computation_device=self.device,
235
+ ),
236
+ )
237
+ if self.vace is not None:
238
+ device = "cpu" if vram_limit is not None else self.device
239
+ enable_vram_management(
240
+ self.vace,
241
+ module_map = {
242
+ torch.nn.Linear: AutoWrappedLinear,
243
+ torch.nn.Conv3d: AutoWrappedModule,
244
+ torch.nn.LayerNorm: AutoWrappedModule,
245
+ RMSNorm: AutoWrappedModule,
246
+ },
247
+ module_config = dict(
248
+ offload_dtype=dtype,
249
+ offload_device="cpu",
250
+ onload_dtype=dtype,
251
+ onload_device=device,
252
+ computation_dtype=self.torch_dtype,
253
+ computation_device=self.device,
254
+ ),
255
+ vram_limit=vram_limit,
256
+ )
257
+
258
+
259
+ def initialize_usp(self):
260
+ import torch.distributed as dist
261
+ from xfuser.core.distributed import initialize_model_parallel, init_distributed_environment
262
+ dist.init_process_group(backend="nccl", init_method="env://")
263
+ init_distributed_environment(rank=dist.get_rank(), world_size=dist.get_world_size())
264
+ initialize_model_parallel(
265
+ sequence_parallel_degree=dist.get_world_size(),
266
+ ring_degree=1,
267
+ ulysses_degree=dist.get_world_size(),
268
+ )
269
+ torch.cuda.set_device(dist.get_rank())
270
+
271
+
272
+ def enable_usp(self):
273
+ from xfuser.core.distributed import get_sequence_parallel_world_size
274
+ from ..distributed.xdit_context_parallel import usp_attn_forward, usp_dit_forward
275
+
276
+ for block in self.dit.blocks:
277
+ block.self_attn.forward = types.MethodType(usp_attn_forward, block.self_attn)
278
+ self.dit.forward = types.MethodType(usp_dit_forward, self.dit)
279
+ if self.dit2 is not None:
280
+ for block in self.dit2.blocks:
281
+ block.self_attn.forward = types.MethodType(usp_attn_forward, block.self_attn)
282
+ self.dit2.forward = types.MethodType(usp_dit_forward, self.dit2)
283
+ self.sp_size = get_sequence_parallel_world_size()
284
+ self.use_unified_sequence_parallel = True
285
+
286
+
287
+ @staticmethod
288
+ def from_pretrained(
289
+ torch_dtype: torch.dtype = torch.bfloat16,
290
+ device: Union[str, torch.device] = "cuda",
291
+ model_configs: list[ModelConfig] = [],
292
+ tokenizer_config: ModelConfig = ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="google/*"),
293
+ redirect_common_files: bool = True,
294
+ use_usp=False,
295
+ ):
296
+ # Redirect model path
297
+ if redirect_common_files:
298
+ redirect_dict = {
299
+ "models_t5_umt5-xxl-enc-bf16.pth": "Wan-AI/Wan2.1-T2V-1.3B",
300
+ "Wan2.1_VAE.pth": "Wan-AI/Wan2.1-T2V-1.3B",
301
+ "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth": "Wan-AI/Wan2.1-I2V-14B-480P",
302
+ }
303
+ for model_config in model_configs:
304
+ if model_config.origin_file_pattern is None or model_config.model_id is None:
305
+ continue
306
+ if model_config.origin_file_pattern in redirect_dict and model_config.model_id != redirect_dict[model_config.origin_file_pattern]:
307
+ print(f"To avoid repeatedly downloading model files, ({model_config.model_id}, {model_config.origin_file_pattern}) is redirected to ({redirect_dict[model_config.origin_file_pattern]}, {model_config.origin_file_pattern}). You can use `redirect_common_files=False` to disable file redirection.")
308
+ model_config.model_id = redirect_dict[model_config.origin_file_pattern]
309
+
310
+ # Initialize pipeline
311
+ pipe = WanVideoPipeline(device=device, torch_dtype=torch_dtype)
312
+ if use_usp: pipe.initialize_usp()
313
+
314
+ # Download and load models
315
+ model_manager = ModelManager()
316
+ for model_config in model_configs:
317
+ model_config.download_if_necessary(use_usp=use_usp)
318
+ model_manager.load_model(
319
+ model_config.path,
320
+ device=model_config.offload_device or device,
321
+ torch_dtype=model_config.offload_dtype or torch_dtype
322
+ )
323
+
324
+ # Load models
325
+ pipe.text_encoder = model_manager.fetch_model("wan_video_text_encoder")
326
+ dit = model_manager.fetch_model("wan_video_dit", index=2)
327
+ if isinstance(dit, list):
328
+ pipe.dit, pipe.dit2 = dit
329
+ else:
330
+ pipe.dit = dit
331
+ pipe.vae = model_manager.fetch_model("wan_video_vae")
332
+ pipe.image_encoder = model_manager.fetch_model("wan_video_image_encoder")
333
+ pipe.motion_controller = model_manager.fetch_model("wan_video_motion_controller")
334
+ pipe.vace = model_manager.fetch_model("wan_video_vace")
335
+
336
+ # Size division factor
337
+ if pipe.vae is not None:
338
+ pipe.height_division_factor = pipe.vae.upsampling_factor * 2
339
+ pipe.width_division_factor = pipe.vae.upsampling_factor * 2
340
+
341
+ # Initialize tokenizer
342
+ tokenizer_config.download_if_necessary(use_usp=use_usp)
343
+ pipe.prompter.fetch_models(pipe.text_encoder)
344
+ pipe.prompter.fetch_tokenizer(tokenizer_config.path)
345
+
346
+ # Unified Sequence Parallel
347
+ if use_usp: pipe.enable_usp()
348
+ return pipe
349
+
350
+
351
+ @torch.no_grad()
352
+ def __call__(
353
+ self,
354
+ # Prompt
355
+ prompt: str,
356
+ negative_prompt: Optional[str] = "",
357
+ # Image-to-video
358
+ input_image: Optional[Image.Image] = None,
359
+ # First-last-frame-to-video
360
+ end_image: Optional[Image.Image] = None,
361
+ # Video-to-video
362
+ input_video: Optional[list[Image.Image]] = None,
363
+ denoising_strength: Optional[float] = 1.0,
364
+ # ControlNet
365
+ control_video: Optional[list[Image.Image]] = None,
366
+ reference_image: Optional[Image.Image] = None,
367
+ # Camera control
368
+ camera_control_direction: Optional[Literal["Left", "Right", "Up", "Down", "LeftUp", "LeftDown", "RightUp", "RightDown"]] = None,
369
+ camera_control_speed: Optional[float] = 1/54,
370
+ camera_control_origin: Optional[tuple] = (0, 0.532139961, 0.946026558, 0.5, 0.5, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0),
371
+ # VACE
372
+ vace_video: Optional[list[Image.Image]] = None,
373
+ vace_video_mask: Optional[Image.Image] = None,
374
+ vace_reference_image: Optional[Image.Image] = None,
375
+ vace_scale: Optional[float] = 1.0,
376
+ # Randomness
377
+ seed: Optional[int] = None,
378
+ rand_device: Optional[str] = "cpu",
379
+ # Shape
380
+ height: Optional[int] = 480,
381
+ width: Optional[int] = 832,
382
+ num_frames=81,
383
+ # Classifier-free guidance
384
+ cfg_scale: Optional[float] = 5.0,
385
+ cfg_merge: Optional[bool] = False,
386
+ # Boundary
387
+ switch_DiT_boundary: Optional[float] = 0.875,
388
+ # Scheduler
389
+ num_inference_steps: Optional[int] = 50,
390
+ sigma_shift: Optional[float] = 5.0,
391
+ # Speed control
392
+ motion_bucket_id: Optional[int] = None,
393
+ # VAE tiling
394
+ tiled: Optional[bool] = True,
395
+ tile_size: Optional[tuple[int, int]] = (30, 52),
396
+ tile_stride: Optional[tuple[int, int]] = (15, 26),
397
+ # Sliding window
398
+ sliding_window_size: Optional[int] = None,
399
+ sliding_window_stride: Optional[int] = None,
400
+ # Teacache
401
+ tea_cache_l1_thresh: Optional[float] = None,
402
+ tea_cache_model_id: Optional[str] = "",
403
+ # progress_bar
404
+ progress_bar_cmd=tqdm,
405
+ ):
406
+ # Scheduler
407
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, shift=sigma_shift)
408
+
409
+ # Inputs
410
+ inputs_posi = {
411
+ "prompt": prompt,
412
+ "tea_cache_l1_thresh": tea_cache_l1_thresh, "tea_cache_model_id": tea_cache_model_id, "num_inference_steps": num_inference_steps,
413
+ }
414
+ inputs_nega = {
415
+ "negative_prompt": negative_prompt,
416
+ "tea_cache_l1_thresh": tea_cache_l1_thresh, "tea_cache_model_id": tea_cache_model_id, "num_inference_steps": num_inference_steps,
417
+ }
418
+ inputs_shared = {
419
+ "input_image": input_image,
420
+ "end_image": end_image,
421
+ "input_video": input_video, "denoising_strength": denoising_strength,
422
+ "control_video": control_video, "reference_image": reference_image,
423
+ "camera_control_direction": camera_control_direction, "camera_control_speed": camera_control_speed, "camera_control_origin": camera_control_origin,
424
+ "vace_video": vace_video, "vace_video_mask": vace_video_mask, "vace_reference_image": vace_reference_image, "vace_scale": vace_scale,
425
+ "seed": seed, "rand_device": rand_device,
426
+ "height": height, "width": width, "num_frames": num_frames,
427
+ "cfg_scale": cfg_scale, "cfg_merge": cfg_merge,
428
+ "sigma_shift": sigma_shift,
429
+ "motion_bucket_id": motion_bucket_id,
430
+ "tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride,
431
+ "sliding_window_size": sliding_window_size, "sliding_window_stride": sliding_window_stride,
432
+ }
433
+ for unit in self.units:
434
+ inputs_shared, inputs_posi, inputs_nega = self.unit_runner(unit, self, inputs_shared, inputs_posi, inputs_nega)
435
+
436
+ # Denoise
437
+ self.load_models_to_device(self.in_iteration_models)
438
+ models = {name: getattr(self, name) for name in self.in_iteration_models}
439
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
440
+ # Switch DiT if necessary
441
+ if timestep.item() < switch_DiT_boundary * self.scheduler.num_train_timesteps and self.dit2 is not None and not models["dit"] is self.dit2:
442
+ self.load_models_to_device(self.in_iteration_models_2)
443
+ models["dit"] = self.dit2
444
+
445
+ # Timestep
446
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
447
+
448
+ # Inference
449
+ noise_pred_posi = self.model_fn(**models, **inputs_shared, **inputs_posi, timestep=timestep)
450
+ if cfg_scale != 1.0:
451
+ if cfg_merge:
452
+ noise_pred_posi, noise_pred_nega = noise_pred_posi.chunk(2, dim=0)
453
+ else:
454
+ noise_pred_nega = self.model_fn(**models, **inputs_shared, **inputs_nega, timestep=timestep)
455
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
456
+ else:
457
+ noise_pred = noise_pred_posi
458
+
459
+ # Scheduler
460
+ inputs_shared["latents"] = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], inputs_shared["latents"])
461
+ if "first_frame_latents" in inputs_shared:
462
+ inputs_shared["latents"][:, :, 0:1] = inputs_shared["first_frame_latents"]
463
+
464
+ # VACE (TODO: remove it)
465
+ if vace_reference_image is not None:
466
+ inputs_shared["latents"] = inputs_shared["latents"][:, :, 1:]
467
+
468
+ # Decode
469
+ self.load_models_to_device(['vae'])
470
+ video = self.vae.decode(inputs_shared["latents"], device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
471
+ video = self.vae_output_to_video(video)
472
+ self.load_models_to_device([])
473
+
474
+ return video
475
+
476
+
477
+
478
+ class WanVideoUnit_ShapeChecker(PipelineUnit):
479
+ def __init__(self):
480
+ super().__init__(input_params=("height", "width", "num_frames"))
481
+
482
+ def process(self, pipe: WanVideoPipeline, height, width, num_frames):
483
+ height, width, num_frames = pipe.check_resize_height_width(height, width, num_frames)
484
+ return {"height": height, "width": width, "num_frames": num_frames}
485
+
486
+
487
+
488
+ class WanVideoUnit_NoiseInitializer(PipelineUnit):
489
+ def __init__(self):
490
+ super().__init__(input_params=("height", "width", "num_frames", "seed", "rand_device", "vace_reference_image"))
491
+
492
+ def process(self, pipe: WanVideoPipeline, height, width, num_frames, seed, rand_device, vace_reference_image):
493
+ length = (num_frames - 1) // 4 + 1
494
+ if vace_reference_image is not None:
495
+ length += 1
496
+ shape = (1, pipe.vae.model.z_dim, length, height // pipe.vae.upsampling_factor, width // pipe.vae.upsampling_factor)
497
+ noise = pipe.generate_noise(shape, seed=seed, rand_device=rand_device)
498
+ if vace_reference_image is not None:
499
+ noise = torch.concat((noise[:, :, -1:], noise[:, :, :-1]), dim=2)
500
+ return {"noise": noise}
501
+
502
+
503
+
504
+ class WanVideoUnit_InputVideoEmbedder(PipelineUnit):
505
+ def __init__(self):
506
+ super().__init__(
507
+ input_params=("input_video", "noise", "tiled", "tile_size", "tile_stride", "vace_reference_image"),
508
+ onload_model_names=("vae",)
509
+ )
510
+
511
+ def process(self, pipe: WanVideoPipeline, input_video, noise, tiled, tile_size, tile_stride, vace_reference_image):
512
+ if input_video is None:
513
+ return {"latents": noise}
514
+ pipe.load_models_to_device(["vae"])
515
+ input_video = pipe.preprocess_video(input_video)
516
+ input_latents = pipe.vae.encode(input_video, device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=pipe.torch_dtype, device=pipe.device)
517
+ if vace_reference_image is not None:
518
+ vace_reference_image = pipe.preprocess_video([vace_reference_image])
519
+ vace_reference_latents = pipe.vae.encode(vace_reference_image, device=pipe.device).to(dtype=pipe.torch_dtype, device=pipe.device)
520
+ input_latents = torch.concat([vace_reference_latents, input_latents], dim=2)
521
+ if pipe.scheduler.training:
522
+ return {"latents": noise, "input_latents": input_latents}
523
+ else:
524
+ latents = pipe.scheduler.add_noise(input_latents, noise, timestep=pipe.scheduler.timesteps[0])
525
+ return {"latents": latents}
526
+
527
+
528
+
529
+ class WanVideoUnit_PromptEmbedder(PipelineUnit):
530
+ def __init__(self):
531
+ super().__init__(
532
+ seperate_cfg=True,
533
+ input_params_posi={"prompt": "prompt", "positive": "positive"},
534
+ input_params_nega={"prompt": "negative_prompt", "positive": "positive"},
535
+ onload_model_names=("text_encoder",)
536
+ )
537
+
538
+ def process(self, pipe: WanVideoPipeline, prompt, positive) -> dict:
539
+ pipe.load_models_to_device(self.onload_model_names)
540
+ prompt_emb = pipe.prompter.encode_prompt(prompt, positive=positive, device=pipe.device)
541
+ return {"context": prompt_emb}
542
+
543
+
544
+
545
+ class WanVideoUnit_ImageEmbedder(PipelineUnit):
546
+ """
547
+ Deprecated
548
+ """
549
+ def __init__(self):
550
+ super().__init__(
551
+ input_params=("input_image", "end_image", "num_frames", "height", "width", "tiled", "tile_size", "tile_stride"),
552
+ onload_model_names=("image_encoder", "vae")
553
+ )
554
+
555
+ def process(self, pipe: WanVideoPipeline, input_image, end_image, num_frames, height, width, tiled, tile_size, tile_stride):
556
+ if input_image is None or pipe.image_encoder is None:
557
+ return {}
558
+ pipe.load_models_to_device(self.onload_model_names)
559
+ image = pipe.preprocess_image(input_image.resize((width, height))).to(pipe.device)
560
+ clip_context = pipe.image_encoder.encode_image([image])
561
+ msk = torch.ones(1, num_frames, height//8, width//8, device=pipe.device)
562
+ msk[:, 1:] = 0
563
+ if end_image is not None:
564
+ end_image = pipe.preprocess_image(end_image.resize((width, height))).to(pipe.device)
565
+ vae_input = torch.concat([image.transpose(0,1), torch.zeros(3, num_frames-2, height, width).to(image.device), end_image.transpose(0,1)],dim=1)
566
+ if pipe.dit.has_image_pos_emb:
567
+ clip_context = torch.concat([clip_context, pipe.image_encoder.encode_image([end_image])], dim=1)
568
+ msk[:, -1:] = 1
569
+ else:
570
+ vae_input = torch.concat([image.transpose(0, 1), torch.zeros(3, num_frames-1, height, width).to(image.device)], dim=1)
571
+
572
+ msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
573
+ msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
574
+ msk = msk.transpose(1, 2)[0]
575
+
576
+ y = pipe.vae.encode([vae_input.to(dtype=pipe.torch_dtype, device=pipe.device)], device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)[0]
577
+ y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
578
+ y = torch.concat([msk, y])
579
+ y = y.unsqueeze(0)
580
+ clip_context = clip_context.to(dtype=pipe.torch_dtype, device=pipe.device)
581
+ y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
582
+ return {"clip_feature": clip_context, "y": y}
583
+
584
+
585
+
586
+ class WanVideoUnit_ImageEmbedderCLIP(PipelineUnit):
587
+ def __init__(self):
588
+ super().__init__(
589
+ input_params=("input_image", "end_image", "height", "width"),
590
+ onload_model_names=("image_encoder",)
591
+ )
592
+
593
+ def process(self, pipe: WanVideoPipeline, input_image, end_image, height, width):
594
+ if input_image is None or pipe.image_encoder is None or not pipe.dit.require_clip_embedding:
595
+ return {}
596
+ pipe.load_models_to_device(self.onload_model_names)
597
+ image = pipe.preprocess_image(input_image.resize((width, height))).to(pipe.device)
598
+ clip_context = pipe.image_encoder.encode_image([image])
599
+ if end_image is not None:
600
+ end_image = pipe.preprocess_image(end_image.resize((width, height))).to(pipe.device)
601
+ if pipe.dit.has_image_pos_emb:
602
+ clip_context = torch.concat([clip_context, pipe.image_encoder.encode_image([end_image])], dim=1)
603
+ clip_context = clip_context.to(dtype=pipe.torch_dtype, device=pipe.device)
604
+ return {"clip_feature": clip_context}
605
+
606
+
607
+
608
+ class WanVideoUnit_ImageEmbedderVAE(PipelineUnit):
609
+ def __init__(self):
610
+ super().__init__(
611
+ input_params=("input_image", "end_image", "num_frames", "height", "width", "tiled", "tile_size", "tile_stride"),
612
+ onload_model_names=("vae",)
613
+ )
614
+
615
+ def process(self, pipe: WanVideoPipeline, input_image, end_image, num_frames, height, width, tiled, tile_size, tile_stride):
616
+ if input_image is None or not pipe.dit.require_vae_embedding:
617
+ return {}
618
+ pipe.load_models_to_device(self.onload_model_names)
619
+ image = pipe.preprocess_image(input_image.resize((width, height))).to(pipe.device)
620
+ msk = torch.ones(1, num_frames, height//8, width//8, device=pipe.device)
621
+ msk[:, 1:] = 0
622
+ if end_image is not None:
623
+ end_image = pipe.preprocess_image(end_image.resize((width, height))).to(pipe.device)
624
+ vae_input = torch.concat([image.transpose(0,1), torch.zeros(3, num_frames-2, height, width).to(image.device), end_image.transpose(0,1)],dim=1)
625
+ msk[:, -1:] = 1
626
+ else:
627
+ vae_input = torch.concat([image.transpose(0, 1), torch.zeros(3, num_frames-1, height, width).to(image.device)], dim=1)
628
+
629
+ msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
630
+ msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
631
+ msk = msk.transpose(1, 2)[0]
632
+
633
+ y = pipe.vae.encode([vae_input.to(dtype=pipe.torch_dtype, device=pipe.device)], device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)[0]
634
+ y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
635
+ y = torch.concat([msk, y])
636
+ y = y.unsqueeze(0)
637
+ y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
638
+ return {"y": y}
639
+
640
+
641
+
642
+ class WanVideoUnit_ImageEmbedderFused(PipelineUnit):
643
+ """
644
+ Encode input image to latents using VAE. This unit is for Wan-AI/Wan2.2-TI2V-5B.
645
+ """
646
+ def __init__(self):
647
+ super().__init__(
648
+ input_params=("input_image", "latents", "height", "width", "tiled", "tile_size", "tile_stride"),
649
+ onload_model_names=("vae",)
650
+ )
651
+
652
+ def process(self, pipe: WanVideoPipeline, input_image, latents, height, width, tiled, tile_size, tile_stride):
653
+ if input_image is None or not pipe.dit.fuse_vae_embedding_in_latents:
654
+ return {}
655
+ pipe.load_models_to_device(self.onload_model_names)
656
+ image = pipe.preprocess_image(input_image.resize((width, height))).transpose(0, 1)
657
+ z = pipe.vae.encode([image], device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
658
+ latents[:, :, 0: 1] = z
659
+ return {"latents": latents, "fuse_vae_embedding_in_latents": True, "first_frame_latents": z}
660
+
661
+
662
+
663
+ class WanVideoUnit_FunControl(PipelineUnit):
664
+ def __init__(self):
665
+ super().__init__(
666
+ input_params=("control_video", "num_frames", "height", "width", "tiled", "tile_size", "tile_stride", "clip_feature", "y"),
667
+ onload_model_names=("vae",)
668
+ )
669
+
670
+ def process(self, pipe: WanVideoPipeline, control_video, num_frames, height, width, tiled, tile_size, tile_stride, clip_feature, y):
671
+ if control_video is None:
672
+ return {}
673
+ pipe.load_models_to_device(self.onload_model_names)
674
+ control_video = pipe.preprocess_video(control_video)
675
+ control_latents = pipe.vae.encode(control_video, device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=pipe.torch_dtype, device=pipe.device)
676
+ control_latents = control_latents.to(dtype=pipe.torch_dtype, device=pipe.device)
677
+ if clip_feature is None or y is None:
678
+ clip_feature = torch.zeros((1, 257, 1280), dtype=pipe.torch_dtype, device=pipe.device)
679
+ y = torch.zeros((1, 16, (num_frames - 1) // 4 + 1, height//8, width//8), dtype=pipe.torch_dtype, device=pipe.device)
680
+ else:
681
+ y = y[:, -16:]
682
+ y = torch.concat([control_latents, y], dim=1)
683
+ return {"clip_feature": clip_feature, "y": y}
684
+
685
+
686
+
687
+ class WanVideoUnit_FunReference(PipelineUnit):
688
+ def __init__(self):
689
+ super().__init__(
690
+ input_params=("reference_image", "height", "width", "reference_image"),
691
+ onload_model_names=("vae",)
692
+ )
693
+
694
+ def process(self, pipe: WanVideoPipeline, reference_image, height, width):
695
+ if reference_image is None:
696
+ return {}
697
+ pipe.load_models_to_device(["vae"])
698
+ reference_image = reference_image.resize((width, height))
699
+ reference_latents = pipe.preprocess_video([reference_image])
700
+ reference_latents = pipe.vae.encode(reference_latents, device=pipe.device)
701
+ clip_feature = pipe.preprocess_image(reference_image)
702
+ clip_feature = pipe.image_encoder.encode_image([clip_feature])
703
+ return {"reference_latents": reference_latents, "clip_feature": clip_feature}
704
+
705
+
706
+
707
+ class WanVideoUnit_FunCameraControl(PipelineUnit):
708
+ def __init__(self):
709
+ super().__init__(
710
+ input_params=("height", "width", "num_frames", "camera_control_direction", "camera_control_speed", "camera_control_origin", "latents", "input_image"),
711
+ onload_model_names=("vae",)
712
+ )
713
+
714
+ def process(self, pipe: WanVideoPipeline, height, width, num_frames, camera_control_direction, camera_control_speed, camera_control_origin, latents, input_image):
715
+ if camera_control_direction is None:
716
+ return {}
717
+ camera_control_plucker_embedding = pipe.dit.control_adapter.process_camera_coordinates(
718
+ camera_control_direction, num_frames, height, width, camera_control_speed, camera_control_origin)
719
+
720
+ control_camera_video = camera_control_plucker_embedding[:num_frames].permute([3, 0, 1, 2]).unsqueeze(0)
721
+ control_camera_latents = torch.concat(
722
+ [
723
+ torch.repeat_interleave(control_camera_video[:, :, 0:1], repeats=4, dim=2),
724
+ control_camera_video[:, :, 1:]
725
+ ], dim=2
726
+ ).transpose(1, 2)
727
+ b, f, c, h, w = control_camera_latents.shape
728
+ control_camera_latents = control_camera_latents.contiguous().view(b, f // 4, 4, c, h, w).transpose(2, 3)
729
+ control_camera_latents = control_camera_latents.contiguous().view(b, f // 4, c * 4, h, w).transpose(1, 2)
730
+ control_camera_latents_input = control_camera_latents.to(device=pipe.device, dtype=pipe.torch_dtype)
731
+
732
+ input_image = input_image.resize((width, height))
733
+ input_latents = pipe.preprocess_video([input_image])
734
+ pipe.load_models_to_device(self.onload_model_names)
735
+ input_latents = pipe.vae.encode(input_latents, device=pipe.device)
736
+ y = torch.zeros_like(latents).to(pipe.device)
737
+ y[:, :, :1] = input_latents
738
+ y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
739
+ return {"control_camera_latents_input": control_camera_latents_input, "y": y}
740
+
741
+
742
+
743
+ class WanVideoUnit_SpeedControl(PipelineUnit):
744
+ def __init__(self):
745
+ super().__init__(input_params=("motion_bucket_id",))
746
+
747
+ def process(self, pipe: WanVideoPipeline, motion_bucket_id):
748
+ if motion_bucket_id is None:
749
+ return {}
750
+ motion_bucket_id = torch.Tensor((motion_bucket_id,)).to(dtype=pipe.torch_dtype, device=pipe.device)
751
+ return {"motion_bucket_id": motion_bucket_id}
752
+
753
+
754
+
755
+ class WanVideoUnit_VACE(PipelineUnit):
756
+ def __init__(self):
757
+ super().__init__(
758
+ input_params=("vace_video", "vace_video_mask", "vace_reference_image", "vace_scale", "height", "width", "num_frames", "tiled", "tile_size", "tile_stride"),
759
+ onload_model_names=("vae",)
760
+ )
761
+
762
+ def process(
763
+ self,
764
+ pipe: WanVideoPipeline,
765
+ vace_video, vace_video_mask, vace_reference_image, vace_scale,
766
+ height, width, num_frames,
767
+ tiled, tile_size, tile_stride
768
+ ):
769
+ if vace_video is not None or vace_video_mask is not None or vace_reference_image is not None:
770
+ pipe.load_models_to_device(["vae"])
771
+ if vace_video is None:
772
+ vace_video = torch.zeros((1, 3, num_frames, height, width), dtype=pipe.torch_dtype, device=pipe.device)
773
+ else:
774
+ vace_video = pipe.preprocess_video(vace_video)
775
+
776
+ if vace_video_mask is None:
777
+ vace_video_mask = torch.ones_like(vace_video)
778
+ else:
779
+ vace_video_mask = pipe.preprocess_video(vace_video_mask, min_value=0, max_value=1)
780
+
781
+ inactive = vace_video * (1 - vace_video_mask) + 0 * vace_video_mask
782
+ reactive = vace_video * vace_video_mask + 0 * (1 - vace_video_mask)
783
+ inactive = pipe.vae.encode(inactive, device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=pipe.torch_dtype, device=pipe.device)
784
+ reactive = pipe.vae.encode(reactive, device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=pipe.torch_dtype, device=pipe.device)
785
+ vace_video_latents = torch.concat((inactive, reactive), dim=1)
786
+
787
+ vace_mask_latents = rearrange(vace_video_mask[0,0], "T (H P) (W Q) -> 1 (P Q) T H W", P=8, Q=8)
788
+ vace_mask_latents = torch.nn.functional.interpolate(vace_mask_latents, size=((vace_mask_latents.shape[2] + 3) // 4, vace_mask_latents.shape[3], vace_mask_latents.shape[4]), mode='nearest-exact')
789
+
790
+ if vace_reference_image is None:
791
+ pass
792
+ else:
793
+ vace_reference_image = pipe.preprocess_video([vace_reference_image])
794
+ vace_reference_latents = pipe.vae.encode(vace_reference_image, device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=pipe.torch_dtype, device=pipe.device)
795
+ vace_reference_latents = torch.concat((vace_reference_latents, torch.zeros_like(vace_reference_latents)), dim=1)
796
+ vace_video_latents = torch.concat((vace_reference_latents, vace_video_latents), dim=2)
797
+ vace_mask_latents = torch.concat((torch.zeros_like(vace_mask_latents[:, :, :1]), vace_mask_latents), dim=2)
798
+
799
+ vace_context = torch.concat((vace_video_latents, vace_mask_latents), dim=1)
800
+ return {"vace_context": vace_context, "vace_scale": vace_scale}
801
+ else:
802
+ return {"vace_context": None, "vace_scale": vace_scale}
803
+
804
+
805
+
806
+ class WanVideoUnit_UnifiedSequenceParallel(PipelineUnit):
807
+ def __init__(self):
808
+ super().__init__(input_params=())
809
+
810
+ def process(self, pipe: WanVideoPipeline):
811
+ if hasattr(pipe, "use_unified_sequence_parallel"):
812
+ if pipe.use_unified_sequence_parallel:
813
+ return {"use_unified_sequence_parallel": True}
814
+ return {}
815
+
816
+
817
+
818
+ class WanVideoUnit_TeaCache(PipelineUnit):
819
+ def __init__(self):
820
+ super().__init__(
821
+ seperate_cfg=True,
822
+ input_params_posi={"num_inference_steps": "num_inference_steps", "tea_cache_l1_thresh": "tea_cache_l1_thresh", "tea_cache_model_id": "tea_cache_model_id"},
823
+ input_params_nega={"num_inference_steps": "num_inference_steps", "tea_cache_l1_thresh": "tea_cache_l1_thresh", "tea_cache_model_id": "tea_cache_model_id"},
824
+ )
825
+
826
+ def process(self, pipe: WanVideoPipeline, num_inference_steps, tea_cache_l1_thresh, tea_cache_model_id):
827
+ if tea_cache_l1_thresh is None:
828
+ return {}
829
+ return {"tea_cache": TeaCache(num_inference_steps, rel_l1_thresh=tea_cache_l1_thresh, model_id=tea_cache_model_id)}
830
+
831
+
832
+
833
+ class WanVideoUnit_CfgMerger(PipelineUnit):
834
+ def __init__(self):
835
+ super().__init__(take_over=True)
836
+ self.concat_tensor_names = ["context", "clip_feature", "y", "reference_latents"]
837
+
838
+ def process(self, pipe: WanVideoPipeline, inputs_shared, inputs_posi, inputs_nega):
839
+ if not inputs_shared["cfg_merge"]:
840
+ return inputs_shared, inputs_posi, inputs_nega
841
+ for name in self.concat_tensor_names:
842
+ tensor_posi = inputs_posi.get(name)
843
+ tensor_nega = inputs_nega.get(name)
844
+ tensor_shared = inputs_shared.get(name)
845
+ if tensor_posi is not None and tensor_nega is not None:
846
+ inputs_shared[name] = torch.concat((tensor_posi, tensor_nega), dim=0)
847
+ elif tensor_shared is not None:
848
+ inputs_shared[name] = torch.concat((tensor_shared, tensor_shared), dim=0)
849
+ inputs_posi.clear()
850
+ inputs_nega.clear()
851
+ return inputs_shared, inputs_posi, inputs_nega
852
+
853
+
854
+
855
+ class TeaCache:
856
+ def __init__(self, num_inference_steps, rel_l1_thresh, model_id):
857
+ self.num_inference_steps = num_inference_steps
858
+ self.step = 0
859
+ self.accumulated_rel_l1_distance = 0
860
+ self.previous_modulated_input = None
861
+ self.rel_l1_thresh = rel_l1_thresh
862
+ self.previous_residual = None
863
+ self.previous_hidden_states = None
864
+
865
+ self.coefficients_dict = {
866
+ "Wan2.1-T2V-1.3B": [-5.21862437e+04, 9.23041404e+03, -5.28275948e+02, 1.36987616e+01, -4.99875664e-02],
867
+ "Wan2.1-T2V-14B": [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01],
868
+ "Wan2.1-I2V-14B-480P": [2.57151496e+05, -3.54229917e+04, 1.40286849e+03, -1.35890334e+01, 1.32517977e-01],
869
+ "Wan2.1-I2V-14B-720P": [ 8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02],
870
+ }
871
+ if model_id not in self.coefficients_dict:
872
+ supported_model_ids = ", ".join([i for i in self.coefficients_dict])
873
+ raise ValueError(f"{model_id} is not a supported TeaCache model id. Please choose a valid model id in ({supported_model_ids}).")
874
+ self.coefficients = self.coefficients_dict[model_id]
875
+
876
+ def check(self, dit: WanModel, x, t_mod):
877
+ modulated_inp = t_mod.clone()
878
+ if self.step == 0 or self.step == self.num_inference_steps - 1:
879
+ should_calc = True
880
+ self.accumulated_rel_l1_distance = 0
881
+ else:
882
+ coefficients = self.coefficients
883
+ rescale_func = np.poly1d(coefficients)
884
+ self.accumulated_rel_l1_distance += rescale_func(((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()).cpu().item())
885
+ if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
886
+ should_calc = False
887
+ else:
888
+ should_calc = True
889
+ self.accumulated_rel_l1_distance = 0
890
+ self.previous_modulated_input = modulated_inp
891
+ self.step += 1
892
+ if self.step == self.num_inference_steps:
893
+ self.step = 0
894
+ if should_calc:
895
+ self.previous_hidden_states = x.clone()
896
+ return not should_calc
897
+
898
+ def store(self, hidden_states):
899
+ self.previous_residual = hidden_states - self.previous_hidden_states
900
+ self.previous_hidden_states = None
901
+
902
+ def update(self, hidden_states):
903
+ hidden_states = hidden_states + self.previous_residual
904
+ return hidden_states
905
+
906
+
907
+
908
+ class TemporalTiler_BCTHW:
909
+ def __init__(self):
910
+ pass
911
+
912
+ def build_1d_mask(self, length, left_bound, right_bound, border_width):
913
+ x = torch.ones((length,))
914
+ if border_width == 0:
915
+ return x
916
+
917
+ shift = 0.5
918
+ if not left_bound:
919
+ x[:border_width] = (torch.arange(border_width) + shift) / border_width
920
+ if not right_bound:
921
+ x[-border_width:] = torch.flip((torch.arange(border_width) + shift) / border_width, dims=(0,))
922
+ return x
923
+
924
+ def build_mask(self, data, is_bound, border_width):
925
+ _, _, T, _, _ = data.shape
926
+ t = self.build_1d_mask(T, is_bound[0], is_bound[1], border_width[0])
927
+ mask = repeat(t, "T -> 1 1 T 1 1")
928
+ return mask
929
+
930
+ def run(self, model_fn, sliding_window_size, sliding_window_stride, computation_device, computation_dtype, model_kwargs, tensor_names, batch_size=None):
931
+ tensor_names = [tensor_name for tensor_name in tensor_names if model_kwargs.get(tensor_name) is not None]
932
+ tensor_dict = {tensor_name: model_kwargs[tensor_name] for tensor_name in tensor_names}
933
+ B, C, T, H, W = tensor_dict[tensor_names[0]].shape
934
+ if batch_size is not None:
935
+ B *= batch_size
936
+ data_device, data_dtype = tensor_dict[tensor_names[0]].device, tensor_dict[tensor_names[0]].dtype
937
+ value = torch.zeros((B, C, T, H, W), device=data_device, dtype=data_dtype)
938
+ weight = torch.zeros((1, 1, T, 1, 1), device=data_device, dtype=data_dtype)
939
+ for t in range(0, T, sliding_window_stride):
940
+ if t - sliding_window_stride >= 0 and t - sliding_window_stride + sliding_window_size >= T:
941
+ continue
942
+ t_ = min(t + sliding_window_size, T)
943
+ model_kwargs.update({
944
+ tensor_name: tensor_dict[tensor_name][:, :, t: t_:, :].to(device=computation_device, dtype=computation_dtype) \
945
+ for tensor_name in tensor_names
946
+ })
947
+ model_output = model_fn(**model_kwargs).to(device=data_device, dtype=data_dtype)
948
+ mask = self.build_mask(
949
+ model_output,
950
+ is_bound=(t == 0, t_ == T),
951
+ border_width=(sliding_window_size - sliding_window_stride,)
952
+ ).to(device=data_device, dtype=data_dtype)
953
+ value[:, :, t: t_, :, :] += model_output * mask
954
+ weight[:, :, t: t_, :, :] += mask
955
+ value /= weight
956
+ model_kwargs.update(tensor_dict)
957
+ return value
958
+
959
+
960
+
961
+ def model_fn_wan_video(
962
+ dit: WanModel,
963
+ motion_controller: WanMotionControllerModel = None,
964
+ vace: VaceWanModel = None,
965
+ latents: torch.Tensor = None,
966
+ timestep: torch.Tensor = None,
967
+ context: torch.Tensor = None,
968
+ clip_feature: Optional[torch.Tensor] = None,
969
+ y: Optional[torch.Tensor] = None,
970
+ reference_latents = None,
971
+ vace_context = None,
972
+ vace_scale = 1.0,
973
+ tea_cache: TeaCache = None,
974
+ use_unified_sequence_parallel: bool = False,
975
+ motion_bucket_id: Optional[torch.Tensor] = None,
976
+ sliding_window_size: Optional[int] = None,
977
+ sliding_window_stride: Optional[int] = None,
978
+ cfg_merge: bool = False,
979
+ use_gradient_checkpointing: bool = False,
980
+ use_gradient_checkpointing_offload: bool = False,
981
+ control_camera_latents_input = None,
982
+ fuse_vae_embedding_in_latents: bool = False,
983
+ **kwargs,
984
+ ):
985
+ if sliding_window_size is not None and sliding_window_stride is not None:
986
+ model_kwargs = dict(
987
+ dit=dit,
988
+ motion_controller=motion_controller,
989
+ vace=vace,
990
+ latents=latents,
991
+ timestep=timestep,
992
+ context=context,
993
+ clip_feature=clip_feature,
994
+ y=y,
995
+ reference_latents=reference_latents,
996
+ vace_context=vace_context,
997
+ vace_scale=vace_scale,
998
+ tea_cache=tea_cache,
999
+ use_unified_sequence_parallel=use_unified_sequence_parallel,
1000
+ motion_bucket_id=motion_bucket_id,
1001
+ )
1002
+ return TemporalTiler_BCTHW().run(
1003
+ model_fn_wan_video,
1004
+ sliding_window_size, sliding_window_stride,
1005
+ latents.device, latents.dtype,
1006
+ model_kwargs=model_kwargs,
1007
+ tensor_names=["latents", "y"],
1008
+ batch_size=2 if cfg_merge else 1
1009
+ )
1010
+
1011
+ if use_unified_sequence_parallel:
1012
+ import torch.distributed as dist
1013
+ from xfuser.core.distributed import (get_sequence_parallel_rank,
1014
+ get_sequence_parallel_world_size,
1015
+ get_sp_group)
1016
+
1017
+ # Timestep
1018
+ if dit.seperated_timestep and fuse_vae_embedding_in_latents:
1019
+ timestep = torch.concat([
1020
+ torch.zeros((1, latents.shape[3] * latents.shape[4] // 4), dtype=latents.dtype, device=latents.device),
1021
+ torch.ones((latents.shape[2] - 1, latents.shape[3] * latents.shape[4] // 4), dtype=latents.dtype, device=latents.device) * timestep
1022
+ ]).flatten()
1023
+ t = dit.time_embedding(sinusoidal_embedding_1d(dit.freq_dim, timestep).unsqueeze(0))
1024
+ t_mod = dit.time_projection(t).unflatten(2, (6, dit.dim))
1025
+ else:
1026
+ t = dit.time_embedding(sinusoidal_embedding_1d(dit.freq_dim, timestep))
1027
+ t_mod = dit.time_projection(t).unflatten(1, (6, dit.dim))
1028
+
1029
+ # Motion Controller
1030
+ if motion_bucket_id is not None and motion_controller is not None:
1031
+ t_mod = t_mod + motion_controller(motion_bucket_id).unflatten(1, (6, dit.dim))
1032
+ context = dit.text_embedding(context)
1033
+
1034
+ x = latents
1035
+ # Merged cfg
1036
+ if x.shape[0] != context.shape[0]:
1037
+ x = torch.concat([x] * context.shape[0], dim=0)
1038
+ if timestep.shape[0] != context.shape[0]:
1039
+ timestep = torch.concat([timestep] * context.shape[0], dim=0)
1040
+
1041
+ # Image Embedding
1042
+ # import pdb; pdb.set_trace()
1043
+ if y is not None and dit.require_vae_embedding: # torch.Size([1, 20, 21, 60, 104])
1044
+ x = torch.cat([x, y], dim=1) # torch.Size([1, 36, 21, 60, 104])
1045
+ if clip_feature is not None and dit.require_clip_embedding:
1046
+ clip_embdding = dit.img_emb(clip_feature)
1047
+ context = torch.cat([clip_embdding, context], dim=1)
1048
+
1049
+ # Add camera control
1050
+ x, (f, h, w) = dit.patchify(x, control_camera_latents_input)
1051
+
1052
+ # Reference image
1053
+ if reference_latents is not None:
1054
+ if len(reference_latents.shape) == 5:
1055
+ reference_latents = reference_latents[:, :, 0]
1056
+ reference_latents = dit.ref_conv(reference_latents).flatten(2).transpose(1, 2)
1057
+ x = torch.concat([reference_latents, x], dim=1)
1058
+ f += 1
1059
+
1060
+ freqs = torch.cat([
1061
+ dit.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
1062
+ dit.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
1063
+ dit.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
1064
+ ], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
1065
+
1066
+ # TeaCache
1067
+ if tea_cache is not None:
1068
+ tea_cache_update = tea_cache.check(dit, x, t_mod)
1069
+ else:
1070
+ tea_cache_update = False
1071
+
1072
+ if vace_context is not None:
1073
+ vace_hints = vace(x, vace_context, context, t_mod, freqs)
1074
+
1075
+ # blocks
1076
+ if use_unified_sequence_parallel:
1077
+ if dist.is_initialized() and dist.get_world_size() > 1:
1078
+ chunks = torch.chunk(x, get_sequence_parallel_world_size(), dim=1)
1079
+ pad_shape = chunks[0].shape[1] - chunks[-1].shape[1]
1080
+ chunks = [torch.nn.functional.pad(chunk, (0, 0, 0, chunks[0].shape[1]-chunk.shape[1]), value=0) for chunk in chunks]
1081
+ x = chunks[get_sequence_parallel_rank()]
1082
+ if tea_cache_update:
1083
+ x = tea_cache.update(x)
1084
+ else:
1085
+ def create_custom_forward(module):
1086
+ def custom_forward(*inputs):
1087
+ return module(*inputs)
1088
+ return custom_forward
1089
+
1090
+ for block_id, block in enumerate(dit.blocks):
1091
+ if use_gradient_checkpointing_offload:
1092
+ with torch.autograd.graph.save_on_cpu():
1093
+ x = torch.utils.checkpoint.checkpoint(
1094
+ create_custom_forward(block),
1095
+ x, context, t_mod, freqs,
1096
+ use_reentrant=False,
1097
+ )
1098
+ elif use_gradient_checkpointing:
1099
+ x = torch.utils.checkpoint.checkpoint(
1100
+ create_custom_forward(block),
1101
+ x, context, t_mod, freqs,
1102
+ use_reentrant=False,
1103
+ )
1104
+ else:
1105
+ x = block(x, context, t_mod, freqs)
1106
+ if vace_context is not None and block_id in vace.vace_layers_mapping:
1107
+ current_vace_hint = vace_hints[vace.vace_layers_mapping[block_id]]
1108
+ if use_unified_sequence_parallel and dist.is_initialized() and dist.get_world_size() > 1:
1109
+ current_vace_hint = torch.chunk(current_vace_hint, get_sequence_parallel_world_size(), dim=1)[get_sequence_parallel_rank()]
1110
+ current_vace_hint = torch.nn.functional.pad(current_vace_hint, (0, 0, 0, chunks[0].shape[1] - current_vace_hint.shape[1]), value=0)
1111
+ x = x + current_vace_hint * vace_scale
1112
+ if tea_cache is not None:
1113
+ tea_cache.store(x)
1114
+
1115
+ x = dit.head(x, t)
1116
+ if use_unified_sequence_parallel:
1117
+ if dist.is_initialized() and dist.get_world_size() > 1:
1118
+ x = get_sp_group().all_gather(x, dim=1)
1119
+ x = x[:, :-pad_shape] if pad_shape > 0 else x
1120
+ # Remove reference latents
1121
+ if reference_latents is not None:
1122
+ x = x[:, reference_latents.shape[1]:]
1123
+ f -= 1
1124
+ x = dit.unpatchify(x, (f, h, w))
1125
+ return x
diffsynth/pipelines/wan_video_relit_live.py ADDED
@@ -0,0 +1,525 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models import ModelManager
2
+ from ..models.wan_video_dit_relit_live import WanModel
3
+ from ..models.wan_video_text_encoder import WanTextEncoder
4
+ from ..models.wan_video_vae import WanVideoVAE
5
+ from ..models.wan_video_image_encoder import WanImageEncoder
6
+ from ..schedulers.flow_match import FlowMatchScheduler
7
+ from .base import BasePipeline
8
+ from ..prompters import WanPrompter
9
+ import torch, os
10
+ from einops import rearrange
11
+ import numpy as np
12
+ from PIL import Image
13
+ from tqdm import tqdm
14
+ from typing import Optional
15
+ import torch.nn.functional as F
16
+
17
+ from ..vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
18
+ from ..models.wan_video_text_encoder import T5RelativeEmbedding, T5LayerNorm
19
+ from ..models.wan_video_dit_relit_live import RMSNorm, sinusoidal_embedding_1d
20
+ from ..models.wan_video_vae import RMS_norm, CausalConv3d, Upsample
21
+
22
+ class WanVideoRelitlivePipeline(BasePipeline):
23
+
24
+ def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None):
25
+ super().__init__(device=device, torch_dtype=torch_dtype)
26
+ self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
27
+ self.prompter = WanPrompter(tokenizer_path=tokenizer_path)
28
+ self.text_encoder: WanTextEncoder = None
29
+ self.image_encoder: WanImageEncoder = None
30
+ self.dit: WanModel = None
31
+ self.vae: WanVideoVAE = None
32
+ self.model_names = ['text_encoder', 'dit', 'vae']
33
+ self.height_division_factor = 16
34
+ self.width_division_factor = 16
35
+
36
+ def enable_vram_management(self, num_persistent_param_in_dit=None):
37
+ dtype = next(iter(self.text_encoder.parameters())).dtype
38
+ enable_vram_management(
39
+ self.text_encoder,
40
+ module_map = {
41
+ torch.nn.Linear: AutoWrappedLinear,
42
+ torch.nn.Embedding: AutoWrappedModule,
43
+ T5RelativeEmbedding: AutoWrappedModule,
44
+ T5LayerNorm: AutoWrappedModule,
45
+ },
46
+ module_config = dict(
47
+ offload_dtype=dtype,
48
+ offload_device="cpu",
49
+ onload_dtype=dtype,
50
+ onload_device="cpu",
51
+ computation_dtype=self.torch_dtype,
52
+ computation_device=self.device,
53
+ ),
54
+ )
55
+ dtype = next(iter(self.dit.parameters())).dtype
56
+ enable_vram_management(
57
+ self.dit,
58
+ module_map = {
59
+ torch.nn.Linear: AutoWrappedLinear,
60
+ torch.nn.Conv3d: AutoWrappedModule,
61
+ torch.nn.LayerNorm: AutoWrappedModule,
62
+ RMSNorm: AutoWrappedModule,
63
+ },
64
+ module_config = dict(
65
+ offload_dtype=dtype,
66
+ offload_device="cpu",
67
+ onload_dtype=dtype,
68
+ onload_device=self.device,
69
+ computation_dtype=self.torch_dtype,
70
+ computation_device=self.device,
71
+ ),
72
+ max_num_param=num_persistent_param_in_dit,
73
+ overflow_module_config = dict(
74
+ offload_dtype=dtype,
75
+ offload_device="cpu",
76
+ onload_dtype=dtype,
77
+ onload_device="cpu",
78
+ computation_dtype=self.torch_dtype,
79
+ computation_device=self.device,
80
+ ),
81
+ )
82
+ dtype = next(iter(self.vae.parameters())).dtype
83
+ enable_vram_management(
84
+ self.vae,
85
+ module_map = {
86
+ torch.nn.Linear: AutoWrappedLinear,
87
+ torch.nn.Conv2d: AutoWrappedModule,
88
+ RMS_norm: AutoWrappedModule,
89
+ CausalConv3d: AutoWrappedModule,
90
+ Upsample: AutoWrappedModule,
91
+ torch.nn.SiLU: AutoWrappedModule,
92
+ torch.nn.Dropout: AutoWrappedModule,
93
+ },
94
+ module_config = dict(
95
+ offload_dtype=dtype,
96
+ offload_device="cpu",
97
+ onload_dtype=dtype,
98
+ onload_device=self.device,
99
+ computation_dtype=self.torch_dtype,
100
+ computation_device=self.device,
101
+ ),
102
+ )
103
+ if self.image_encoder is not None:
104
+ dtype = next(iter(self.image_encoder.parameters())).dtype
105
+ enable_vram_management(
106
+ self.image_encoder,
107
+ module_map = {
108
+ torch.nn.Linear: AutoWrappedLinear,
109
+ torch.nn.Conv2d: AutoWrappedModule,
110
+ torch.nn.LayerNorm: AutoWrappedModule,
111
+ },
112
+ module_config = dict(
113
+ offload_dtype=dtype,
114
+ offload_device="cpu",
115
+ onload_dtype=dtype,
116
+ onload_device="cpu",
117
+ computation_dtype=dtype,
118
+ computation_device=self.device,
119
+ ),
120
+ )
121
+ self.enable_cpu_offload()
122
+
123
+
124
+ def fetch_models(self, model_manager: ModelManager):
125
+ text_encoder_model_and_path = model_manager.fetch_model("wan_video_text_encoder", require_model_path=True)
126
+ if text_encoder_model_and_path is not None:
127
+ self.text_encoder, tokenizer_path = text_encoder_model_and_path
128
+ self.prompter.fetch_models(self.text_encoder)
129
+ self.prompter.fetch_tokenizer(os.path.join(os.path.dirname(tokenizer_path), "google/umt5-xxl"))
130
+ self.dit = model_manager.fetch_model("wan_video_dit")
131
+ self.vae = model_manager.fetch_model("wan_video_vae")
132
+ self.image_encoder = model_manager.fetch_model("wan_video_image_encoder")
133
+
134
+
135
+ @staticmethod
136
+ def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None):
137
+ if device is None: device = model_manager.device
138
+ if torch_dtype is None: torch_dtype = model_manager.torch_dtype
139
+ pipe = WanVideoRelitlivePipeline(device=device, torch_dtype=torch_dtype)
140
+ pipe.fetch_models(model_manager)
141
+ return pipe
142
+
143
+
144
+ def denoising_model(self):
145
+ return self.dit
146
+
147
+
148
+ def encode_prompt(self, prompt, positive=True):
149
+ prompt_emb = self.prompter.encode_prompt(prompt, positive=positive)
150
+ return {"context": prompt_emb}
151
+
152
+
153
+ def encode_image(self, image, num_frames, height, width):
154
+ image = self.preprocess_image(image.resize((width, height))).to(self.device)
155
+ clip_context = self.image_encoder.encode_image([image])
156
+ msk = torch.ones(1, num_frames, height//8, width//8, device=self.device)
157
+ msk[:, 1:] = 0
158
+ msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
159
+ msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
160
+ msk = msk.transpose(1, 2)[0]
161
+
162
+ vae_input = torch.concat([image.transpose(0, 1), torch.zeros(3, num_frames-1, height, width).to(image.device)], dim=1)
163
+ y = self.vae.encode([vae_input.to(dtype=self.torch_dtype, device=self.device)], device=self.device)[0]
164
+ y = torch.concat([msk, y])
165
+ y = y.unsqueeze(0)
166
+ clip_context = clip_context.to(dtype=self.torch_dtype, device=self.device)
167
+ y = y.to(dtype=self.torch_dtype, device=self.device)
168
+ return {"clip_feature": clip_context, "y": y}
169
+
170
+
171
+ def tensor2video(self, frames, return_float_array=False):
172
+ frames = rearrange(frames, "C T H W -> T H W C") # (-1, 1)
173
+ if return_float_array:
174
+ frames_float_array = ((frames.float() + 1) * 0.5).clip(0, 1).cpu().numpy().astype(np.float32)
175
+ frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
176
+ frames = [Image.fromarray(frame) for frame in frames]
177
+ if return_float_array:
178
+ return frames, frames_float_array
179
+ return frames
180
+
181
+
182
+ def prepare_extra_input(self, latents=None):
183
+ return {}
184
+
185
+
186
+ def encode_video(self, input_video, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
187
+ latents = self.vae.encode(input_video, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
188
+ return latents
189
+
190
+
191
+ def decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
192
+ frames = self.vae.decode(latents, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
193
+ return frames
194
+
195
+
196
+ @torch.no_grad()
197
+ def __call__(
198
+ self,
199
+ prompt,
200
+ negative_prompt="",
201
+ batch=None,
202
+ input_image=None,
203
+ input_video=None,
204
+ denoising_strength=1.0,
205
+ seed=None,
206
+ rand_device="cpu",
207
+ height=480,
208
+ width=832,
209
+ num_frames=81,
210
+ cfg_scale=5.0,
211
+ num_inference_steps=50,
212
+ sigma_shift=5.0,
213
+ tiled=True,
214
+ tile_size=(30, 52),
215
+ tile_stride=(15, 26),
216
+ tea_cache_l1_thresh=None,
217
+ tea_cache_model_id="",
218
+ progress_bar_cmd=tqdm,
219
+ progress_bar_st=None,
220
+ wo_ref_weight=0.0,
221
+ use_muti_ref_image=False,
222
+ ):
223
+
224
+ basecolor = batch["basecolor"]
225
+ depth = batch["depth"]
226
+ metallic = batch["metallic"]
227
+ normal = batch["normal"]
228
+ roughness = batch["roughness"]
229
+ ldr = batch["ldr"]
230
+ hdr_log = batch["hdr_log"]
231
+ env_dir = batch["env_dir"]
232
+ metallic_weight = batch["metallic_weight"]
233
+ roughness_weight = batch["roughness_weight"]
234
+ env_self_weight = batch["env_self_weight"]
235
+ env_cross_weight = batch["env_cross_weight"]
236
+
237
+ b = basecolor.shape[0]
238
+
239
+ # Parameter check
240
+ height, width = self.check_resize_height_width(height, width)
241
+ if num_frames % 4 != 1:
242
+ num_frames = (num_frames + 2) // 4 * 4 + 1
243
+ print(f"Only `num_frames % 4 != 1` is acceptable. We round it up to {num_frames}.")
244
+
245
+ # Tiler parameters
246
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
247
+
248
+ # Scheduler
249
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, shift=sigma_shift)
250
+
251
+ # Initialize noise
252
+ noise = self.generate_noise((b, 16, ((num_frames - 1) // 4 + 1)*2, height//8, width//8), seed=seed, device=rand_device, dtype=torch.float32) # noise
253
+ noise = noise.to(dtype=self.torch_dtype, device=self.device)
254
+ if input_video is not None:
255
+ self.load_models_to_device(['vae'])
256
+ input_video = self.preprocess_images(input_video)
257
+ input_video = torch.stack(input_video, dim=2).to(dtype=self.torch_dtype, device=self.device)
258
+ latents = self.encode_video(input_video, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
259
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
260
+ else:
261
+ latents = noise
262
+
263
+ # Encode Gbuffers (diffusionrenderer)
264
+ self.load_models_to_device(['vae'])
265
+ basecolor = basecolor.to(dtype=self.torch_dtype, device=self.device)
266
+ depth = depth.to(dtype=self.torch_dtype, device=self.device)
267
+ metallic = metallic.to(dtype=self.torch_dtype, device=self.device)
268
+ normal = normal.to(dtype=self.torch_dtype, device=self.device)
269
+ roughness = roughness.to(dtype=self.torch_dtype, device=self.device)
270
+
271
+ basecolor_latents = self.encode_video(basecolor, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
272
+ depth_latents = self.encode_video(depth, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
273
+ metallic_latents = self.encode_video(metallic, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
274
+ normal_latents = self.encode_video(normal, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
275
+ roughness_latents = self.encode_video(roughness, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
276
+
277
+ # If there are significant deviations in metallic and roughness, we can set their weights to 0 to avoid potential negative effects.
278
+ metallic_weight = metallic_weight.to(dtype=self.torch_dtype, device=self.device)
279
+ roughness_weight = roughness_weight.to(dtype=self.torch_dtype, device=self.device)
280
+ metallic_latents = metallic_latents * metallic_weight.reshape(-1, 1, 1, 1, 1)
281
+ roughness_latents = roughness_latents * roughness_weight.reshape(-1, 1, 1, 1, 1)
282
+
283
+ # Process env map (diffusionrenderer)
284
+ ldr = ldr.to(dtype=self.torch_dtype, device=self.device)
285
+ hdr_log = hdr_log.to(dtype=self.torch_dtype, device=self.device)
286
+ env_dir = env_dir.to(dtype=self.torch_dtype, device=self.device)
287
+
288
+ h, w = ldr.shape[-2:]
289
+ h_std, w_std = 256, 512
290
+ if (h, w) != (h_std, w_std):
291
+ use_same_env = False
292
+ ldr_cross = ldr.permute(0, 2, 1, 3, 4).contiguous().view(-1, ldr.size(1), ldr.size(3), ldr.size(4)) # b c t h w -> b*t c h w
293
+ ldr_cross = F.interpolate(ldr_cross, size=(h_std, w_std), mode='nearest')
294
+ ldr_cross = ldr_cross.view(ldr.size(0), ldr.size(2), ldr.size(1), h_std, w_std).permute(0, 2, 1, 3, 4) # b*t c h_std w_std -> b c t h_std w_std
295
+ hdr_log_cross = hdr_log.permute(0, 2, 1, 3, 4).contiguous().view(-1, hdr_log.size(1), hdr_log.size(3), hdr_log.size(4))
296
+ hdr_log_cross = F.interpolate(hdr_log_cross, size=(h_std, w_std), mode='nearest')
297
+ hdr_log_cross = hdr_log_cross.view(hdr_log.size(0), hdr_log.size(2), hdr_log.size(1), h_std, w_std).permute(0, 2, 1, 3, 4)
298
+ env_dir_cross = env_dir.permute(0, 2, 1, 3, 4).contiguous().view(-1, env_dir.size(1), env_dir.size(3), env_dir.size(4))
299
+ env_dir_cross = F.interpolate(env_dir_cross, size=(h_std, w_std), mode='nearest')
300
+ env_dir_cross = env_dir_cross.view(env_dir.size(0), env_dir.size(2), env_dir.size(1), h_std, w_std).permute(0, 2, 1, 3, 4)
301
+ else:
302
+ use_same_env = True
303
+
304
+ ldr_latents = self.encode_video(ldr, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device) # (b c f h w) -> (b c_encode f/4 h/8 w/8)
305
+ hdr_log_latents = self.encode_video(hdr_log, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
306
+ env_dir_latents = self.encode_video(env_dir, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
307
+ env_emb = torch.cat([ldr_latents, hdr_log_latents, env_dir_latents], dim=1)
308
+ env_self_weight = env_self_weight.to(dtype=self.torch_dtype, device=self.device)
309
+ env_cross_weight = env_cross_weight.to(dtype=self.torch_dtype, device=self.device)
310
+ env_emb = env_emb * env_self_weight.reshape(-1, 1, 1, 1, 1)
311
+
312
+ if not use_same_env:
313
+ ldr_cross = self.encode_video(ldr_cross, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
314
+ hdr_log_cross = self.encode_video(hdr_log_cross, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
315
+ env_dir_cross = self.encode_video(env_dir_cross, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
316
+ env_emb_cross = torch.cat([ldr_cross, hdr_log_cross, env_dir_cross], dim=1)
317
+ env_emb_cross = env_emb_cross * env_cross_weight.reshape(-1, 1, 1, 1, 1)
318
+
319
+ if "ref_image" in batch.keys():
320
+ ref_image = batch["ref_image"].to(dtype=self.torch_dtype, device=self.device)
321
+ ref_latent = self.encode_video(ref_image, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
322
+ else:
323
+ ref_image = batch["basecolor"][:,:,0:1,:,:].to(dtype=self.torch_dtype, device=self.device)
324
+ ref_latent = self.encode_video(ref_image, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device) * 0.0
325
+
326
+ # Encode prompts
327
+ self.load_models_to_device(["text_encoder"])
328
+ prompt_emb_posi = self.encode_prompt(prompt, positive=True)
329
+ if cfg_scale != 1.0:
330
+ prompt_emb_nega = self.encode_prompt(negative_prompt, positive=False)
331
+
332
+ # Encode image
333
+ if input_image is not None and self.image_encoder is not None:
334
+ self.load_models_to_device(["image_encoder", "vae"])
335
+ image_emb = self.encode_image(input_image, num_frames, height, width)
336
+ else:
337
+ image_emb = {}
338
+
339
+ # Extra input
340
+ extra_input = self.prepare_extra_input(latents)
341
+
342
+ # TeaCache
343
+ tea_cache_posi = {"tea_cache": TeaCache(num_inference_steps, rel_l1_thresh=tea_cache_l1_thresh, model_id=tea_cache_model_id) if tea_cache_l1_thresh is not None else None}
344
+ tea_cache_nega = {"tea_cache": TeaCache(num_inference_steps, rel_l1_thresh=tea_cache_l1_thresh, model_id=tea_cache_model_id) if tea_cache_l1_thresh is not None else None}
345
+
346
+ # Denoise
347
+ self.load_models_to_device(["dit"])
348
+ tgt_latent_length = latents.shape[2]
349
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
350
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
351
+
352
+ latents_input = torch.cat([latents, basecolor_latents, roughness_latents], dim=2)
353
+ geo_latents = torch.cat([depth_latents, normal_latents], dim=2)
354
+ use_context = True
355
+ if use_same_env:
356
+ other_latents_list = [('latent', metallic_latents)]
357
+ else:
358
+ other_latents_list = [('latent', metallic_latents), ('env', env_emb_cross)]
359
+
360
+ if use_muti_ref_image:
361
+ ref_frames = batch["ref_video"].to(dtype=self.torch_dtype, device=self.device) # b c t h w
362
+ ref_index = min(int(progress_id / num_inference_steps * num_frames), num_frames-1)
363
+ ref_image_now = ref_frames[:,:,ref_index:ref_index+1,...]
364
+ ref_latent = self.encode_video(ref_image_now, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
365
+
366
+ # Inference
367
+ noise_pred_posi = model_fn_wan_video(self.dit, latents_input, timestep=timestep, geo_latents=geo_latents, env_emb=env_emb, ref_latent=ref_latent, use_context=use_context, other_latents_list=other_latents_list, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi)
368
+ if cfg_scale != 1.0:
369
+ noise_pred_nega = model_fn_wan_video(self.dit, latents_input, timestep=timestep, geo_latents=geo_latents, env_emb=env_emb, ref_latent=ref_latent, use_context=use_context, other_latents_list=other_latents_list, **prompt_emb_nega, **image_emb, **extra_input, **tea_cache_nega)
370
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
371
+ else:
372
+ noise_pred = noise_pred_posi
373
+
374
+ if wo_ref_weight != 0.0 and "ref_image" in batch.keys():
375
+ noise_pred_wo_ref = model_fn_wan_video(self.dit, latents_input, timestep=timestep, geo_latents=geo_latents, env_emb=env_emb, ref_latent=ref_latent*0.0, use_context=use_context, other_latents_list=other_latents_list, **prompt_emb_posi, **image_emb, **extra_input, **tea_cache_posi)
376
+ if wo_ref_weight >= 0:
377
+ noise_pred = noise_pred / (1+wo_ref_weight) + noise_pred_wo_ref * wo_ref_weight / (1+wo_ref_weight)
378
+ else:
379
+ noise_pred = noise_pred * (-wo_ref_weight) / (1-wo_ref_weight) + noise_pred_wo_ref / (1-wo_ref_weight)
380
+
381
+ # Scheduler
382
+ latents = self.scheduler.step(noise_pred[:,:,:tgt_latent_length,...], self.scheduler.timesteps[progress_id], latents_input[:,:,:tgt_latent_length,...])
383
+
384
+ # Decode
385
+ self.load_models_to_device(['vae'])
386
+ frames = self.decode_video(latents[:,:,:(tgt_latent_length//2),...], **tiler_kwargs)
387
+ envs = self.decode_video(latents[:,:,(tgt_latent_length//2):(tgt_latent_length),...], **tiler_kwargs)
388
+
389
+ self.load_models_to_device([])
390
+ frames, frames_float_array = self.tensor2video(frames[0], return_float_array=True)
391
+ envs, envs_float_array = self.tensor2video(envs[0], return_float_array=True)
392
+
393
+ return frames, envs
394
+
395
+ class TeaCache:
396
+ def __init__(self, num_inference_steps, rel_l1_thresh, model_id):
397
+ self.num_inference_steps = num_inference_steps
398
+ self.step = 0
399
+ self.accumulated_rel_l1_distance = 0
400
+ self.previous_modulated_input = None
401
+ self.rel_l1_thresh = rel_l1_thresh
402
+ self.previous_residual = None
403
+ self.previous_hidden_states = None
404
+
405
+ self.coefficients_dict = {
406
+ "Wan2.1-T2V-1.3B": [-5.21862437e+04, 9.23041404e+03, -5.28275948e+02, 1.36987616e+01, -4.99875664e-02],
407
+ "Wan2.1-T2V-14B": [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01],
408
+ "Wan2.1-I2V-14B-480P": [2.57151496e+05, -3.54229917e+04, 1.40286849e+03, -1.35890334e+01, 1.32517977e-01],
409
+ "Wan2.1-I2V-14B-720P": [ 8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02],
410
+ }
411
+ if model_id not in self.coefficients_dict:
412
+ supported_model_ids = ", ".join([i for i in self.coefficients_dict])
413
+ raise ValueError(f"{model_id} is not a supported TeaCache model id. Please choose a valid model id in ({supported_model_ids}).")
414
+ self.coefficients = self.coefficients_dict[model_id]
415
+
416
+ def check(self, dit: WanModel, x, t_mod):
417
+ modulated_inp = t_mod.clone()
418
+ if self.step == 0 or self.step == self.num_inference_steps - 1:
419
+ should_calc = True
420
+ self.accumulated_rel_l1_distance = 0
421
+ else:
422
+ coefficients = self.coefficients
423
+ rescale_func = np.poly1d(coefficients)
424
+ self.accumulated_rel_l1_distance += rescale_func(((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()).cpu().item())
425
+ if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
426
+ should_calc = False
427
+ else:
428
+ should_calc = True
429
+ self.accumulated_rel_l1_distance = 0
430
+ self.previous_modulated_input = modulated_inp
431
+ self.step += 1
432
+ if self.step == self.num_inference_steps:
433
+ self.step = 0
434
+ if should_calc:
435
+ self.previous_hidden_states = x.clone()
436
+ return not should_calc
437
+
438
+ def store(self, hidden_states):
439
+ self.previous_residual = hidden_states - self.previous_hidden_states
440
+ self.previous_hidden_states = None
441
+
442
+ def update(self, hidden_states):
443
+ hidden_states = hidden_states + self.previous_residual
444
+ return hidden_states
445
+
446
+
447
+
448
+ def model_fn_wan_video(
449
+ dit: WanModel,
450
+ x: torch.Tensor,
451
+ timestep: torch.Tensor,
452
+ geo_latents: torch.Tensor,
453
+ env_emb: torch.Tensor,
454
+ context: torch.Tensor,
455
+ ref_latent: torch.Tensor,
456
+ other_latents_list: Optional[list] = None,
457
+ use_context: bool = False,
458
+ clip_feature: Optional[torch.Tensor] = None,
459
+ y: Optional[torch.Tensor] = None,
460
+ tea_cache: TeaCache = None,
461
+ **kwargs,
462
+ ):
463
+
464
+ t = dit.time_embedding(sinusoidal_embedding_1d(dit.freq_dim, timestep)) # time_embedding
465
+ t_mod = dit.time_projection(t).unflatten(1, (6, dit.dim))
466
+
467
+ if use_context:
468
+ context = dit.text_embedding(context) # prompt_embedding
469
+ else:
470
+ context = None
471
+
472
+ x, (f, h, w) = dit.patchify(x)
473
+
474
+ # Reference image
475
+ if len(ref_latent.shape) == 5:
476
+ ref_latent = ref_latent[:, :, 0] # (b c_encode 1 h/8 w/8) -> (b c_encode h/8 w/8)
477
+ ref_latent = dit.ref_conv(ref_latent).flatten(2).transpose(1, 2) # (b c_encode h/8 w/8) -> (b dim h/16? w/16?) -> (b h/16?*w/16? dim)
478
+ x = torch.concat([ref_latent, x], dim=1)
479
+ f += 1
480
+
481
+ if geo_latents is not None:
482
+ geo_latents = dit.patch_embedding(geo_latents)
483
+ if env_emb.shape[1] == 16:
484
+ env_emb = dit.patch_embedding(env_emb)
485
+ if other_latents_list is not None:
486
+ other_latents_list = [dit.patch_embedding(latent) if flag == 'latent' else latent for (flag, latent) in other_latents_list]
487
+
488
+ f_sample_for_each_latent = (f-1) // 4
489
+ freqs_for_ref = dit.freqs[0][:1]
490
+ freqs_for_latent1 = torch.stack([dit.freqs[0][1+i//5*21+i%5] for i in range(f_sample_for_each_latent)])
491
+ freqs_for_latent2 = torch.stack([dit.freqs[0][6+i//5*21+i%5] for i in range(f_sample_for_each_latent)])
492
+ freqs_for_latent3 = torch.stack([dit.freqs[0][11+i//5*21+i%5] for i in range(f_sample_for_each_latent)])
493
+ freqs_for_latent4 = torch.stack([dit.freqs[0][16+i//5*21+i%5] for i in range(f_sample_for_each_latent)])
494
+ freqs_for_f = torch.concat([freqs_for_ref, freqs_for_latent1, freqs_for_latent2, freqs_for_latent3, freqs_for_latent4], dim=0)
495
+ freqs = torch.cat([
496
+ freqs_for_f.view(f, 1, 1, -1).expand(f, h, w, -1),
497
+ dit.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
498
+ dit.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
499
+ ], dim=-1).reshape(f * h * w, 1, -1).to(x.device) # 3d ROPE
500
+
501
+ # TeaCache
502
+ tea_cache=None
503
+ if tea_cache is not None:
504
+ tea_cache_update = tea_cache.check(dit, x, t_mod)
505
+ else:
506
+ tea_cache_update = False
507
+
508
+ if tea_cache_update:
509
+ x = tea_cache.update(x)
510
+ else:
511
+ # blocks
512
+ for block in dit.blocks:
513
+ x = block(x, geo_latents, env_emb, ref_latent, context, other_latents_list, t_mod, freqs)
514
+ if tea_cache is not None:
515
+ tea_cache.store(x)
516
+
517
+ x = dit.head(x, t)
518
+
519
+ # Remove reference latents
520
+ if ref_latent is not None:
521
+ x = x[:, ref_latent.shape[1]:]
522
+ f -= 1
523
+
524
+ x = dit.unpatchify(x, (f, h, w))
525
+ return x
diffsynth/processors/FastBlend.py ADDED
@@ -0,0 +1,142 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from PIL import Image
2
+ import cupy as cp
3
+ import numpy as np
4
+ from tqdm import tqdm
5
+ from ..extensions.FastBlend.patch_match import PyramidPatchMatcher
6
+ from ..extensions.FastBlend.runners.fast import TableManager
7
+ from .base import VideoProcessor
8
+
9
+
10
+ class FastBlendSmoother(VideoProcessor):
11
+ def __init__(
12
+ self,
13
+ inference_mode="fast", batch_size=8, window_size=60,
14
+ minimum_patch_size=5, threads_per_block=8, num_iter=5, gpu_id=0, guide_weight=10.0, initialize="identity", tracking_window_size=0
15
+ ):
16
+ self.inference_mode = inference_mode
17
+ self.batch_size = batch_size
18
+ self.window_size = window_size
19
+ self.ebsynth_config = {
20
+ "minimum_patch_size": minimum_patch_size,
21
+ "threads_per_block": threads_per_block,
22
+ "num_iter": num_iter,
23
+ "gpu_id": gpu_id,
24
+ "guide_weight": guide_weight,
25
+ "initialize": initialize,
26
+ "tracking_window_size": tracking_window_size
27
+ }
28
+
29
+ @staticmethod
30
+ def from_model_manager(model_manager, **kwargs):
31
+ # TODO: fetch GPU ID from model_manager
32
+ return FastBlendSmoother(**kwargs)
33
+
34
+ def inference_fast(self, frames_guide, frames_style):
35
+ table_manager = TableManager()
36
+ patch_match_engine = PyramidPatchMatcher(
37
+ image_height=frames_style[0].shape[0],
38
+ image_width=frames_style[0].shape[1],
39
+ channel=3,
40
+ **self.ebsynth_config
41
+ )
42
+ # left part
43
+ table_l = table_manager.build_remapping_table(frames_guide, frames_style, patch_match_engine, self.batch_size, desc="Fast Mode Step 1/4")
44
+ table_l = table_manager.remapping_table_to_blending_table(table_l)
45
+ table_l = table_manager.process_window_sum(frames_guide, table_l, patch_match_engine, self.window_size, self.batch_size, desc="Fast Mode Step 2/4")
46
+ # right part
47
+ table_r = table_manager.build_remapping_table(frames_guide[::-1], frames_style[::-1], patch_match_engine, self.batch_size, desc="Fast Mode Step 3/4")
48
+ table_r = table_manager.remapping_table_to_blending_table(table_r)
49
+ table_r = table_manager.process_window_sum(frames_guide[::-1], table_r, patch_match_engine, self.window_size, self.batch_size, desc="Fast Mode Step 4/4")[::-1]
50
+ # merge
51
+ frames = []
52
+ for (frame_l, weight_l), frame_m, (frame_r, weight_r) in zip(table_l, frames_style, table_r):
53
+ weight_m = -1
54
+ weight = weight_l + weight_m + weight_r
55
+ frame = frame_l * (weight_l / weight) + frame_m * (weight_m / weight) + frame_r * (weight_r / weight)
56
+ frames.append(frame)
57
+ frames = [frame.clip(0, 255).astype("uint8") for frame in frames]
58
+ frames = [Image.fromarray(frame) for frame in frames]
59
+ return frames
60
+
61
+ def inference_balanced(self, frames_guide, frames_style):
62
+ patch_match_engine = PyramidPatchMatcher(
63
+ image_height=frames_style[0].shape[0],
64
+ image_width=frames_style[0].shape[1],
65
+ channel=3,
66
+ **self.ebsynth_config
67
+ )
68
+ output_frames = []
69
+ # tasks
70
+ n = len(frames_style)
71
+ tasks = []
72
+ for target in range(n):
73
+ for source in range(target - self.window_size, target + self.window_size + 1):
74
+ if source >= 0 and source < n and source != target:
75
+ tasks.append((source, target))
76
+ # run
77
+ frames = [(None, 1) for i in range(n)]
78
+ for batch_id in tqdm(range(0, len(tasks), self.batch_size), desc="Balanced Mode"):
79
+ tasks_batch = tasks[batch_id: min(batch_id+self.batch_size, len(tasks))]
80
+ source_guide = np.stack([frames_guide[source] for source, target in tasks_batch])
81
+ target_guide = np.stack([frames_guide[target] for source, target in tasks_batch])
82
+ source_style = np.stack([frames_style[source] for source, target in tasks_batch])
83
+ _, target_style = patch_match_engine.estimate_nnf(source_guide, target_guide, source_style)
84
+ for (source, target), result in zip(tasks_batch, target_style):
85
+ frame, weight = frames[target]
86
+ if frame is None:
87
+ frame = frames_style[target]
88
+ frames[target] = (
89
+ frame * (weight / (weight + 1)) + result / (weight + 1),
90
+ weight + 1
91
+ )
92
+ if weight + 1 == min(n, target + self.window_size + 1) - max(0, target - self.window_size):
93
+ frame = frame.clip(0, 255).astype("uint8")
94
+ output_frames.append(Image.fromarray(frame))
95
+ frames[target] = (None, 1)
96
+ return output_frames
97
+
98
+ def inference_accurate(self, frames_guide, frames_style):
99
+ patch_match_engine = PyramidPatchMatcher(
100
+ image_height=frames_style[0].shape[0],
101
+ image_width=frames_style[0].shape[1],
102
+ channel=3,
103
+ use_mean_target_style=True,
104
+ **self.ebsynth_config
105
+ )
106
+ output_frames = []
107
+ # run
108
+ n = len(frames_style)
109
+ for target in tqdm(range(n), desc="Accurate Mode"):
110
+ l, r = max(target - self.window_size, 0), min(target + self.window_size + 1, n)
111
+ remapped_frames = []
112
+ for i in range(l, r, self.batch_size):
113
+ j = min(i + self.batch_size, r)
114
+ source_guide = np.stack([frames_guide[source] for source in range(i, j)])
115
+ target_guide = np.stack([frames_guide[target]] * (j - i))
116
+ source_style = np.stack([frames_style[source] for source in range(i, j)])
117
+ _, target_style = patch_match_engine.estimate_nnf(source_guide, target_guide, source_style)
118
+ remapped_frames.append(target_style)
119
+ frame = np.concatenate(remapped_frames, axis=0).mean(axis=0)
120
+ frame = frame.clip(0, 255).astype("uint8")
121
+ output_frames.append(Image.fromarray(frame))
122
+ return output_frames
123
+
124
+ def release_vram(self):
125
+ mempool = cp.get_default_memory_pool()
126
+ pinned_mempool = cp.get_default_pinned_memory_pool()
127
+ mempool.free_all_blocks()
128
+ pinned_mempool.free_all_blocks()
129
+
130
+ def __call__(self, rendered_frames, original_frames=None, **kwargs):
131
+ rendered_frames = [np.array(frame) for frame in rendered_frames]
132
+ original_frames = [np.array(frame) for frame in original_frames]
133
+ if self.inference_mode == "fast":
134
+ output_frames = self.inference_fast(original_frames, rendered_frames)
135
+ elif self.inference_mode == "balanced":
136
+ output_frames = self.inference_balanced(original_frames, rendered_frames)
137
+ elif self.inference_mode == "accurate":
138
+ output_frames = self.inference_accurate(original_frames, rendered_frames)
139
+ else:
140
+ raise ValueError("inference_mode must be fast, balanced or accurate")
141
+ self.release_vram()
142
+ return output_frames
diffsynth/processors/PILEditor.py ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from PIL import ImageEnhance
2
+ from .base import VideoProcessor
3
+
4
+
5
+ class ContrastEditor(VideoProcessor):
6
+ def __init__(self, rate=1.5):
7
+ self.rate = rate
8
+
9
+ @staticmethod
10
+ def from_model_manager(model_manager, **kwargs):
11
+ return ContrastEditor(**kwargs)
12
+
13
+ def __call__(self, rendered_frames, **kwargs):
14
+ rendered_frames = [ImageEnhance.Contrast(i).enhance(self.rate) for i in rendered_frames]
15
+ return rendered_frames
16
+
17
+
18
+ class SharpnessEditor(VideoProcessor):
19
+ def __init__(self, rate=1.5):
20
+ self.rate = rate
21
+
22
+ @staticmethod
23
+ def from_model_manager(model_manager, **kwargs):
24
+ return SharpnessEditor(**kwargs)
25
+
26
+ def __call__(self, rendered_frames, **kwargs):
27
+ rendered_frames = [ImageEnhance.Sharpness(i).enhance(self.rate) for i in rendered_frames]
28
+ return rendered_frames
diffsynth/processors/RIFE.py ADDED
@@ -0,0 +1,77 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import numpy as np
3
+ from PIL import Image
4
+ from .base import VideoProcessor
5
+
6
+
7
+ class RIFESmoother(VideoProcessor):
8
+ def __init__(self, model, device="cuda", scale=1.0, batch_size=4, interpolate=True):
9
+ self.model = model
10
+ self.device = device
11
+
12
+ # IFNet only does not support float16
13
+ self.torch_dtype = torch.float32
14
+
15
+ # Other parameters
16
+ self.scale = scale
17
+ self.batch_size = batch_size
18
+ self.interpolate = interpolate
19
+
20
+ @staticmethod
21
+ def from_model_manager(model_manager, **kwargs):
22
+ return RIFESmoother(model_manager.RIFE, device=model_manager.device, **kwargs)
23
+
24
+ def process_image(self, image):
25
+ width, height = image.size
26
+ if width % 32 != 0 or height % 32 != 0:
27
+ width = (width + 31) // 32
28
+ height = (height + 31) // 32
29
+ image = image.resize((width, height))
30
+ image = torch.Tensor(np.array(image, dtype=np.float32)[:, :, [2,1,0]] / 255).permute(2, 0, 1)
31
+ return image
32
+
33
+ def process_images(self, images):
34
+ images = [self.process_image(image) for image in images]
35
+ images = torch.stack(images)
36
+ return images
37
+
38
+ def decode_images(self, images):
39
+ images = (images[:, [2,1,0]].permute(0, 2, 3, 1) * 255).clip(0, 255).numpy().astype(np.uint8)
40
+ images = [Image.fromarray(image) for image in images]
41
+ return images
42
+
43
+ def process_tensors(self, input_tensor, scale=1.0, batch_size=4):
44
+ output_tensor = []
45
+ for batch_id in range(0, input_tensor.shape[0], batch_size):
46
+ batch_id_ = min(batch_id + batch_size, input_tensor.shape[0])
47
+ batch_input_tensor = input_tensor[batch_id: batch_id_]
48
+ batch_input_tensor = batch_input_tensor.to(device=self.device, dtype=self.torch_dtype)
49
+ flow, mask, merged = self.model(batch_input_tensor, [4/scale, 2/scale, 1/scale])
50
+ output_tensor.append(merged[2].cpu())
51
+ output_tensor = torch.concat(output_tensor, dim=0)
52
+ return output_tensor
53
+
54
+ @torch.no_grad()
55
+ def __call__(self, rendered_frames, **kwargs):
56
+ # Preprocess
57
+ processed_images = self.process_images(rendered_frames)
58
+
59
+ # Input
60
+ input_tensor = torch.cat((processed_images[:-2], processed_images[2:]), dim=1)
61
+
62
+ # Interpolate
63
+ output_tensor = self.process_tensors(input_tensor, scale=self.scale, batch_size=self.batch_size)
64
+
65
+ if self.interpolate:
66
+ # Blend
67
+ input_tensor = torch.cat((processed_images[1:-1], output_tensor), dim=1)
68
+ output_tensor = self.process_tensors(input_tensor, scale=self.scale, batch_size=self.batch_size)
69
+ processed_images[1:-1] = output_tensor
70
+ else:
71
+ processed_images[1:-1] = (processed_images[1:-1] + output_tensor) / 2
72
+
73
+ # To images
74
+ output_images = self.decode_images(processed_images)
75
+ if output_images[0].size != rendered_frames[0].size:
76
+ output_images = [image.resize(rendered_frames[0].size) for image in output_images]
77
+ return output_images
diffsynth/processors/__init__.py ADDED
File without changes
diffsynth/processors/base.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ class VideoProcessor:
2
+ def __init__(self):
3
+ pass
4
+
5
+ def __call__(self):
6
+ raise NotImplementedError
diffsynth/processors/sequencial_processor.py ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base import VideoProcessor
2
+
3
+
4
+ class AutoVideoProcessor(VideoProcessor):
5
+ def __init__(self):
6
+ pass
7
+
8
+ @staticmethod
9
+ def from_model_manager(model_manager, processor_type, **kwargs):
10
+ if processor_type == "FastBlend":
11
+ from .FastBlend import FastBlendSmoother
12
+ return FastBlendSmoother.from_model_manager(model_manager, **kwargs)
13
+ elif processor_type == "Contrast":
14
+ from .PILEditor import ContrastEditor
15
+ return ContrastEditor.from_model_manager(model_manager, **kwargs)
16
+ elif processor_type == "Sharpness":
17
+ from .PILEditor import SharpnessEditor
18
+ return SharpnessEditor.from_model_manager(model_manager, **kwargs)
19
+ elif processor_type == "RIFE":
20
+ from .RIFE import RIFESmoother
21
+ return RIFESmoother.from_model_manager(model_manager, **kwargs)
22
+ else:
23
+ raise ValueError(f"invalid processor_type: {processor_type}")
24
+
25
+
26
+ class SequencialProcessor(VideoProcessor):
27
+ def __init__(self, processors=[]):
28
+ self.processors = processors
29
+
30
+ @staticmethod
31
+ def from_model_manager(model_manager, configs):
32
+ processors = [
33
+ AutoVideoProcessor.from_model_manager(model_manager, config["processor_type"], **config["config"])
34
+ for config in configs
35
+ ]
36
+ return SequencialProcessor(processors)
37
+
38
+ def __call__(self, rendered_frames, **kwargs):
39
+ for processor in self.processors:
40
+ rendered_frames = processor(rendered_frames, **kwargs)
41
+ return rendered_frames
diffsynth/prompters/__init__.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .prompt_refiners import Translator, BeautifulPrompt, QwenPrompt
2
+ from .sd_prompter import SDPrompter
3
+ from .sdxl_prompter import SDXLPrompter
4
+ from .sd3_prompter import SD3Prompter
5
+ from .hunyuan_dit_prompter import HunyuanDiTPrompter
6
+ from .kolors_prompter import KolorsPrompter
7
+ from .flux_prompter import FluxPrompter
8
+ from .omost import OmostPromter
9
+ from .cog_prompter import CogPrompter
10
+ from .hunyuan_video_prompter import HunyuanVideoPrompter
11
+ from .stepvideo_prompter import StepVideoPrompter
12
+ from .wan_prompter import WanPrompter
diffsynth/prompters/base_prompter.py ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ..models.model_manager import ModelManager
2
+ import torch
3
+
4
+
5
+
6
+ def tokenize_long_prompt(tokenizer, prompt, max_length=None):
7
+ # Get model_max_length from self.tokenizer
8
+ length = tokenizer.model_max_length if max_length is None else max_length
9
+
10
+ # To avoid the warning. set self.tokenizer.model_max_length to +oo.
11
+ tokenizer.model_max_length = 99999999
12
+
13
+ # Tokenize it!
14
+ input_ids = tokenizer(prompt, return_tensors="pt").input_ids
15
+
16
+ # Determine the real length.
17
+ max_length = (input_ids.shape[1] + length - 1) // length * length
18
+
19
+ # Restore tokenizer.model_max_length
20
+ tokenizer.model_max_length = length
21
+
22
+ # Tokenize it again with fixed length.
23
+ input_ids = tokenizer(
24
+ prompt,
25
+ return_tensors="pt",
26
+ padding="max_length",
27
+ max_length=max_length,
28
+ truncation=True
29
+ ).input_ids
30
+
31
+ # Reshape input_ids to fit the text encoder.
32
+ num_sentence = input_ids.shape[1] // length
33
+ input_ids = input_ids.reshape((num_sentence, length))
34
+
35
+ return input_ids
36
+
37
+
38
+
39
+ class BasePrompter:
40
+ def __init__(self):
41
+ self.refiners = []
42
+ self.extenders = []
43
+
44
+
45
+ def load_prompt_refiners(self, model_manager: ModelManager, refiner_classes=[]):
46
+ for refiner_class in refiner_classes:
47
+ refiner = refiner_class.from_model_manager(model_manager)
48
+ self.refiners.append(refiner)
49
+
50
+ def load_prompt_extenders(self,model_manager:ModelManager,extender_classes=[]):
51
+ for extender_class in extender_classes:
52
+ extender = extender_class.from_model_manager(model_manager)
53
+ self.extenders.append(extender)
54
+
55
+
56
+ @torch.no_grad()
57
+ def process_prompt(self, prompt, positive=True):
58
+ if isinstance(prompt, list):
59
+ prompt = [self.process_prompt(prompt_, positive=positive) for prompt_ in prompt]
60
+ else:
61
+ for refiner in self.refiners:
62
+ prompt = refiner(prompt, positive=positive)
63
+ return prompt
64
+
65
+ @torch.no_grad()
66
+ def extend_prompt(self, prompt:str, positive=True):
67
+ extended_prompt = dict(prompt=prompt)
68
+ for extender in self.extenders:
69
+ extended_prompt = extender(extended_prompt)
70
+ return extended_prompt
diffsynth/prompters/cog_prompter.py ADDED
@@ -0,0 +1,46 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.flux_text_encoder import FluxTextEncoder2
3
+ from transformers import T5TokenizerFast
4
+ import os
5
+
6
+
7
+ class CogPrompter(BasePrompter):
8
+ def __init__(
9
+ self,
10
+ tokenizer_path=None
11
+ ):
12
+ if tokenizer_path is None:
13
+ base_path = os.path.dirname(os.path.dirname(__file__))
14
+ tokenizer_path = os.path.join(base_path, "tokenizer_configs/cog/tokenizer")
15
+ super().__init__()
16
+ self.tokenizer = T5TokenizerFast.from_pretrained(tokenizer_path)
17
+ self.text_encoder: FluxTextEncoder2 = None
18
+
19
+
20
+ def fetch_models(self, text_encoder: FluxTextEncoder2 = None):
21
+ self.text_encoder = text_encoder
22
+
23
+
24
+ def encode_prompt_using_t5(self, prompt, text_encoder, tokenizer, max_length, device):
25
+ input_ids = tokenizer(
26
+ prompt,
27
+ return_tensors="pt",
28
+ padding="max_length",
29
+ max_length=max_length,
30
+ truncation=True,
31
+ ).input_ids.to(device)
32
+ prompt_emb = text_encoder(input_ids)
33
+ prompt_emb = prompt_emb.reshape((1, prompt_emb.shape[0]*prompt_emb.shape[1], -1))
34
+
35
+ return prompt_emb
36
+
37
+
38
+ def encode_prompt(
39
+ self,
40
+ prompt,
41
+ positive=True,
42
+ device="cuda"
43
+ ):
44
+ prompt = self.process_prompt(prompt, positive=positive)
45
+ prompt_emb = self.encode_prompt_using_t5(prompt, self.text_encoder, self.tokenizer, 226, device)
46
+ return prompt_emb
diffsynth/prompters/flux_prompter.py ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.flux_text_encoder import FluxTextEncoder2
3
+ from ..models.sd3_text_encoder import SD3TextEncoder1
4
+ from transformers import CLIPTokenizer, T5TokenizerFast
5
+ import os, torch
6
+
7
+
8
+ class FluxPrompter(BasePrompter):
9
+ def __init__(
10
+ self,
11
+ tokenizer_1_path=None,
12
+ tokenizer_2_path=None
13
+ ):
14
+ if tokenizer_1_path is None:
15
+ base_path = os.path.dirname(os.path.dirname(__file__))
16
+ tokenizer_1_path = os.path.join(base_path, "tokenizer_configs/flux/tokenizer_1")
17
+ if tokenizer_2_path is None:
18
+ base_path = os.path.dirname(os.path.dirname(__file__))
19
+ tokenizer_2_path = os.path.join(base_path, "tokenizer_configs/flux/tokenizer_2")
20
+ super().__init__()
21
+ self.tokenizer_1 = CLIPTokenizer.from_pretrained(tokenizer_1_path)
22
+ self.tokenizer_2 = T5TokenizerFast.from_pretrained(tokenizer_2_path)
23
+ self.text_encoder_1: SD3TextEncoder1 = None
24
+ self.text_encoder_2: FluxTextEncoder2 = None
25
+
26
+
27
+ def fetch_models(self, text_encoder_1: SD3TextEncoder1 = None, text_encoder_2: FluxTextEncoder2 = None):
28
+ self.text_encoder_1 = text_encoder_1
29
+ self.text_encoder_2 = text_encoder_2
30
+
31
+
32
+ def encode_prompt_using_clip(self, prompt, text_encoder, tokenizer, max_length, device):
33
+ input_ids = tokenizer(
34
+ prompt,
35
+ return_tensors="pt",
36
+ padding="max_length",
37
+ max_length=max_length,
38
+ truncation=True
39
+ ).input_ids.to(device)
40
+ pooled_prompt_emb, _ = text_encoder(input_ids)
41
+ return pooled_prompt_emb
42
+
43
+
44
+ def encode_prompt_using_t5(self, prompt, text_encoder, tokenizer, max_length, device):
45
+ input_ids = tokenizer(
46
+ prompt,
47
+ return_tensors="pt",
48
+ padding="max_length",
49
+ max_length=max_length,
50
+ truncation=True,
51
+ ).input_ids.to(device)
52
+ prompt_emb = text_encoder(input_ids)
53
+ return prompt_emb
54
+
55
+
56
+ def encode_prompt(
57
+ self,
58
+ prompt,
59
+ positive=True,
60
+ device="cuda",
61
+ t5_sequence_length=512,
62
+ ):
63
+ prompt = self.process_prompt(prompt, positive=positive)
64
+
65
+ # CLIP
66
+ pooled_prompt_emb = self.encode_prompt_using_clip(prompt, self.text_encoder_1, self.tokenizer_1, 77, device)
67
+
68
+ # T5
69
+ prompt_emb = self.encode_prompt_using_t5(prompt, self.text_encoder_2, self.tokenizer_2, t5_sequence_length, device)
70
+
71
+ # text_ids
72
+ text_ids = torch.zeros(prompt_emb.shape[0], prompt_emb.shape[1], 3).to(device=device, dtype=prompt_emb.dtype)
73
+
74
+ return prompt_emb, pooled_prompt_emb, text_ids
diffsynth/prompters/hunyuan_dit_prompter.py ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.model_manager import ModelManager
3
+ from ..models import HunyuanDiTCLIPTextEncoder, HunyuanDiTT5TextEncoder
4
+ from transformers import BertTokenizer, AutoTokenizer
5
+ import warnings, os
6
+
7
+
8
+ class HunyuanDiTPrompter(BasePrompter):
9
+ def __init__(
10
+ self,
11
+ tokenizer_path=None,
12
+ tokenizer_t5_path=None
13
+ ):
14
+ if tokenizer_path is None:
15
+ base_path = os.path.dirname(os.path.dirname(__file__))
16
+ tokenizer_path = os.path.join(base_path, "tokenizer_configs/hunyuan_dit/tokenizer")
17
+ if tokenizer_t5_path is None:
18
+ base_path = os.path.dirname(os.path.dirname(__file__))
19
+ tokenizer_t5_path = os.path.join(base_path, "tokenizer_configs/hunyuan_dit/tokenizer_t5")
20
+ super().__init__()
21
+ self.tokenizer = BertTokenizer.from_pretrained(tokenizer_path)
22
+ with warnings.catch_warnings():
23
+ warnings.simplefilter("ignore")
24
+ self.tokenizer_t5 = AutoTokenizer.from_pretrained(tokenizer_t5_path)
25
+ self.text_encoder: HunyuanDiTCLIPTextEncoder = None
26
+ self.text_encoder_t5: HunyuanDiTT5TextEncoder = None
27
+
28
+
29
+ def fetch_models(self, text_encoder: HunyuanDiTCLIPTextEncoder = None, text_encoder_t5: HunyuanDiTT5TextEncoder = None):
30
+ self.text_encoder = text_encoder
31
+ self.text_encoder_t5 = text_encoder_t5
32
+
33
+
34
+ def encode_prompt_using_signle_model(self, prompt, text_encoder, tokenizer, max_length, clip_skip, device):
35
+ text_inputs = tokenizer(
36
+ prompt,
37
+ padding="max_length",
38
+ max_length=max_length,
39
+ truncation=True,
40
+ return_attention_mask=True,
41
+ return_tensors="pt",
42
+ )
43
+ text_input_ids = text_inputs.input_ids
44
+ attention_mask = text_inputs.attention_mask.to(device)
45
+ prompt_embeds = text_encoder(
46
+ text_input_ids.to(device),
47
+ attention_mask=attention_mask,
48
+ clip_skip=clip_skip
49
+ )
50
+ return prompt_embeds, attention_mask
51
+
52
+
53
+ def encode_prompt(
54
+ self,
55
+ prompt,
56
+ clip_skip=1,
57
+ clip_skip_2=1,
58
+ positive=True,
59
+ device="cuda"
60
+ ):
61
+ prompt = self.process_prompt(prompt, positive=positive)
62
+
63
+ # CLIP
64
+ prompt_emb, attention_mask = self.encode_prompt_using_signle_model(prompt, self.text_encoder, self.tokenizer, self.tokenizer.model_max_length, clip_skip, device)
65
+
66
+ # T5
67
+ prompt_emb_t5, attention_mask_t5 = self.encode_prompt_using_signle_model(prompt, self.text_encoder_t5, self.tokenizer_t5, self.tokenizer_t5.model_max_length, clip_skip_2, device)
68
+
69
+ return prompt_emb, attention_mask, prompt_emb_t5, attention_mask_t5
diffsynth/prompters/hunyuan_video_prompter.py ADDED
@@ -0,0 +1,275 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.sd3_text_encoder import SD3TextEncoder1
3
+ from ..models.hunyuan_video_text_encoder import HunyuanVideoLLMEncoder, HunyuanVideoMLLMEncoder
4
+ from transformers import CLIPTokenizer, LlamaTokenizerFast, CLIPImageProcessor
5
+ import os, torch
6
+ from typing import Union
7
+
8
+ PROMPT_TEMPLATE_ENCODE = (
9
+ "<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
10
+ "quantity, text, spatial relationships of the objects and background:<|eot_id|>"
11
+ "<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
12
+
13
+ PROMPT_TEMPLATE_ENCODE_VIDEO = (
14
+ "<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
15
+ "1. The main content and theme of the video."
16
+ "2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
17
+ "3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
18
+ "4. background environment, light, style and atmosphere."
19
+ "5. camera angles, movements, and transitions used in the video:<|eot_id|>"
20
+ "<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
21
+
22
+ PROMPT_TEMPLATE_ENCODE_I2V = (
23
+ "<|start_header_id|>system<|end_header_id|>\n\n<image>\nDescribe the image by detailing the color, shape, size, texture, "
24
+ "quantity, text, spatial relationships of the objects and background:<|eot_id|>"
25
+ "<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
26
+ "<|start_header_id|>assistant<|end_header_id|>\n\n"
27
+ )
28
+
29
+ PROMPT_TEMPLATE_ENCODE_VIDEO_I2V = (
30
+ "<|start_header_id|>system<|end_header_id|>\n\n<image>\nDescribe the video by detailing the following aspects according to the reference image: "
31
+ "1. The main content and theme of the video."
32
+ "2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
33
+ "3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
34
+ "4. background environment, light, style and atmosphere."
35
+ "5. camera angles, movements, and transitions used in the video:<|eot_id|>\n\n"
36
+ "<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
37
+ "<|start_header_id|>assistant<|end_header_id|>\n\n"
38
+ )
39
+
40
+ PROMPT_TEMPLATE = {
41
+ "dit-llm-encode": {
42
+ "template": PROMPT_TEMPLATE_ENCODE,
43
+ "crop_start": 36,
44
+ },
45
+ "dit-llm-encode-video": {
46
+ "template": PROMPT_TEMPLATE_ENCODE_VIDEO,
47
+ "crop_start": 95,
48
+ },
49
+ "dit-llm-encode-i2v": {
50
+ "template": PROMPT_TEMPLATE_ENCODE_I2V,
51
+ "crop_start": 36,
52
+ "image_emb_start": 5,
53
+ "image_emb_end": 581,
54
+ "image_emb_len": 576,
55
+ "double_return_token_id": 271
56
+ },
57
+ "dit-llm-encode-video-i2v": {
58
+ "template": PROMPT_TEMPLATE_ENCODE_VIDEO_I2V,
59
+ "crop_start": 103,
60
+ "image_emb_start": 5,
61
+ "image_emb_end": 581,
62
+ "image_emb_len": 576,
63
+ "double_return_token_id": 271
64
+ },
65
+ }
66
+
67
+ NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
68
+
69
+
70
+ class HunyuanVideoPrompter(BasePrompter):
71
+
72
+ def __init__(
73
+ self,
74
+ tokenizer_1_path=None,
75
+ tokenizer_2_path=None,
76
+ ):
77
+ if tokenizer_1_path is None:
78
+ base_path = os.path.dirname(os.path.dirname(__file__))
79
+ tokenizer_1_path = os.path.join(
80
+ base_path, "tokenizer_configs/hunyuan_video/tokenizer_1")
81
+ if tokenizer_2_path is None:
82
+ base_path = os.path.dirname(os.path.dirname(__file__))
83
+ tokenizer_2_path = os.path.join(
84
+ base_path, "tokenizer_configs/hunyuan_video/tokenizer_2")
85
+ super().__init__()
86
+ self.tokenizer_1 = CLIPTokenizer.from_pretrained(tokenizer_1_path)
87
+ self.tokenizer_2 = LlamaTokenizerFast.from_pretrained(tokenizer_2_path, padding_side='right')
88
+ self.text_encoder_1: SD3TextEncoder1 = None
89
+ self.text_encoder_2: HunyuanVideoLLMEncoder = None
90
+
91
+ self.prompt_template = PROMPT_TEMPLATE['dit-llm-encode']
92
+ self.prompt_template_video = PROMPT_TEMPLATE['dit-llm-encode-video']
93
+
94
+ def fetch_models(self,
95
+ text_encoder_1: SD3TextEncoder1 = None,
96
+ text_encoder_2: Union[HunyuanVideoLLMEncoder, HunyuanVideoMLLMEncoder] = None):
97
+ self.text_encoder_1 = text_encoder_1
98
+ self.text_encoder_2 = text_encoder_2
99
+ if isinstance(text_encoder_2, HunyuanVideoMLLMEncoder):
100
+ # processor
101
+ # TODO: may need to replace processor with local implementation
102
+ base_path = os.path.dirname(os.path.dirname(__file__))
103
+ tokenizer_2_path = os.path.join(base_path, "tokenizer_configs/hunyuan_video/tokenizer_2")
104
+ self.processor = CLIPImageProcessor.from_pretrained(tokenizer_2_path)
105
+ # template
106
+ self.prompt_template = PROMPT_TEMPLATE['dit-llm-encode-i2v']
107
+ self.prompt_template_video = PROMPT_TEMPLATE['dit-llm-encode-video-i2v']
108
+
109
+ def apply_text_to_template(self, text, template):
110
+ assert isinstance(template, str)
111
+ if isinstance(text, list):
112
+ return [self.apply_text_to_template(text_) for text_ in text]
113
+ elif isinstance(text, str):
114
+ # Will send string to tokenizer. Used for llm
115
+ return template.format(text)
116
+ else:
117
+ raise TypeError(f"Unsupported prompt type: {type(text)}")
118
+
119
+ def encode_prompt_using_clip(self, prompt, max_length, device):
120
+ tokenized_result = self.tokenizer_1(
121
+ prompt,
122
+ return_tensors="pt",
123
+ padding="max_length",
124
+ max_length=max_length,
125
+ truncation=True,
126
+ return_attention_mask=True
127
+ )
128
+ input_ids = tokenized_result.input_ids.to(device)
129
+ attention_mask = tokenized_result.attention_mask.to(device)
130
+ return self.text_encoder_1(input_ids=input_ids, extra_mask=attention_mask)[0]
131
+
132
+ def encode_prompt_using_llm(self,
133
+ prompt,
134
+ max_length,
135
+ device,
136
+ crop_start,
137
+ hidden_state_skip_layer=2,
138
+ use_attention_mask=True):
139
+ max_length += crop_start
140
+ inputs = self.tokenizer_2(prompt,
141
+ return_tensors="pt",
142
+ padding="max_length",
143
+ max_length=max_length,
144
+ truncation=True)
145
+ input_ids = inputs.input_ids.to(device)
146
+ attention_mask = inputs.attention_mask.to(device)
147
+ last_hidden_state = self.text_encoder_2(input_ids, attention_mask, hidden_state_skip_layer)
148
+
149
+ # crop out
150
+ if crop_start > 0:
151
+ last_hidden_state = last_hidden_state[:, crop_start:]
152
+ attention_mask = (attention_mask[:, crop_start:] if use_attention_mask else None)
153
+
154
+ return last_hidden_state, attention_mask
155
+
156
+ def encode_prompt_using_mllm(self,
157
+ prompt,
158
+ images,
159
+ max_length,
160
+ device,
161
+ crop_start,
162
+ hidden_state_skip_layer=2,
163
+ use_attention_mask=True,
164
+ image_embed_interleave=4):
165
+ image_outputs = self.processor(images, return_tensors="pt")["pixel_values"].to(device)
166
+ max_length += crop_start
167
+ inputs = self.tokenizer_2(prompt,
168
+ return_tensors="pt",
169
+ padding="max_length",
170
+ max_length=max_length,
171
+ truncation=True)
172
+ input_ids = inputs.input_ids.to(device)
173
+ attention_mask = inputs.attention_mask.to(device)
174
+ last_hidden_state = self.text_encoder_2(input_ids=input_ids,
175
+ attention_mask=attention_mask,
176
+ hidden_state_skip_layer=hidden_state_skip_layer,
177
+ pixel_values=image_outputs)
178
+
179
+ text_crop_start = (crop_start - 1 + self.prompt_template_video.get("image_emb_len", 576))
180
+ image_crop_start = self.prompt_template_video.get("image_emb_start", 5)
181
+ image_crop_end = self.prompt_template_video.get("image_emb_end", 581)
182
+ batch_indices, last_double_return_token_indices = torch.where(
183
+ input_ids == self.prompt_template_video.get("double_return_token_id", 271))
184
+ if last_double_return_token_indices.shape[0] == 3:
185
+ # in case the prompt is too long
186
+ last_double_return_token_indices = torch.cat((
187
+ last_double_return_token_indices,
188
+ torch.tensor([input_ids.shape[-1]]),
189
+ ))
190
+ batch_indices = torch.cat((batch_indices, torch.tensor([0])))
191
+ last_double_return_token_indices = (last_double_return_token_indices.reshape(input_ids.shape[0], -1)[:, -1])
192
+ batch_indices = batch_indices.reshape(input_ids.shape[0], -1)[:, -1]
193
+ assistant_crop_start = (last_double_return_token_indices - 1 + self.prompt_template_video.get("image_emb_len", 576) - 4)
194
+ assistant_crop_end = (last_double_return_token_indices - 1 + self.prompt_template_video.get("image_emb_len", 576))
195
+ attention_mask_assistant_crop_start = (last_double_return_token_indices - 4)
196
+ attention_mask_assistant_crop_end = last_double_return_token_indices
197
+ text_last_hidden_state = []
198
+ text_attention_mask = []
199
+ image_last_hidden_state = []
200
+ image_attention_mask = []
201
+ for i in range(input_ids.shape[0]):
202
+ text_last_hidden_state.append(
203
+ torch.cat([
204
+ last_hidden_state[i, text_crop_start:assistant_crop_start[i].item()],
205
+ last_hidden_state[i, assistant_crop_end[i].item():],
206
+ ]))
207
+ text_attention_mask.append(
208
+ torch.cat([
209
+ attention_mask[
210
+ i,
211
+ crop_start:attention_mask_assistant_crop_start[i].item(),
212
+ ],
213
+ attention_mask[i, attention_mask_assistant_crop_end[i].item():],
214
+ ]) if use_attention_mask else None)
215
+ image_last_hidden_state.append(last_hidden_state[i, image_crop_start:image_crop_end])
216
+ image_attention_mask.append(
217
+ torch.ones(image_last_hidden_state[-1].shape[0]).to(last_hidden_state.device).
218
+ to(attention_mask.dtype) if use_attention_mask else None)
219
+
220
+ text_last_hidden_state = torch.stack(text_last_hidden_state)
221
+ text_attention_mask = torch.stack(text_attention_mask)
222
+ image_last_hidden_state = torch.stack(image_last_hidden_state)
223
+ image_attention_mask = torch.stack(image_attention_mask)
224
+
225
+ image_last_hidden_state = image_last_hidden_state[:, ::image_embed_interleave, :]
226
+ image_attention_mask = image_attention_mask[:, ::image_embed_interleave]
227
+
228
+ assert (text_last_hidden_state.shape[0] == text_attention_mask.shape[0] and
229
+ image_last_hidden_state.shape[0] == image_attention_mask.shape[0])
230
+
231
+ last_hidden_state = torch.cat([image_last_hidden_state, text_last_hidden_state], dim=1)
232
+ attention_mask = torch.cat([image_attention_mask, text_attention_mask], dim=1)
233
+
234
+ return last_hidden_state, attention_mask
235
+
236
+ def encode_prompt(self,
237
+ prompt,
238
+ images=None,
239
+ positive=True,
240
+ device="cuda",
241
+ clip_sequence_length=77,
242
+ llm_sequence_length=256,
243
+ data_type='video',
244
+ use_template=True,
245
+ hidden_state_skip_layer=2,
246
+ use_attention_mask=True,
247
+ image_embed_interleave=4):
248
+
249
+ prompt = self.process_prompt(prompt, positive=positive)
250
+
251
+ # apply template
252
+ if use_template:
253
+ template = self.prompt_template_video if data_type == 'video' else self.prompt_template
254
+ prompt_formated = self.apply_text_to_template(prompt, template['template'])
255
+ else:
256
+ prompt_formated = prompt
257
+ # Text encoder
258
+ if data_type == 'video':
259
+ crop_start = self.prompt_template_video.get("crop_start", 0)
260
+ else:
261
+ crop_start = self.prompt_template.get("crop_start", 0)
262
+
263
+ # CLIP
264
+ pooled_prompt_emb = self.encode_prompt_using_clip(prompt, clip_sequence_length, device)
265
+
266
+ # LLM
267
+ if images is None:
268
+ prompt_emb, attention_mask = self.encode_prompt_using_llm(prompt_formated, llm_sequence_length, device, crop_start,
269
+ hidden_state_skip_layer, use_attention_mask)
270
+ else:
271
+ prompt_emb, attention_mask = self.encode_prompt_using_mllm(prompt_formated, images, llm_sequence_length, device,
272
+ crop_start, hidden_state_skip_layer, use_attention_mask,
273
+ image_embed_interleave)
274
+
275
+ return prompt_emb, pooled_prompt_emb, attention_mask
diffsynth/prompters/kolors_prompter.py ADDED
@@ -0,0 +1,354 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.model_manager import ModelManager
3
+ import json, os, re
4
+ from typing import List, Optional, Union, Dict
5
+ from sentencepiece import SentencePieceProcessor
6
+ from transformers import PreTrainedTokenizer
7
+ from transformers.utils import PaddingStrategy
8
+ from transformers.tokenization_utils_base import EncodedInput, BatchEncoding
9
+ from ..models.kolors_text_encoder import ChatGLMModel
10
+
11
+
12
+ class SPTokenizer:
13
+ def __init__(self, model_path: str):
14
+ # reload tokenizer
15
+ assert os.path.isfile(model_path), model_path
16
+ self.sp_model = SentencePieceProcessor(model_file=model_path)
17
+
18
+ # BOS / EOS token IDs
19
+ self.n_words: int = self.sp_model.vocab_size()
20
+ self.bos_id: int = self.sp_model.bos_id()
21
+ self.eos_id: int = self.sp_model.eos_id()
22
+ self.pad_id: int = self.sp_model.unk_id()
23
+ assert self.sp_model.vocab_size() == self.sp_model.get_piece_size()
24
+
25
+ role_special_tokens = ["<|system|>", "<|user|>", "<|assistant|>", "<|observation|>"]
26
+ special_tokens = ["[MASK]", "[gMASK]", "[sMASK]", "sop", "eop"] + role_special_tokens
27
+ self.special_tokens = {}
28
+ self.index_special_tokens = {}
29
+ for token in special_tokens:
30
+ self.special_tokens[token] = self.n_words
31
+ self.index_special_tokens[self.n_words] = token
32
+ self.n_words += 1
33
+ self.role_special_token_expression = "|".join([re.escape(token) for token in role_special_tokens])
34
+
35
+ def tokenize(self, s: str, encode_special_tokens=False):
36
+ if encode_special_tokens:
37
+ last_index = 0
38
+ t = []
39
+ for match in re.finditer(self.role_special_token_expression, s):
40
+ if last_index < match.start():
41
+ t.extend(self.sp_model.EncodeAsPieces(s[last_index:match.start()]))
42
+ t.append(s[match.start():match.end()])
43
+ last_index = match.end()
44
+ if last_index < len(s):
45
+ t.extend(self.sp_model.EncodeAsPieces(s[last_index:]))
46
+ return t
47
+ else:
48
+ return self.sp_model.EncodeAsPieces(s)
49
+
50
+ def encode(self, s: str, bos: bool = False, eos: bool = False) -> List[int]:
51
+ assert type(s) is str
52
+ t = self.sp_model.encode(s)
53
+ if bos:
54
+ t = [self.bos_id] + t
55
+ if eos:
56
+ t = t + [self.eos_id]
57
+ return t
58
+
59
+ def decode(self, t: List[int]) -> str:
60
+ text, buffer = "", []
61
+ for token in t:
62
+ if token in self.index_special_tokens:
63
+ if buffer:
64
+ text += self.sp_model.decode(buffer)
65
+ buffer = []
66
+ text += self.index_special_tokens[token]
67
+ else:
68
+ buffer.append(token)
69
+ if buffer:
70
+ text += self.sp_model.decode(buffer)
71
+ return text
72
+
73
+ def decode_tokens(self, tokens: List[str]) -> str:
74
+ text = self.sp_model.DecodePieces(tokens)
75
+ return text
76
+
77
+ def convert_token_to_id(self, token):
78
+ """ Converts a token (str) in an id using the vocab. """
79
+ if token in self.special_tokens:
80
+ return self.special_tokens[token]
81
+ return self.sp_model.PieceToId(token)
82
+
83
+ def convert_id_to_token(self, index):
84
+ """Converts an index (integer) in a token (str) using the vocab."""
85
+ if index in self.index_special_tokens:
86
+ return self.index_special_tokens[index]
87
+ if index in [self.eos_id, self.bos_id, self.pad_id] or index < 0:
88
+ return ""
89
+ return self.sp_model.IdToPiece(index)
90
+
91
+
92
+
93
+ class ChatGLMTokenizer(PreTrainedTokenizer):
94
+ vocab_files_names = {"vocab_file": "tokenizer.model"}
95
+
96
+ model_input_names = ["input_ids", "attention_mask", "position_ids"]
97
+
98
+ def __init__(self, vocab_file, padding_side="left", clean_up_tokenization_spaces=False, encode_special_tokens=False,
99
+ **kwargs):
100
+ self.name = "GLMTokenizer"
101
+
102
+ self.vocab_file = vocab_file
103
+ self.tokenizer = SPTokenizer(vocab_file)
104
+ self.special_tokens = {
105
+ "<bos>": self.tokenizer.bos_id,
106
+ "<eos>": self.tokenizer.eos_id,
107
+ "<pad>": self.tokenizer.pad_id
108
+ }
109
+ self.encode_special_tokens = encode_special_tokens
110
+ super().__init__(padding_side=padding_side, clean_up_tokenization_spaces=clean_up_tokenization_spaces,
111
+ encode_special_tokens=encode_special_tokens,
112
+ **kwargs)
113
+
114
+ def get_command(self, token):
115
+ if token in self.special_tokens:
116
+ return self.special_tokens[token]
117
+ assert token in self.tokenizer.special_tokens, f"{token} is not a special token for {self.name}"
118
+ return self.tokenizer.special_tokens[token]
119
+
120
+ @property
121
+ def unk_token(self) -> str:
122
+ return "<unk>"
123
+
124
+ @property
125
+ def pad_token(self) -> str:
126
+ return "<unk>"
127
+
128
+ @property
129
+ def pad_token_id(self):
130
+ return self.get_command("<pad>")
131
+
132
+ @property
133
+ def eos_token(self) -> str:
134
+ return "</s>"
135
+
136
+ @property
137
+ def eos_token_id(self):
138
+ return self.get_command("<eos>")
139
+
140
+ @property
141
+ def vocab_size(self):
142
+ return self.tokenizer.n_words
143
+
144
+ def get_vocab(self):
145
+ """ Returns vocab as a dict """
146
+ vocab = {self._convert_id_to_token(i): i for i in range(self.vocab_size)}
147
+ vocab.update(self.added_tokens_encoder)
148
+ return vocab
149
+
150
+ def _tokenize(self, text, **kwargs):
151
+ return self.tokenizer.tokenize(text, encode_special_tokens=self.encode_special_tokens)
152
+
153
+ def _convert_token_to_id(self, token):
154
+ """ Converts a token (str) in an id using the vocab. """
155
+ return self.tokenizer.convert_token_to_id(token)
156
+
157
+ def _convert_id_to_token(self, index):
158
+ """Converts an index (integer) in a token (str) using the vocab."""
159
+ return self.tokenizer.convert_id_to_token(index)
160
+
161
+ def convert_tokens_to_string(self, tokens: List[str]) -> str:
162
+ return self.tokenizer.decode_tokens(tokens)
163
+
164
+ def save_vocabulary(self, save_directory, filename_prefix=None):
165
+ """
166
+ Save the vocabulary and special tokens file to a directory.
167
+
168
+ Args:
169
+ save_directory (`str`):
170
+ The directory in which to save the vocabulary.
171
+ filename_prefix (`str`, *optional*):
172
+ An optional prefix to add to the named of the saved files.
173
+
174
+ Returns:
175
+ `Tuple(str)`: Paths to the files saved.
176
+ """
177
+ if os.path.isdir(save_directory):
178
+ vocab_file = os.path.join(
179
+ save_directory, self.vocab_files_names["vocab_file"]
180
+ )
181
+ else:
182
+ vocab_file = save_directory
183
+
184
+ with open(self.vocab_file, 'rb') as fin:
185
+ proto_str = fin.read()
186
+
187
+ with open(vocab_file, "wb") as writer:
188
+ writer.write(proto_str)
189
+
190
+ return (vocab_file,)
191
+
192
+ def get_prefix_tokens(self):
193
+ prefix_tokens = [self.get_command("[gMASK]"), self.get_command("sop")]
194
+ return prefix_tokens
195
+
196
+ def build_single_message(self, role, metadata, message):
197
+ assert role in ["system", "user", "assistant", "observation"], role
198
+ role_tokens = [self.get_command(f"<|{role}|>")] + self.tokenizer.encode(f"{metadata}\n")
199
+ message_tokens = self.tokenizer.encode(message)
200
+ tokens = role_tokens + message_tokens
201
+ return tokens
202
+
203
+ def build_chat_input(self, query, history=None, role="user"):
204
+ if history is None:
205
+ history = []
206
+ input_ids = []
207
+ for item in history:
208
+ content = item["content"]
209
+ if item["role"] == "system" and "tools" in item:
210
+ content = content + "\n" + json.dumps(item["tools"], indent=4, ensure_ascii=False)
211
+ input_ids.extend(self.build_single_message(item["role"], item.get("metadata", ""), content))
212
+ input_ids.extend(self.build_single_message(role, "", query))
213
+ input_ids.extend([self.get_command("<|assistant|>")])
214
+ return self.batch_encode_plus([input_ids], return_tensors="pt", is_split_into_words=True)
215
+
216
+ def build_inputs_with_special_tokens(
217
+ self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None
218
+ ) -> List[int]:
219
+ """
220
+ Build model inputs from a sequence or a pair of sequence for sequence classification tasks by concatenating and
221
+ adding special tokens. A BERT sequence has the following format:
222
+
223
+ - single sequence: `[CLS] X [SEP]`
224
+ - pair of sequences: `[CLS] A [SEP] B [SEP]`
225
+
226
+ Args:
227
+ token_ids_0 (`List[int]`):
228
+ List of IDs to which the special tokens will be added.
229
+ token_ids_1 (`List[int]`, *optional*):
230
+ Optional second list of IDs for sequence pairs.
231
+
232
+ Returns:
233
+ `List[int]`: List of [input IDs](../glossary#input-ids) with the appropriate special tokens.
234
+ """
235
+ prefix_tokens = self.get_prefix_tokens()
236
+ token_ids_0 = prefix_tokens + token_ids_0
237
+ if token_ids_1 is not None:
238
+ token_ids_0 = token_ids_0 + token_ids_1 + [self.get_command("<eos>")]
239
+ return token_ids_0
240
+
241
+ def _pad(
242
+ self,
243
+ encoded_inputs: Union[Dict[str, EncodedInput], BatchEncoding],
244
+ max_length: Optional[int] = None,
245
+ padding_strategy: PaddingStrategy = PaddingStrategy.DO_NOT_PAD,
246
+ pad_to_multiple_of: Optional[int] = None,
247
+ return_attention_mask: Optional[bool] = None,
248
+ padding_side: Optional[str] = None,
249
+ ) -> dict:
250
+ """
251
+ Pad encoded inputs (on left/right and up to predefined length or max length in the batch)
252
+
253
+ Args:
254
+ encoded_inputs:
255
+ Dictionary of tokenized inputs (`List[int]`) or batch of tokenized inputs (`List[List[int]]`).
256
+ max_length: maximum length of the returned list and optionally padding length (see below).
257
+ Will truncate by taking into account the special tokens.
258
+ padding_strategy: PaddingStrategy to use for padding.
259
+
260
+ - PaddingStrategy.LONGEST Pad to the longest sequence in the batch
261
+ - PaddingStrategy.MAX_LENGTH: Pad to the max length (default)
262
+ - PaddingStrategy.DO_NOT_PAD: Do not pad
263
+ The tokenizer padding sides are defined in self.padding_side:
264
+
265
+ - 'left': pads on the left of the sequences
266
+ - 'right': pads on the right of the sequences
267
+ pad_to_multiple_of: (optional) Integer if set will pad the sequence to a multiple of the provided value.
268
+ This is especially useful to enable the use of Tensor Core on NVIDIA hardware with compute capability
269
+ `>= 7.5` (Volta).
270
+ return_attention_mask:
271
+ (optional) Set to False to avoid returning attention mask (default: set to model specifics)
272
+ """
273
+ # Load from model defaults
274
+ assert self.padding_side == "left"
275
+
276
+ required_input = encoded_inputs[self.model_input_names[0]]
277
+ seq_length = len(required_input)
278
+
279
+ if padding_strategy == PaddingStrategy.LONGEST:
280
+ max_length = len(required_input)
281
+
282
+ if max_length is not None and pad_to_multiple_of is not None and (max_length % pad_to_multiple_of != 0):
283
+ max_length = ((max_length // pad_to_multiple_of) + 1) * pad_to_multiple_of
284
+
285
+ needs_to_be_padded = padding_strategy != PaddingStrategy.DO_NOT_PAD and len(required_input) != max_length
286
+
287
+ # Initialize attention mask if not present.
288
+ if "attention_mask" not in encoded_inputs:
289
+ encoded_inputs["attention_mask"] = [1] * seq_length
290
+
291
+ if "position_ids" not in encoded_inputs:
292
+ encoded_inputs["position_ids"] = list(range(seq_length))
293
+
294
+ if needs_to_be_padded:
295
+ difference = max_length - len(required_input)
296
+
297
+ if "attention_mask" in encoded_inputs:
298
+ encoded_inputs["attention_mask"] = [0] * difference + encoded_inputs["attention_mask"]
299
+ if "position_ids" in encoded_inputs:
300
+ encoded_inputs["position_ids"] = [0] * difference + encoded_inputs["position_ids"]
301
+ encoded_inputs[self.model_input_names[0]] = [self.pad_token_id] * difference + required_input
302
+
303
+ return encoded_inputs
304
+
305
+
306
+
307
+ class KolorsPrompter(BasePrompter):
308
+ def __init__(
309
+ self,
310
+ tokenizer_path=None
311
+ ):
312
+ if tokenizer_path is None:
313
+ base_path = os.path.dirname(os.path.dirname(__file__))
314
+ tokenizer_path = os.path.join(base_path, "tokenizer_configs/kolors/tokenizer")
315
+ super().__init__()
316
+ self.tokenizer = ChatGLMTokenizer.from_pretrained(tokenizer_path)
317
+ self.text_encoder: ChatGLMModel = None
318
+
319
+
320
+ def fetch_models(self, text_encoder: ChatGLMModel = None):
321
+ self.text_encoder = text_encoder
322
+
323
+
324
+ def encode_prompt_using_ChatGLM(self, prompt, text_encoder, tokenizer, max_length, clip_skip, device):
325
+ text_inputs = tokenizer(
326
+ prompt,
327
+ padding="max_length",
328
+ max_length=max_length,
329
+ truncation=True,
330
+ return_tensors="pt",
331
+ ).to(device)
332
+ output = text_encoder(
333
+ input_ids=text_inputs['input_ids'] ,
334
+ attention_mask=text_inputs['attention_mask'],
335
+ position_ids=text_inputs['position_ids'],
336
+ output_hidden_states=True
337
+ )
338
+ prompt_emb = output.hidden_states[-clip_skip].permute(1, 0, 2).clone()
339
+ pooled_prompt_emb = output.hidden_states[-1][-1, :, :].clone()
340
+ return prompt_emb, pooled_prompt_emb
341
+
342
+
343
+ def encode_prompt(
344
+ self,
345
+ prompt,
346
+ clip_skip=1,
347
+ clip_skip_2=2,
348
+ positive=True,
349
+ device="cuda"
350
+ ):
351
+ prompt = self.process_prompt(prompt, positive=positive)
352
+ prompt_emb, pooled_prompt_emb = self.encode_prompt_using_ChatGLM(prompt, self.text_encoder, self.tokenizer, 256, clip_skip_2, device)
353
+
354
+ return pooled_prompt_emb, prompt_emb
diffsynth/prompters/omnigen_prompter.py ADDED
@@ -0,0 +1,356 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import re
3
+ from typing import Dict, List
4
+
5
+ import torch
6
+ from PIL import Image
7
+ from torchvision import transforms
8
+ from transformers import AutoTokenizer
9
+ from huggingface_hub import snapshot_download
10
+ import numpy as np
11
+
12
+
13
+
14
+ def crop_arr(pil_image, max_image_size):
15
+ while min(*pil_image.size) >= 2 * max_image_size:
16
+ pil_image = pil_image.resize(
17
+ tuple(x // 2 for x in pil_image.size), resample=Image.BOX
18
+ )
19
+
20
+ if max(*pil_image.size) > max_image_size:
21
+ scale = max_image_size / max(*pil_image.size)
22
+ pil_image = pil_image.resize(
23
+ tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC
24
+ )
25
+
26
+ if min(*pil_image.size) < 16:
27
+ scale = 16 / min(*pil_image.size)
28
+ pil_image = pil_image.resize(
29
+ tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC
30
+ )
31
+
32
+ arr = np.array(pil_image)
33
+ crop_y1 = (arr.shape[0] % 16) // 2
34
+ crop_y2 = arr.shape[0] % 16 - crop_y1
35
+
36
+ crop_x1 = (arr.shape[1] % 16) // 2
37
+ crop_x2 = arr.shape[1] % 16 - crop_x1
38
+
39
+ arr = arr[crop_y1:arr.shape[0]-crop_y2, crop_x1:arr.shape[1]-crop_x2]
40
+ return Image.fromarray(arr)
41
+
42
+
43
+
44
+ class OmniGenPrompter:
45
+ def __init__(self,
46
+ text_tokenizer,
47
+ max_image_size: int=1024):
48
+ self.text_tokenizer = text_tokenizer
49
+ self.max_image_size = max_image_size
50
+
51
+ self.image_transform = transforms.Compose([
52
+ transforms.Lambda(lambda pil_image: crop_arr(pil_image, max_image_size)),
53
+ transforms.ToTensor(),
54
+ transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
55
+ ])
56
+
57
+ self.collator = OmniGenCollator()
58
+ self.separate_collator = OmniGenSeparateCollator()
59
+
60
+ @classmethod
61
+ def from_pretrained(cls, model_name):
62
+ if not os.path.exists(model_name):
63
+ cache_folder = os.getenv('HF_HUB_CACHE')
64
+ model_name = snapshot_download(repo_id=model_name,
65
+ cache_dir=cache_folder,
66
+ allow_patterns="*.json")
67
+ text_tokenizer = AutoTokenizer.from_pretrained(model_name)
68
+
69
+ return cls(text_tokenizer)
70
+
71
+
72
+ def process_image(self, image):
73
+ return self.image_transform(image)
74
+
75
+ def process_multi_modal_prompt(self, text, input_images):
76
+ text = self.add_prefix_instruction(text)
77
+ if input_images is None or len(input_images) == 0:
78
+ model_inputs = self.text_tokenizer(text)
79
+ return {"input_ids": model_inputs.input_ids, "pixel_values": None, "image_sizes": None}
80
+
81
+ pattern = r"<\|image_\d+\|>"
82
+ prompt_chunks = [self.text_tokenizer(chunk).input_ids for chunk in re.split(pattern, text)]
83
+
84
+ for i in range(1, len(prompt_chunks)):
85
+ if prompt_chunks[i][0] == 1:
86
+ prompt_chunks[i] = prompt_chunks[i][1:]
87
+
88
+ image_tags = re.findall(pattern, text)
89
+ image_ids = [int(s.split("|")[1].split("_")[-1]) for s in image_tags]
90
+
91
+ unique_image_ids = sorted(list(set(image_ids)))
92
+ assert unique_image_ids == list(range(1, len(unique_image_ids)+1)), f"image_ids must start from 1, and must be continuous int, e.g. [1, 2, 3], cannot be {unique_image_ids}"
93
+ # total images must be the same as the number of image tags
94
+ assert len(unique_image_ids) == len(input_images), f"total images must be the same as the number of image tags, got {len(unique_image_ids)} image tags and {len(input_images)} images"
95
+
96
+ input_images = [input_images[x-1] for x in image_ids]
97
+
98
+ all_input_ids = []
99
+ img_inx = []
100
+ idx = 0
101
+ for i in range(len(prompt_chunks)):
102
+ all_input_ids.extend(prompt_chunks[i])
103
+ if i != len(prompt_chunks) -1:
104
+ start_inx = len(all_input_ids)
105
+ size = input_images[i].size(-2) * input_images[i].size(-1) // 16 // 16
106
+ img_inx.append([start_inx, start_inx+size])
107
+ all_input_ids.extend([0]*size)
108
+
109
+ return {"input_ids": all_input_ids, "pixel_values": input_images, "image_sizes": img_inx}
110
+
111
+
112
+ def add_prefix_instruction(self, prompt):
113
+ user_prompt = '<|user|>\n'
114
+ generation_prompt = 'Generate an image according to the following instructions\n'
115
+ assistant_prompt = '<|assistant|>\n<|diffusion|>'
116
+ prompt_suffix = "<|end|>\n"
117
+ prompt = f"{user_prompt}{generation_prompt}{prompt}{prompt_suffix}{assistant_prompt}"
118
+ return prompt
119
+
120
+
121
+ def __call__(self,
122
+ instructions: List[str],
123
+ input_images: List[List[str]] = None,
124
+ height: int = 1024,
125
+ width: int = 1024,
126
+ negative_prompt: str = "low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers.",
127
+ use_img_cfg: bool = True,
128
+ separate_cfg_input: bool = False,
129
+ use_input_image_size_as_output: bool=False,
130
+ ) -> Dict:
131
+
132
+ if input_images is None:
133
+ use_img_cfg = False
134
+ if isinstance(instructions, str):
135
+ instructions = [instructions]
136
+ input_images = [input_images]
137
+
138
+ input_data = []
139
+ for i in range(len(instructions)):
140
+ cur_instruction = instructions[i]
141
+ cur_input_images = None if input_images is None else input_images[i]
142
+ if cur_input_images is not None and len(cur_input_images) > 0:
143
+ cur_input_images = [self.process_image(x) for x in cur_input_images]
144
+ else:
145
+ cur_input_images = None
146
+ assert "<img><|image_1|></img>" not in cur_instruction
147
+
148
+ mllm_input = self.process_multi_modal_prompt(cur_instruction, cur_input_images)
149
+
150
+
151
+ neg_mllm_input, img_cfg_mllm_input = None, None
152
+ neg_mllm_input = self.process_multi_modal_prompt(negative_prompt, None)
153
+ if use_img_cfg:
154
+ if cur_input_images is not None and len(cur_input_images) >= 1:
155
+ img_cfg_prompt = [f"<img><|image_{i+1}|></img>" for i in range(len(cur_input_images))]
156
+ img_cfg_mllm_input = self.process_multi_modal_prompt(" ".join(img_cfg_prompt), cur_input_images)
157
+ else:
158
+ img_cfg_mllm_input = neg_mllm_input
159
+
160
+ if use_input_image_size_as_output:
161
+ input_data.append((mllm_input, neg_mllm_input, img_cfg_mllm_input, [mllm_input['pixel_values'][0].size(-2), mllm_input['pixel_values'][0].size(-1)]))
162
+ else:
163
+ input_data.append((mllm_input, neg_mllm_input, img_cfg_mllm_input, [height, width]))
164
+
165
+ if separate_cfg_input:
166
+ return self.separate_collator(input_data)
167
+ return self.collator(input_data)
168
+
169
+
170
+
171
+
172
+ class OmniGenCollator:
173
+ def __init__(self, pad_token_id=2, hidden_size=3072):
174
+ self.pad_token_id = pad_token_id
175
+ self.hidden_size = hidden_size
176
+
177
+ def create_position(self, attention_mask, num_tokens_for_output_images):
178
+ position_ids = []
179
+ text_length = attention_mask.size(-1)
180
+ img_length = max(num_tokens_for_output_images)
181
+ for mask in attention_mask:
182
+ temp_l = torch.sum(mask)
183
+ temp_position = [0]*(text_length-temp_l) + [i for i in range(temp_l+img_length+1)] # we add a time embedding into the sequence, so add one more token
184
+ position_ids.append(temp_position)
185
+ return torch.LongTensor(position_ids)
186
+
187
+ def create_mask(self, attention_mask, num_tokens_for_output_images):
188
+ extended_mask = []
189
+ padding_images = []
190
+ text_length = attention_mask.size(-1)
191
+ img_length = max(num_tokens_for_output_images)
192
+ seq_len = text_length + img_length + 1 # we add a time embedding into the sequence, so add one more token
193
+ inx = 0
194
+ for mask in attention_mask:
195
+ temp_l = torch.sum(mask)
196
+ pad_l = text_length - temp_l
197
+
198
+ temp_mask = torch.tril(torch.ones(size=(temp_l+1, temp_l+1)))
199
+
200
+ image_mask = torch.zeros(size=(temp_l+1, img_length))
201
+ temp_mask = torch.cat([temp_mask, image_mask], dim=-1)
202
+
203
+ image_mask = torch.ones(size=(img_length, temp_l+img_length+1))
204
+ temp_mask = torch.cat([temp_mask, image_mask], dim=0)
205
+
206
+ if pad_l > 0:
207
+ pad_mask = torch.zeros(size=(temp_l+1+img_length, pad_l))
208
+ temp_mask = torch.cat([pad_mask, temp_mask], dim=-1)
209
+
210
+ pad_mask = torch.ones(size=(pad_l, seq_len))
211
+ temp_mask = torch.cat([pad_mask, temp_mask], dim=0)
212
+
213
+ true_img_length = num_tokens_for_output_images[inx]
214
+ pad_img_length = img_length - true_img_length
215
+ if pad_img_length > 0:
216
+ temp_mask[:, -pad_img_length:] = 0
217
+ temp_padding_imgs = torch.zeros(size=(1, pad_img_length, self.hidden_size))
218
+ else:
219
+ temp_padding_imgs = None
220
+
221
+ extended_mask.append(temp_mask.unsqueeze(0))
222
+ padding_images.append(temp_padding_imgs)
223
+ inx += 1
224
+ return torch.cat(extended_mask, dim=0), padding_images
225
+
226
+ def adjust_attention_for_input_images(self, attention_mask, image_sizes):
227
+ for b_inx in image_sizes.keys():
228
+ for start_inx, end_inx in image_sizes[b_inx]:
229
+ attention_mask[b_inx][start_inx:end_inx, start_inx:end_inx] = 1
230
+
231
+ return attention_mask
232
+
233
+ def pad_input_ids(self, input_ids, image_sizes):
234
+ max_l = max([len(x) for x in input_ids])
235
+ padded_ids = []
236
+ attention_mask = []
237
+ new_image_sizes = []
238
+
239
+ for i in range(len(input_ids)):
240
+ temp_ids = input_ids[i]
241
+ temp_l = len(temp_ids)
242
+ pad_l = max_l - temp_l
243
+ if pad_l == 0:
244
+ attention_mask.append([1]*max_l)
245
+ padded_ids.append(temp_ids)
246
+ else:
247
+ attention_mask.append([0]*pad_l+[1]*temp_l)
248
+ padded_ids.append([self.pad_token_id]*pad_l+temp_ids)
249
+
250
+ if i in image_sizes:
251
+ new_inx = []
252
+ for old_inx in image_sizes[i]:
253
+ new_inx.append([x+pad_l for x in old_inx])
254
+ image_sizes[i] = new_inx
255
+
256
+ return torch.LongTensor(padded_ids), torch.LongTensor(attention_mask), image_sizes
257
+
258
+
259
+ def process_mllm_input(self, mllm_inputs, target_img_size):
260
+ num_tokens_for_output_images = []
261
+ for img_size in target_img_size:
262
+ num_tokens_for_output_images.append(img_size[0]*img_size[1]//16//16)
263
+
264
+ pixel_values, image_sizes = [], {}
265
+ b_inx = 0
266
+ for x in mllm_inputs:
267
+ if x['pixel_values'] is not None:
268
+ pixel_values.extend(x['pixel_values'])
269
+ for size in x['image_sizes']:
270
+ if b_inx not in image_sizes:
271
+ image_sizes[b_inx] = [size]
272
+ else:
273
+ image_sizes[b_inx].append(size)
274
+ b_inx += 1
275
+ pixel_values = [x.unsqueeze(0) for x in pixel_values]
276
+
277
+
278
+ input_ids = [x['input_ids'] for x in mllm_inputs]
279
+ padded_input_ids, attention_mask, image_sizes = self.pad_input_ids(input_ids, image_sizes)
280
+ position_ids = self.create_position(attention_mask, num_tokens_for_output_images)
281
+ attention_mask, padding_images = self.create_mask(attention_mask, num_tokens_for_output_images)
282
+ attention_mask = self.adjust_attention_for_input_images(attention_mask, image_sizes)
283
+
284
+ return padded_input_ids, position_ids, attention_mask, padding_images, pixel_values, image_sizes
285
+
286
+
287
+ def __call__(self, features):
288
+ mllm_inputs = [f[0] for f in features]
289
+ cfg_mllm_inputs = [f[1] for f in features]
290
+ img_cfg_mllm_input = [f[2] for f in features]
291
+ target_img_size = [f[3] for f in features]
292
+
293
+
294
+ if img_cfg_mllm_input[0] is not None:
295
+ mllm_inputs = mllm_inputs + cfg_mllm_inputs + img_cfg_mllm_input
296
+ target_img_size = target_img_size + target_img_size + target_img_size
297
+ else:
298
+ mllm_inputs = mllm_inputs + cfg_mllm_inputs
299
+ target_img_size = target_img_size + target_img_size
300
+
301
+
302
+ all_padded_input_ids, all_position_ids, all_attention_mask, all_padding_images, all_pixel_values, all_image_sizes = self.process_mllm_input(mllm_inputs, target_img_size)
303
+
304
+ data = {"input_ids": all_padded_input_ids,
305
+ "attention_mask": all_attention_mask,
306
+ "position_ids": all_position_ids,
307
+ "input_pixel_values": all_pixel_values,
308
+ "input_image_sizes": all_image_sizes,
309
+ "padding_images": all_padding_images,
310
+ }
311
+ return data
312
+
313
+
314
+ class OmniGenSeparateCollator(OmniGenCollator):
315
+ def __call__(self, features):
316
+ mllm_inputs = [f[0] for f in features]
317
+ cfg_mllm_inputs = [f[1] for f in features]
318
+ img_cfg_mllm_input = [f[2] for f in features]
319
+ target_img_size = [f[3] for f in features]
320
+
321
+ all_padded_input_ids, all_attention_mask, all_position_ids, all_pixel_values, all_image_sizes, all_padding_images = [], [], [], [], [], []
322
+
323
+
324
+ padded_input_ids, position_ids, attention_mask, padding_images, pixel_values, image_sizes = self.process_mllm_input(mllm_inputs, target_img_size)
325
+ all_padded_input_ids.append(padded_input_ids)
326
+ all_attention_mask.append(attention_mask)
327
+ all_position_ids.append(position_ids)
328
+ all_pixel_values.append(pixel_values)
329
+ all_image_sizes.append(image_sizes)
330
+ all_padding_images.append(padding_images)
331
+
332
+ if cfg_mllm_inputs[0] is not None:
333
+ padded_input_ids, position_ids, attention_mask, padding_images, pixel_values, image_sizes = self.process_mllm_input(cfg_mllm_inputs, target_img_size)
334
+ all_padded_input_ids.append(padded_input_ids)
335
+ all_attention_mask.append(attention_mask)
336
+ all_position_ids.append(position_ids)
337
+ all_pixel_values.append(pixel_values)
338
+ all_image_sizes.append(image_sizes)
339
+ all_padding_images.append(padding_images)
340
+ if img_cfg_mllm_input[0] is not None:
341
+ padded_input_ids, position_ids, attention_mask, padding_images, pixel_values, image_sizes = self.process_mllm_input(img_cfg_mllm_input, target_img_size)
342
+ all_padded_input_ids.append(padded_input_ids)
343
+ all_attention_mask.append(attention_mask)
344
+ all_position_ids.append(position_ids)
345
+ all_pixel_values.append(pixel_values)
346
+ all_image_sizes.append(image_sizes)
347
+ all_padding_images.append(padding_images)
348
+
349
+ data = {"input_ids": all_padded_input_ids,
350
+ "attention_mask": all_attention_mask,
351
+ "position_ids": all_position_ids,
352
+ "input_pixel_values": all_pixel_values,
353
+ "input_image_sizes": all_image_sizes,
354
+ "padding_images": all_padding_images,
355
+ }
356
+ return data
diffsynth/prompters/omost.py ADDED
@@ -0,0 +1,323 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import AutoTokenizer, TextIteratorStreamer
2
+ import difflib
3
+ import torch
4
+ import numpy as np
5
+ import re
6
+ from ..models.model_manager import ModelManager
7
+ from PIL import Image
8
+
9
+ valid_colors = { # r, g, b
10
+ 'aliceblue': (240, 248, 255), 'antiquewhite': (250, 235, 215), 'aqua': (0, 255, 255),
11
+ 'aquamarine': (127, 255, 212), 'azure': (240, 255, 255), 'beige': (245, 245, 220),
12
+ 'bisque': (255, 228, 196), 'black': (0, 0, 0), 'blanchedalmond': (255, 235, 205), 'blue': (0, 0, 255),
13
+ 'blueviolet': (138, 43, 226), 'brown': (165, 42, 42), 'burlywood': (222, 184, 135),
14
+ 'cadetblue': (95, 158, 160), 'chartreuse': (127, 255, 0), 'chocolate': (210, 105, 30),
15
+ 'coral': (255, 127, 80), 'cornflowerblue': (100, 149, 237), 'cornsilk': (255, 248, 220),
16
+ 'crimson': (220, 20, 60), 'cyan': (0, 255, 255), 'darkblue': (0, 0, 139), 'darkcyan': (0, 139, 139),
17
+ 'darkgoldenrod': (184, 134, 11), 'darkgray': (169, 169, 169), 'darkgrey': (169, 169, 169),
18
+ 'darkgreen': (0, 100, 0), 'darkkhaki': (189, 183, 107), 'darkmagenta': (139, 0, 139),
19
+ 'darkolivegreen': (85, 107, 47), 'darkorange': (255, 140, 0), 'darkorchid': (153, 50, 204),
20
+ 'darkred': (139, 0, 0), 'darksalmon': (233, 150, 122), 'darkseagreen': (143, 188, 143),
21
+ 'darkslateblue': (72, 61, 139), 'darkslategray': (47, 79, 79), 'darkslategrey': (47, 79, 79),
22
+ 'darkturquoise': (0, 206, 209), 'darkviolet': (148, 0, 211), 'deeppink': (255, 20, 147),
23
+ 'deepskyblue': (0, 191, 255), 'dimgray': (105, 105, 105), 'dimgrey': (105, 105, 105),
24
+ 'dodgerblue': (30, 144, 255), 'firebrick': (178, 34, 34), 'floralwhite': (255, 250, 240),
25
+ 'forestgreen': (34, 139, 34), 'fuchsia': (255, 0, 255), 'gainsboro': (220, 220, 220),
26
+ 'ghostwhite': (248, 248, 255), 'gold': (255, 215, 0), 'goldenrod': (218, 165, 32),
27
+ 'gray': (128, 128, 128), 'grey': (128, 128, 128), 'green': (0, 128, 0), 'greenyellow': (173, 255, 47),
28
+ 'honeydew': (240, 255, 240), 'hotpink': (255, 105, 180), 'indianred': (205, 92, 92),
29
+ 'indigo': (75, 0, 130), 'ivory': (255, 255, 240), 'khaki': (240, 230, 140), 'lavender': (230, 230, 250),
30
+ 'lavenderblush': (255, 240, 245), 'lawngreen': (124, 252, 0), 'lemonchiffon': (255, 250, 205),
31
+ 'lightblue': (173, 216, 230), 'lightcoral': (240, 128, 128), 'lightcyan': (224, 255, 255),
32
+ 'lightgoldenrodyellow': (250, 250, 210), 'lightgray': (211, 211, 211), 'lightgrey': (211, 211, 211),
33
+ 'lightgreen': (144, 238, 144), 'lightpink': (255, 182, 193), 'lightsalmon': (255, 160, 122),
34
+ 'lightseagreen': (32, 178, 170), 'lightskyblue': (135, 206, 250), 'lightslategray': (119, 136, 153),
35
+ 'lightslategrey': (119, 136, 153), 'lightsteelblue': (176, 196, 222), 'lightyellow': (255, 255, 224),
36
+ 'lime': (0, 255, 0), 'limegreen': (50, 205, 50), 'linen': (250, 240, 230), 'magenta': (255, 0, 255),
37
+ 'maroon': (128, 0, 0), 'mediumaquamarine': (102, 205, 170), 'mediumblue': (0, 0, 205),
38
+ 'mediumorchid': (186, 85, 211), 'mediumpurple': (147, 112, 219), 'mediumseagreen': (60, 179, 113),
39
+ 'mediumslateblue': (123, 104, 238), 'mediumspringgreen': (0, 250, 154),
40
+ 'mediumturquoise': (72, 209, 204), 'mediumvioletred': (199, 21, 133), 'midnightblue': (25, 25, 112),
41
+ 'mintcream': (245, 255, 250), 'mistyrose': (255, 228, 225), 'moccasin': (255, 228, 181),
42
+ 'navajowhite': (255, 222, 173), 'navy': (0, 0, 128), 'navyblue': (0, 0, 128),
43
+ 'oldlace': (253, 245, 230), 'olive': (128, 128, 0), 'olivedrab': (107, 142, 35),
44
+ 'orange': (255, 165, 0), 'orangered': (255, 69, 0), 'orchid': (218, 112, 214),
45
+ 'palegoldenrod': (238, 232, 170), 'palegreen': (152, 251, 152), 'paleturquoise': (175, 238, 238),
46
+ 'palevioletred': (219, 112, 147), 'papayawhip': (255, 239, 213), 'peachpuff': (255, 218, 185),
47
+ 'peru': (205, 133, 63), 'pink': (255, 192, 203), 'plum': (221, 160, 221), 'powderblue': (176, 224, 230),
48
+ 'purple': (128, 0, 128), 'rebeccapurple': (102, 51, 153), 'red': (255, 0, 0),
49
+ 'rosybrown': (188, 143, 143), 'royalblue': (65, 105, 225), 'saddlebrown': (139, 69, 19),
50
+ 'salmon': (250, 128, 114), 'sandybrown': (244, 164, 96), 'seagreen': (46, 139, 87),
51
+ 'seashell': (255, 245, 238), 'sienna': (160, 82, 45), 'silver': (192, 192, 192),
52
+ 'skyblue': (135, 206, 235), 'slateblue': (106, 90, 205), 'slategray': (112, 128, 144),
53
+ 'slategrey': (112, 128, 144), 'snow': (255, 250, 250), 'springgreen': (0, 255, 127),
54
+ 'steelblue': (70, 130, 180), 'tan': (210, 180, 140), 'teal': (0, 128, 128), 'thistle': (216, 191, 216),
55
+ 'tomato': (255, 99, 71), 'turquoise': (64, 224, 208), 'violet': (238, 130, 238),
56
+ 'wheat': (245, 222, 179), 'white': (255, 255, 255), 'whitesmoke': (245, 245, 245),
57
+ 'yellow': (255, 255, 0), 'yellowgreen': (154, 205, 50)
58
+ }
59
+
60
+ valid_locations = { # x, y in 90*90
61
+ 'in the center': (45, 45),
62
+ 'on the left': (15, 45),
63
+ 'on the right': (75, 45),
64
+ 'on the top': (45, 15),
65
+ 'on the bottom': (45, 75),
66
+ 'on the top-left': (15, 15),
67
+ 'on the top-right': (75, 15),
68
+ 'on the bottom-left': (15, 75),
69
+ 'on the bottom-right': (75, 75)
70
+ }
71
+
72
+ valid_offsets = { # x, y in 90*90
73
+ 'no offset': (0, 0),
74
+ 'slightly to the left': (-10, 0),
75
+ 'slightly to the right': (10, 0),
76
+ 'slightly to the upper': (0, -10),
77
+ 'slightly to the lower': (0, 10),
78
+ 'slightly to the upper-left': (-10, -10),
79
+ 'slightly to the upper-right': (10, -10),
80
+ 'slightly to the lower-left': (-10, 10),
81
+ 'slightly to the lower-right': (10, 10)}
82
+
83
+ valid_areas = { # w, h in 90*90
84
+ "a small square area": (50, 50),
85
+ "a small vertical area": (40, 60),
86
+ "a small horizontal area": (60, 40),
87
+ "a medium-sized square area": (60, 60),
88
+ "a medium-sized vertical area": (50, 80),
89
+ "a medium-sized horizontal area": (80, 50),
90
+ "a large square area": (70, 70),
91
+ "a large vertical area": (60, 90),
92
+ "a large horizontal area": (90, 60)
93
+ }
94
+
95
+ def safe_str(x):
96
+ return x.strip(',. ') + '.'
97
+
98
+ def closest_name(input_str, options):
99
+ input_str = input_str.lower()
100
+
101
+ closest_match = difflib.get_close_matches(input_str, list(options.keys()), n=1, cutoff=0.5)
102
+ assert isinstance(closest_match, list) and len(closest_match) > 0, f'The value [{input_str}] is not valid!'
103
+ result = closest_match[0]
104
+
105
+ if result != input_str:
106
+ print(f'Automatically corrected [{input_str}] -> [{result}].')
107
+
108
+ return result
109
+
110
+ class Canvas:
111
+ @staticmethod
112
+ def from_bot_response(response: str):
113
+
114
+ matched = re.search(r'```python\n(.*?)\n```', response, re.DOTALL)
115
+ assert matched, 'Response does not contain codes!'
116
+ code_content = matched.group(1)
117
+ assert 'canvas = Canvas()' in code_content, 'Code block must include valid canvas var!'
118
+ local_vars = {'Canvas': Canvas}
119
+ exec(code_content, {}, local_vars)
120
+ canvas = local_vars.get('canvas', None)
121
+ assert isinstance(canvas, Canvas), 'Code block must produce valid canvas var!'
122
+ return canvas
123
+
124
+ def __init__(self):
125
+ self.components = []
126
+ self.color = None
127
+ self.record_tags = True
128
+ self.prefixes = []
129
+ self.suffixes = []
130
+ return
131
+
132
+ def set_global_description(self, description: str, detailed_descriptions: list, tags: str,
133
+ HTML_web_color_name: str):
134
+ assert isinstance(description, str), 'Global description is not valid!'
135
+ assert isinstance(detailed_descriptions, list) and all(isinstance(item, str) for item in detailed_descriptions), \
136
+ 'Global detailed_descriptions is not valid!'
137
+ assert isinstance(tags, str), 'Global tags is not valid!'
138
+
139
+ HTML_web_color_name = closest_name(HTML_web_color_name, valid_colors)
140
+ self.color = np.array([[valid_colors[HTML_web_color_name]]], dtype=np.uint8)
141
+
142
+ self.prefixes = [description]
143
+ self.suffixes = detailed_descriptions
144
+
145
+ if self.record_tags:
146
+ self.suffixes = self.suffixes + [tags]
147
+
148
+ self.prefixes = [safe_str(x) for x in self.prefixes]
149
+ self.suffixes = [safe_str(x) for x in self.suffixes]
150
+
151
+ return
152
+
153
+ def add_local_description(self, location: str, offset: str, area: str, distance_to_viewer: float, description: str,
154
+ detailed_descriptions: list, tags: str, atmosphere: str, style: str,
155
+ quality_meta: str, HTML_web_color_name: str):
156
+ assert isinstance(description, str), 'Local description is wrong!'
157
+ assert isinstance(distance_to_viewer, (int, float)) and distance_to_viewer > 0, \
158
+ f'The distance_to_viewer for [{description}] is not positive float number!'
159
+ assert isinstance(detailed_descriptions, list) and all(isinstance(item, str) for item in detailed_descriptions), \
160
+ f'The detailed_descriptions for [{description}] is not valid!'
161
+ assert isinstance(tags, str), f'The tags for [{description}] is not valid!'
162
+ assert isinstance(atmosphere, str), f'The atmosphere for [{description}] is not valid!'
163
+ assert isinstance(style, str), f'The style for [{description}] is not valid!'
164
+ assert isinstance(quality_meta, str), f'The quality_meta for [{description}] is not valid!'
165
+
166
+ location = closest_name(location, valid_locations)
167
+ offset = closest_name(offset, valid_offsets)
168
+ area = closest_name(area, valid_areas)
169
+ HTML_web_color_name = closest_name(HTML_web_color_name, valid_colors)
170
+
171
+ xb, yb = valid_locations[location]
172
+ xo, yo = valid_offsets[offset]
173
+ w, h = valid_areas[area]
174
+ rect = (yb + yo - h // 2, yb + yo + h // 2, xb + xo - w // 2, xb + xo + w // 2)
175
+ rect = [max(0, min(90, i)) for i in rect]
176
+ color = np.array([[valid_colors[HTML_web_color_name]]], dtype=np.uint8)
177
+
178
+ prefixes = self.prefixes + [description]
179
+ suffixes = detailed_descriptions
180
+
181
+ if self.record_tags:
182
+ suffixes = suffixes + [tags, atmosphere, style, quality_meta]
183
+
184
+ prefixes = [safe_str(x) for x in prefixes]
185
+ suffixes = [safe_str(x) for x in suffixes]
186
+
187
+ self.components.append(dict(
188
+ rect=rect,
189
+ distance_to_viewer=distance_to_viewer,
190
+ color=color,
191
+ prefixes=prefixes,
192
+ suffixes=suffixes,
193
+ location=location,
194
+ ))
195
+
196
+ return
197
+
198
+ def process(self):
199
+ # sort components
200
+ self.components = sorted(self.components, key=lambda x: x['distance_to_viewer'], reverse=True)
201
+
202
+ # compute initial latent
203
+ # print(self.color)
204
+ initial_latent = np.zeros(shape=(90, 90, 3), dtype=np.float32) + self.color
205
+
206
+ for component in self.components:
207
+ a, b, c, d = component['rect']
208
+ initial_latent[a:b, c:d] = 0.7 * component['color'] + 0.3 * initial_latent[a:b, c:d]
209
+
210
+ initial_latent = initial_latent.clip(0, 255).astype(np.uint8)
211
+
212
+ # compute conditions
213
+
214
+ bag_of_conditions = [
215
+ dict(mask=np.ones(shape=(90, 90), dtype=np.float32), prefixes=self.prefixes, suffixes=self.suffixes,location= "full")
216
+ ]
217
+
218
+ for i, component in enumerate(self.components):
219
+ a, b, c, d = component['rect']
220
+ m = np.zeros(shape=(90, 90), dtype=np.float32)
221
+ m[a:b, c:d] = 1.0
222
+ bag_of_conditions.append(dict(
223
+ mask = m,
224
+ prefixes = component['prefixes'],
225
+ suffixes = component['suffixes'],
226
+ location = component['location'],
227
+ ))
228
+
229
+ return dict(
230
+ initial_latent = initial_latent,
231
+ bag_of_conditions = bag_of_conditions,
232
+ )
233
+
234
+
235
+ class OmostPromter(torch.nn.Module):
236
+
237
+ def __init__(self,model = None,tokenizer = None, template = "",device="cpu"):
238
+ super().__init__()
239
+ self.model=model
240
+ self.tokenizer = tokenizer
241
+ self.device = device
242
+ if template == "":
243
+ template = r'''You are a helpful AI assistant to compose images using the below python class `Canvas`:
244
+ ```python
245
+ class Canvas:
246
+ def set_global_description(self, description: str, detailed_descriptions: list[str], tags: str, HTML_web_color_name: str):
247
+ pass
248
+
249
+ def add_local_description(self, location: str, offset: str, area: str, distance_to_viewer: float, description: str, detailed_descriptions: list[str], tags: str, atmosphere: str, style: str, quality_meta: str, HTML_web_color_name: str):
250
+ assert location in ["in the center", "on the left", "on the right", "on the top", "on the bottom", "on the top-left", "on the top-right", "on the bottom-left", "on the bottom-right"]
251
+ assert offset in ["no offset", "slightly to the left", "slightly to the right", "slightly to the upper", "slightly to the lower", "slightly to the upper-left", "slightly to the upper-right", "slightly to the lower-left", "slightly to the lower-right"]
252
+ assert area in ["a small square area", "a small vertical area", "a small horizontal area", "a medium-sized square area", "a medium-sized vertical area", "a medium-sized horizontal area", "a large square area", "a large vertical area", "a large horizontal area"]
253
+ assert distance_to_viewer > 0
254
+ pass
255
+ ```'''
256
+ self.template = template
257
+
258
+ @staticmethod
259
+ def from_model_manager(model_manager: ModelManager):
260
+ model, model_path = model_manager.fetch_model("omost_prompt", require_model_path=True)
261
+ tokenizer = AutoTokenizer.from_pretrained(model_path)
262
+ omost = OmostPromter(
263
+ model= model,
264
+ tokenizer = tokenizer,
265
+ device = model_manager.device
266
+ )
267
+ return omost
268
+
269
+
270
+ def __call__(self,prompt_dict:dict):
271
+ raw_prompt=prompt_dict["prompt"]
272
+ conversation = [{"role": "system", "content": self.template}]
273
+ conversation.append({"role": "user", "content": raw_prompt})
274
+
275
+ input_ids = self.tokenizer.apply_chat_template(conversation, return_tensors="pt", add_generation_prompt=True).to(self.device)
276
+ streamer = TextIteratorStreamer(self.tokenizer, timeout=10.0, skip_prompt=True, skip_special_tokens=True)
277
+ attention_mask = torch.ones(input_ids.shape, dtype=torch.bfloat16, device=self.device)
278
+
279
+ generate_kwargs = dict(
280
+ input_ids = input_ids,
281
+ streamer = streamer,
282
+ # stopping_criteria=stopping_criteria,
283
+ # max_new_tokens=max_new_tokens,
284
+ do_sample = True,
285
+ attention_mask = attention_mask,
286
+ pad_token_id = self.tokenizer.eos_token_id,
287
+ # temperature=temperature,
288
+ # top_p=top_p,
289
+ )
290
+ self.model.generate(**generate_kwargs)
291
+ outputs = []
292
+ for text in streamer:
293
+ outputs.append(text)
294
+ llm_outputs = "".join(outputs)
295
+
296
+ canvas = Canvas.from_bot_response(llm_outputs)
297
+ canvas_output = canvas.process()
298
+
299
+ prompts = [" ".join(_["prefixes"]+_["suffixes"][:2]) for _ in canvas_output["bag_of_conditions"]]
300
+ canvas_output["prompt"] = prompts[0]
301
+ canvas_output["prompts"] = prompts[1:]
302
+
303
+ raw_masks = [_["mask"] for _ in canvas_output["bag_of_conditions"]]
304
+ masks=[]
305
+ for mask in raw_masks:
306
+ mask[mask>0.5]=255
307
+ mask = np.stack([mask] * 3, axis=-1).astype("uint8")
308
+ masks.append(Image.fromarray(mask))
309
+
310
+ canvas_output["masks"] = masks
311
+ prompt_dict.update(canvas_output)
312
+ print(f"Your prompt is extended by Omost:\n")
313
+ cnt = 0
314
+ for component,pmt in zip(canvas_output["bag_of_conditions"],prompts):
315
+ loc = component["location"]
316
+ cnt += 1
317
+ print(f"Component {cnt} - Location : {loc}\nPrompt:{pmt}\n")
318
+
319
+ return prompt_dict
320
+
321
+
322
+
323
+
diffsynth/prompters/prompt_refiners.py ADDED
@@ -0,0 +1,130 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import AutoTokenizer
2
+ from ..models.model_manager import ModelManager
3
+ import torch
4
+ from .omost import OmostPromter
5
+
6
+ class BeautifulPrompt(torch.nn.Module):
7
+ def __init__(self, tokenizer_path=None, model=None, template=""):
8
+ super().__init__()
9
+ self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
10
+ self.model = model
11
+ self.template = template
12
+
13
+
14
+ @staticmethod
15
+ def from_model_manager(model_manager: ModelManager):
16
+ model, model_path = model_manager.fetch_model("beautiful_prompt", require_model_path=True)
17
+ template = 'Instruction: Give a simple description of the image to generate a drawing prompt.\nInput: {raw_prompt}\nOutput:'
18
+ if model_path.endswith("v2"):
19
+ template = """Converts a simple image description into a prompt. \
20
+ Prompts are formatted as multiple related tags separated by commas, plus you can use () to increase the weight, [] to decrease the weight, \
21
+ or use a number to specify the weight. You should add appropriate words to make the images described in the prompt more aesthetically pleasing, \
22
+ but make sure there is a correlation between the input and output.\n\
23
+ ### Input: {raw_prompt}\n### Output:"""
24
+ beautiful_prompt = BeautifulPrompt(
25
+ tokenizer_path=model_path,
26
+ model=model,
27
+ template=template
28
+ )
29
+ return beautiful_prompt
30
+
31
+
32
+ def __call__(self, raw_prompt, positive=True, **kwargs):
33
+ if positive:
34
+ model_input = self.template.format(raw_prompt=raw_prompt)
35
+ input_ids = self.tokenizer.encode(model_input, return_tensors='pt').to(self.model.device)
36
+ outputs = self.model.generate(
37
+ input_ids,
38
+ max_new_tokens=384,
39
+ do_sample=True,
40
+ temperature=0.9,
41
+ top_k=50,
42
+ top_p=0.95,
43
+ repetition_penalty=1.1,
44
+ num_return_sequences=1
45
+ )
46
+ prompt = raw_prompt + ", " + self.tokenizer.batch_decode(
47
+ outputs[:, input_ids.size(1):],
48
+ skip_special_tokens=True
49
+ )[0].strip()
50
+ print(f"Your prompt is refined by BeautifulPrompt: {prompt}")
51
+ return prompt
52
+ else:
53
+ return raw_prompt
54
+
55
+
56
+
57
+ class QwenPrompt(torch.nn.Module):
58
+ # This class leverages the open-source Qwen model to translate Chinese prompts into English,
59
+ # with an integrated optimization mechanism for enhanced translation quality.
60
+ def __init__(self, tokenizer_path=None, model=None, system_prompt=""):
61
+ super().__init__()
62
+ self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
63
+ self.model = model
64
+ self.system_prompt = system_prompt
65
+
66
+
67
+ @staticmethod
68
+ def from_model_manager(model_nameger: ModelManager):
69
+ model, model_path = model_nameger.fetch_model("qwen_prompt", require_model_path=True)
70
+ system_prompt = """You are an English image describer. Here are some example image styles:\n\n1. Extreme close-up: Clear focus on a single object with a blurred background, highlighted under natural sunlight.\n2. Vintage: A photograph of a historical scene, using techniques such as Daguerreotype or cyanotype.\n3. Anime: A stylized cartoon image, emphasizing hyper-realistic portraits and luminous brushwork.\n4. Candid: A natural, unposed shot capturing spontaneous moments, often with cinematic qualities.\n5. Landscape: A photorealistic image of natural scenery, such as a sunrise over the sea.\n6. Design: Colorful and detailed illustrations, often in the style of 2D game art or botanical illustrations.\n7. Urban: An ultrarealistic scene in a modern setting, possibly a cityscape viewed from indoors.\n\nYour task is to translate a given Chinese image description into a concise and precise English description. Ensure that the imagery is vivid and descriptive, and include stylistic elements to enrich the description.\nPlease note the following points:\n\n1. Capture the essence and mood of the Chinese description without including direct phrases or words from the examples provided.\n2. You should add appropriate words to make the images described in the prompt more aesthetically pleasing. If the Chinese description does not specify a style, you need to add some stylistic descriptions based on the essence of the Chinese text.\n3. The generated English description should not exceed 200 words.\n\n"""
71
+ qwen_prompt = QwenPrompt(
72
+ tokenizer_path=model_path,
73
+ model=model,
74
+ system_prompt=system_prompt
75
+ )
76
+ return qwen_prompt
77
+
78
+
79
+ def __call__(self, raw_prompt, positive=True, **kwargs):
80
+ if positive:
81
+ messages = [{
82
+ 'role': 'system',
83
+ 'content': self.system_prompt
84
+ }, {
85
+ 'role': 'user',
86
+ 'content': raw_prompt
87
+ }]
88
+ text = self.tokenizer.apply_chat_template(
89
+ messages,
90
+ tokenize=False,
91
+ add_generation_prompt=True
92
+ )
93
+ model_inputs = self.tokenizer([text], return_tensors="pt").to(self.model.device)
94
+
95
+ generated_ids = self.model.generate(
96
+ model_inputs.input_ids,
97
+ max_new_tokens=512
98
+ )
99
+ generated_ids = [
100
+ output_ids[len(input_ids):] for input_ids, output_ids in zip(model_inputs.input_ids, generated_ids)
101
+ ]
102
+
103
+ prompt = self.tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
104
+ print(f"Your prompt is refined by Qwen: {prompt}")
105
+ return prompt
106
+ else:
107
+ return raw_prompt
108
+
109
+
110
+
111
+ class Translator(torch.nn.Module):
112
+ def __init__(self, tokenizer_path=None, model=None):
113
+ super().__init__()
114
+ self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
115
+ self.model = model
116
+
117
+
118
+ @staticmethod
119
+ def from_model_manager(model_manager: ModelManager):
120
+ model, model_path = model_manager.fetch_model("translator", require_model_path=True)
121
+ translator = Translator(tokenizer_path=model_path, model=model)
122
+ return translator
123
+
124
+
125
+ def __call__(self, prompt, **kwargs):
126
+ input_ids = self.tokenizer.encode(prompt, return_tensors='pt').to(self.model.device)
127
+ output_ids = self.model.generate(input_ids)
128
+ prompt = self.tokenizer.batch_decode(output_ids, skip_special_tokens=True)[0]
129
+ print(f"Your prompt is translated: {prompt}")
130
+ return prompt
diffsynth/prompters/sd3_prompter.py ADDED
@@ -0,0 +1,93 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.model_manager import ModelManager
3
+ from ..models import SD3TextEncoder1, SD3TextEncoder2, SD3TextEncoder3
4
+ from transformers import CLIPTokenizer, T5TokenizerFast
5
+ import os, torch
6
+
7
+
8
+ class SD3Prompter(BasePrompter):
9
+ def __init__(
10
+ self,
11
+ tokenizer_1_path=None,
12
+ tokenizer_2_path=None,
13
+ tokenizer_3_path=None
14
+ ):
15
+ if tokenizer_1_path is None:
16
+ base_path = os.path.dirname(os.path.dirname(__file__))
17
+ tokenizer_1_path = os.path.join(base_path, "tokenizer_configs/stable_diffusion_3/tokenizer_1")
18
+ if tokenizer_2_path is None:
19
+ base_path = os.path.dirname(os.path.dirname(__file__))
20
+ tokenizer_2_path = os.path.join(base_path, "tokenizer_configs/stable_diffusion_3/tokenizer_2")
21
+ if tokenizer_3_path is None:
22
+ base_path = os.path.dirname(os.path.dirname(__file__))
23
+ tokenizer_3_path = os.path.join(base_path, "tokenizer_configs/stable_diffusion_3/tokenizer_3")
24
+ super().__init__()
25
+ self.tokenizer_1 = CLIPTokenizer.from_pretrained(tokenizer_1_path)
26
+ self.tokenizer_2 = CLIPTokenizer.from_pretrained(tokenizer_2_path)
27
+ self.tokenizer_3 = T5TokenizerFast.from_pretrained(tokenizer_3_path)
28
+ self.text_encoder_1: SD3TextEncoder1 = None
29
+ self.text_encoder_2: SD3TextEncoder2 = None
30
+ self.text_encoder_3: SD3TextEncoder3 = None
31
+
32
+
33
+ def fetch_models(self, text_encoder_1: SD3TextEncoder1 = None, text_encoder_2: SD3TextEncoder2 = None, text_encoder_3: SD3TextEncoder3 = None):
34
+ self.text_encoder_1 = text_encoder_1
35
+ self.text_encoder_2 = text_encoder_2
36
+ self.text_encoder_3 = text_encoder_3
37
+
38
+
39
+ def encode_prompt_using_clip(self, prompt, text_encoder, tokenizer, max_length, device):
40
+ input_ids = tokenizer(
41
+ prompt,
42
+ return_tensors="pt",
43
+ padding="max_length",
44
+ max_length=max_length,
45
+ truncation=True
46
+ ).input_ids.to(device)
47
+ pooled_prompt_emb, prompt_emb = text_encoder(input_ids)
48
+ return pooled_prompt_emb, prompt_emb
49
+
50
+
51
+ def encode_prompt_using_t5(self, prompt, text_encoder, tokenizer, max_length, device):
52
+ input_ids = tokenizer(
53
+ prompt,
54
+ return_tensors="pt",
55
+ padding="max_length",
56
+ max_length=max_length,
57
+ truncation=True,
58
+ add_special_tokens=True,
59
+ ).input_ids.to(device)
60
+ prompt_emb = text_encoder(input_ids)
61
+ prompt_emb = prompt_emb.reshape((1, prompt_emb.shape[0]*prompt_emb.shape[1], -1))
62
+
63
+ return prompt_emb
64
+
65
+
66
+ def encode_prompt(
67
+ self,
68
+ prompt,
69
+ positive=True,
70
+ device="cuda",
71
+ t5_sequence_length=77,
72
+ ):
73
+ prompt = self.process_prompt(prompt, positive=positive)
74
+
75
+ # CLIP
76
+ pooled_prompt_emb_1, prompt_emb_1 = self.encode_prompt_using_clip(prompt, self.text_encoder_1, self.tokenizer_1, 77, device)
77
+ pooled_prompt_emb_2, prompt_emb_2 = self.encode_prompt_using_clip(prompt, self.text_encoder_2, self.tokenizer_2, 77, device)
78
+
79
+ # T5
80
+ if self.text_encoder_3 is None:
81
+ prompt_emb_3 = torch.zeros((prompt_emb_1.shape[0], t5_sequence_length, 4096), dtype=prompt_emb_1.dtype, device=device)
82
+ else:
83
+ prompt_emb_3 = self.encode_prompt_using_t5(prompt, self.text_encoder_3, self.tokenizer_3, t5_sequence_length, device)
84
+ prompt_emb_3 = prompt_emb_3.to(prompt_emb_1.dtype) # float32 -> float16
85
+
86
+ # Merge
87
+ prompt_emb = torch.cat([
88
+ torch.nn.functional.pad(torch.cat([prompt_emb_1, prompt_emb_2], dim=-1), (0, 4096 - 768 - 1280)),
89
+ prompt_emb_3
90
+ ], dim=-2)
91
+ pooled_prompt_emb = torch.cat([pooled_prompt_emb_1, pooled_prompt_emb_2], dim=-1)
92
+
93
+ return prompt_emb, pooled_prompt_emb
diffsynth/prompters/sd_prompter.py ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter, tokenize_long_prompt
2
+ from ..models.utils import load_state_dict, search_for_embeddings
3
+ from ..models import SDTextEncoder
4
+ from transformers import CLIPTokenizer
5
+ import torch, os
6
+
7
+
8
+
9
+ class SDPrompter(BasePrompter):
10
+ def __init__(self, tokenizer_path=None):
11
+ if tokenizer_path is None:
12
+ base_path = os.path.dirname(os.path.dirname(__file__))
13
+ tokenizer_path = os.path.join(base_path, "tokenizer_configs/stable_diffusion/tokenizer")
14
+ super().__init__()
15
+ self.tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
16
+ self.text_encoder: SDTextEncoder = None
17
+ self.textual_inversion_dict = {}
18
+ self.keyword_dict = {}
19
+
20
+
21
+ def fetch_models(self, text_encoder: SDTextEncoder = None):
22
+ self.text_encoder = text_encoder
23
+
24
+
25
+ def add_textual_inversions_to_model(self, textual_inversion_dict, text_encoder):
26
+ dtype = next(iter(text_encoder.parameters())).dtype
27
+ state_dict = text_encoder.token_embedding.state_dict()
28
+ token_embeddings = [state_dict["weight"]]
29
+ for keyword in textual_inversion_dict:
30
+ _, embeddings = textual_inversion_dict[keyword]
31
+ token_embeddings.append(embeddings.to(dtype=dtype, device=token_embeddings[0].device))
32
+ token_embeddings = torch.concat(token_embeddings, dim=0)
33
+ state_dict["weight"] = token_embeddings
34
+ text_encoder.token_embedding = torch.nn.Embedding(token_embeddings.shape[0], token_embeddings.shape[1])
35
+ text_encoder.token_embedding = text_encoder.token_embedding.to(dtype=dtype, device=token_embeddings[0].device)
36
+ text_encoder.token_embedding.load_state_dict(state_dict)
37
+
38
+
39
+ def add_textual_inversions_to_tokenizer(self, textual_inversion_dict, tokenizer):
40
+ additional_tokens = []
41
+ for keyword in textual_inversion_dict:
42
+ tokens, _ = textual_inversion_dict[keyword]
43
+ additional_tokens += tokens
44
+ self.keyword_dict[keyword] = " " + " ".join(tokens) + " "
45
+ tokenizer.add_tokens(additional_tokens)
46
+
47
+
48
+ def load_textual_inversions(self, model_paths):
49
+ for model_path in model_paths:
50
+ keyword = os.path.splitext(os.path.split(model_path)[-1])[0]
51
+ state_dict = load_state_dict(model_path)
52
+
53
+ # Search for embeddings
54
+ for embeddings in search_for_embeddings(state_dict):
55
+ if len(embeddings.shape) == 2 and embeddings.shape[1] == 768:
56
+ tokens = [f"{keyword}_{i}" for i in range(embeddings.shape[0])]
57
+ self.textual_inversion_dict[keyword] = (tokens, embeddings)
58
+
59
+ self.add_textual_inversions_to_model(self.textual_inversion_dict, self.text_encoder)
60
+ self.add_textual_inversions_to_tokenizer(self.textual_inversion_dict, self.tokenizer)
61
+
62
+
63
+ def encode_prompt(self, prompt, clip_skip=1, device="cuda", positive=True):
64
+ prompt = self.process_prompt(prompt, positive=positive)
65
+ for keyword in self.keyword_dict:
66
+ if keyword in prompt:
67
+ print(f"Textual inversion {keyword} is enabled.")
68
+ prompt = prompt.replace(keyword, self.keyword_dict[keyword])
69
+ input_ids = tokenize_long_prompt(self.tokenizer, prompt).to(device)
70
+ prompt_emb = self.text_encoder(input_ids, clip_skip=clip_skip)
71
+ prompt_emb = prompt_emb.reshape((1, prompt_emb.shape[0]*prompt_emb.shape[1], -1))
72
+
73
+ return prompt_emb
diffsynth/prompters/sdxl_prompter.py ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter, tokenize_long_prompt
2
+ from ..models.model_manager import ModelManager
3
+ from ..models import SDXLTextEncoder, SDXLTextEncoder2
4
+ from transformers import CLIPTokenizer
5
+ import torch, os
6
+
7
+
8
+
9
+ class SDXLPrompter(BasePrompter):
10
+ def __init__(
11
+ self,
12
+ tokenizer_path=None,
13
+ tokenizer_2_path=None
14
+ ):
15
+ if tokenizer_path is None:
16
+ base_path = os.path.dirname(os.path.dirname(__file__))
17
+ tokenizer_path = os.path.join(base_path, "tokenizer_configs/stable_diffusion/tokenizer")
18
+ if tokenizer_2_path is None:
19
+ base_path = os.path.dirname(os.path.dirname(__file__))
20
+ tokenizer_2_path = os.path.join(base_path, "tokenizer_configs/stable_diffusion_xl/tokenizer_2")
21
+ super().__init__()
22
+ self.tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
23
+ self.tokenizer_2 = CLIPTokenizer.from_pretrained(tokenizer_2_path)
24
+ self.text_encoder: SDXLTextEncoder = None
25
+ self.text_encoder_2: SDXLTextEncoder2 = None
26
+
27
+
28
+ def fetch_models(self, text_encoder: SDXLTextEncoder = None, text_encoder_2: SDXLTextEncoder2 = None):
29
+ self.text_encoder = text_encoder
30
+ self.text_encoder_2 = text_encoder_2
31
+
32
+
33
+ def encode_prompt(
34
+ self,
35
+ prompt,
36
+ clip_skip=1,
37
+ clip_skip_2=2,
38
+ positive=True,
39
+ device="cuda"
40
+ ):
41
+ prompt = self.process_prompt(prompt, positive=positive)
42
+
43
+ # 1
44
+ input_ids = tokenize_long_prompt(self.tokenizer, prompt).to(device)
45
+ prompt_emb_1 = self.text_encoder(input_ids, clip_skip=clip_skip)
46
+
47
+ # 2
48
+ input_ids_2 = tokenize_long_prompt(self.tokenizer_2, prompt).to(device)
49
+ add_text_embeds, prompt_emb_2 = self.text_encoder_2(input_ids_2, clip_skip=clip_skip_2)
50
+
51
+ # Merge
52
+ if prompt_emb_1.shape[0] != prompt_emb_2.shape[0]:
53
+ max_batch_size = min(prompt_emb_1.shape[0], prompt_emb_2.shape[0])
54
+ prompt_emb_1 = prompt_emb_1[: max_batch_size]
55
+ prompt_emb_2 = prompt_emb_2[: max_batch_size]
56
+ prompt_emb = torch.concatenate([prompt_emb_1, prompt_emb_2], dim=-1)
57
+
58
+ # For very long prompt, we only use the first 77 tokens to compute `add_text_embeds`.
59
+ add_text_embeds = add_text_embeds[0:1]
60
+ prompt_emb = prompt_emb.reshape((1, prompt_emb.shape[0]*prompt_emb.shape[1], -1))
61
+ return add_text_embeds, prompt_emb
diffsynth/prompters/stepvideo_prompter.py ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.hunyuan_dit_text_encoder import HunyuanDiTCLIPTextEncoder
3
+ from ..models.stepvideo_text_encoder import STEP1TextEncoder
4
+ from transformers import BertTokenizer
5
+ import os, torch
6
+
7
+
8
+ class StepVideoPrompter(BasePrompter):
9
+
10
+ def __init__(
11
+ self,
12
+ tokenizer_1_path=None,
13
+ ):
14
+ if tokenizer_1_path is None:
15
+ base_path = os.path.dirname(os.path.dirname(__file__))
16
+ tokenizer_1_path = os.path.join(
17
+ base_path, "tokenizer_configs/hunyuan_dit/tokenizer")
18
+ super().__init__()
19
+ self.tokenizer_1 = BertTokenizer.from_pretrained(tokenizer_1_path)
20
+
21
+ def fetch_models(self, text_encoder_1: HunyuanDiTCLIPTextEncoder = None, text_encoder_2: STEP1TextEncoder = None):
22
+ self.text_encoder_1 = text_encoder_1
23
+ self.text_encoder_2 = text_encoder_2
24
+
25
+ def encode_prompt_using_clip(self, prompt, max_length, device):
26
+ text_inputs = self.tokenizer_1(
27
+ prompt,
28
+ padding="max_length",
29
+ max_length=max_length,
30
+ truncation=True,
31
+ return_attention_mask=True,
32
+ return_tensors="pt",
33
+ )
34
+ prompt_embeds = self.text_encoder_1(
35
+ text_inputs.input_ids.to(device),
36
+ attention_mask=text_inputs.attention_mask.to(device),
37
+ )
38
+ return prompt_embeds
39
+
40
+ def encode_prompt_using_llm(self, prompt, max_length, device):
41
+ y, y_mask = self.text_encoder_2(prompt, max_length=max_length, device=device)
42
+ return y, y_mask
43
+
44
+ def encode_prompt(self,
45
+ prompt,
46
+ positive=True,
47
+ device="cuda"):
48
+
49
+ prompt = self.process_prompt(prompt, positive=positive)
50
+
51
+ clip_embeds = self.encode_prompt_using_clip(prompt, max_length=77, device=device)
52
+ llm_embeds, llm_mask = self.encode_prompt_using_llm(prompt, max_length=320, device=device)
53
+
54
+ llm_mask = torch.nn.functional.pad(llm_mask, (clip_embeds.shape[1], 0), value=1)
55
+
56
+ return clip_embeds, llm_embeds, llm_mask
diffsynth/prompters/wan_prompter.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_prompter import BasePrompter
2
+ from ..models.wan_video_text_encoder import WanTextEncoder
3
+ from transformers import AutoTokenizer
4
+ import os, torch
5
+ import ftfy
6
+ import html
7
+ import string
8
+ import regex as re
9
+
10
+
11
+ def basic_clean(text):
12
+ text = ftfy.fix_text(text)
13
+ text = html.unescape(html.unescape(text))
14
+ return text.strip()
15
+
16
+
17
+ def whitespace_clean(text):
18
+ text = re.sub(r'\s+', ' ', text)
19
+ text = text.strip()
20
+ return text
21
+
22
+
23
+ def canonicalize(text, keep_punctuation_exact_string=None):
24
+ text = text.replace('_', ' ')
25
+ if keep_punctuation_exact_string:
26
+ text = keep_punctuation_exact_string.join(
27
+ part.translate(str.maketrans('', '', string.punctuation))
28
+ for part in text.split(keep_punctuation_exact_string))
29
+ else:
30
+ text = text.translate(str.maketrans('', '', string.punctuation))
31
+ text = text.lower()
32
+ text = re.sub(r'\s+', ' ', text)
33
+ return text.strip()
34
+
35
+
36
+ class HuggingfaceTokenizer:
37
+
38
+ def __init__(self, name, seq_len=None, clean=None, **kwargs):
39
+ assert clean in (None, 'whitespace', 'lower', 'canonicalize')
40
+ self.name = name
41
+ self.seq_len = seq_len
42
+ self.clean = clean
43
+
44
+ # init tokenizer
45
+ self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
46
+ self.vocab_size = self.tokenizer.vocab_size
47
+
48
+ def __call__(self, sequence, **kwargs):
49
+ return_mask = kwargs.pop('return_mask', False)
50
+
51
+ # arguments
52
+ _kwargs = {'return_tensors': 'pt'}
53
+ if self.seq_len is not None:
54
+ _kwargs.update({
55
+ 'padding': 'max_length',
56
+ 'truncation': True,
57
+ 'max_length': self.seq_len
58
+ })
59
+ _kwargs.update(**kwargs)
60
+
61
+ # tokenization
62
+ if isinstance(sequence, str):
63
+ sequence = [sequence]
64
+ if self.clean:
65
+ sequence = [self._clean(u) for u in sequence]
66
+ ids = self.tokenizer(sequence, **_kwargs)
67
+
68
+ # output
69
+ if return_mask:
70
+ return ids.input_ids, ids.attention_mask
71
+ else:
72
+ return ids.input_ids
73
+
74
+ def _clean(self, text):
75
+ if self.clean == 'whitespace':
76
+ text = whitespace_clean(basic_clean(text))
77
+ elif self.clean == 'lower':
78
+ text = whitespace_clean(basic_clean(text)).lower()
79
+ elif self.clean == 'canonicalize':
80
+ text = canonicalize(basic_clean(text))
81
+ return text
82
+
83
+
84
+ class WanPrompter(BasePrompter):
85
+
86
+ def __init__(self, tokenizer_path=None, text_len=512):
87
+ super().__init__()
88
+ self.text_len = text_len
89
+ self.text_encoder = None
90
+ self.fetch_tokenizer(tokenizer_path)
91
+
92
+ def fetch_tokenizer(self, tokenizer_path=None):
93
+ if tokenizer_path is not None:
94
+ self.tokenizer = HuggingfaceTokenizer(name=tokenizer_path, seq_len=self.text_len, clean='whitespace')
95
+
96
+ def fetch_models(self, text_encoder: WanTextEncoder = None):
97
+ self.text_encoder = text_encoder
98
+
99
+ def encode_prompt(self, prompt, positive=True, device="cuda"):
100
+ prompt = self.process_prompt(prompt, positive=positive)
101
+
102
+ ids, mask = self.tokenizer(prompt, return_mask=True, add_special_tokens=True)
103
+ ids = ids.to(device)
104
+ mask = mask.to(device)
105
+ seq_lens = mask.gt(0).sum(dim=1).long()
106
+ prompt_emb = self.text_encoder(ids, mask)
107
+ for i, v in enumerate(seq_lens):
108
+ prompt_emb[:, v:] = 0
109
+ return prompt_emb
diffsynth/schedulers/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ from .ddim import EnhancedDDIMScheduler
2
+ from .continuous_ode import ContinuousODEScheduler
3
+ from .flow_match import FlowMatchScheduler
4
+ from .flow_match_plf import FlowMatchPLFScheduler
diffsynth/schedulers/continuous_ode.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+
4
+ class ContinuousODEScheduler():
5
+
6
+ def __init__(self, num_inference_steps=100, sigma_max=700.0, sigma_min=0.002, rho=7.0):
7
+ self.sigma_max = sigma_max
8
+ self.sigma_min = sigma_min
9
+ self.rho = rho
10
+ self.set_timesteps(num_inference_steps)
11
+
12
+
13
+ def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, **kwargs):
14
+ ramp = torch.linspace(1-denoising_strength, 1, num_inference_steps)
15
+ min_inv_rho = torch.pow(torch.tensor((self.sigma_min,)), (1 / self.rho))
16
+ max_inv_rho = torch.pow(torch.tensor((self.sigma_max,)), (1 / self.rho))
17
+ self.sigmas = torch.pow(max_inv_rho + ramp * (min_inv_rho - max_inv_rho), self.rho)
18
+ self.timesteps = torch.log(self.sigmas) * 0.25
19
+
20
+
21
+ def step(self, model_output, timestep, sample, to_final=False):
22
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
23
+ sigma = self.sigmas[timestep_id]
24
+ sample *= (sigma*sigma + 1).sqrt()
25
+ estimated_sample = -sigma / (sigma*sigma + 1).sqrt() * model_output + 1 / (sigma*sigma + 1) * sample
26
+ if to_final or timestep_id + 1 >= len(self.timesteps):
27
+ prev_sample = estimated_sample
28
+ else:
29
+ sigma_ = self.sigmas[timestep_id + 1]
30
+ derivative = 1 / sigma * (sample - estimated_sample)
31
+ prev_sample = sample + derivative * (sigma_ - sigma)
32
+ prev_sample /= (sigma_*sigma_ + 1).sqrt()
33
+ return prev_sample
34
+
35
+
36
+ def return_to_timestep(self, timestep, sample, sample_stablized):
37
+ # This scheduler doesn't support this function.
38
+ pass
39
+
40
+
41
+ def add_noise(self, original_samples, noise, timestep):
42
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
43
+ sigma = self.sigmas[timestep_id]
44
+ sample = (original_samples + noise * sigma) / (sigma*sigma + 1).sqrt()
45
+ return sample
46
+
47
+
48
+ def training_target(self, sample, noise, timestep):
49
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
50
+ sigma = self.sigmas[timestep_id]
51
+ target = (-(sigma*sigma + 1).sqrt() / sigma + 1 / (sigma*sigma + 1).sqrt() / sigma) * sample + 1 / (sigma*sigma + 1).sqrt() * noise
52
+ return target
53
+
54
+
55
+ def training_weight(self, timestep):
56
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
57
+ sigma = self.sigmas[timestep_id]
58
+ weight = (1 + sigma*sigma).sqrt() / sigma
59
+ return weight
diffsynth/schedulers/ddim.py ADDED
@@ -0,0 +1,105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, math
2
+
3
+
4
+ class EnhancedDDIMScheduler():
5
+
6
+ def __init__(self, num_train_timesteps=1000, beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", prediction_type="epsilon", rescale_zero_terminal_snr=False):
7
+ self.num_train_timesteps = num_train_timesteps
8
+ if beta_schedule == "scaled_linear":
9
+ betas = torch.square(torch.linspace(math.sqrt(beta_start), math.sqrt(beta_end), num_train_timesteps, dtype=torch.float32))
10
+ elif beta_schedule == "linear":
11
+ betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32)
12
+ else:
13
+ raise NotImplementedError(f"{beta_schedule} is not implemented")
14
+ self.alphas_cumprod = torch.cumprod(1.0 - betas, dim=0)
15
+ if rescale_zero_terminal_snr:
16
+ self.alphas_cumprod = self.rescale_zero_terminal_snr(self.alphas_cumprod)
17
+ self.alphas_cumprod = self.alphas_cumprod.tolist()
18
+ self.set_timesteps(10)
19
+ self.prediction_type = prediction_type
20
+
21
+
22
+ def rescale_zero_terminal_snr(self, alphas_cumprod):
23
+ alphas_bar_sqrt = alphas_cumprod.sqrt()
24
+
25
+ # Store old values.
26
+ alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
27
+ alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
28
+
29
+ # Shift so the last timestep is zero.
30
+ alphas_bar_sqrt -= alphas_bar_sqrt_T
31
+
32
+ # Scale so the first timestep is back to the old value.
33
+ alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
34
+
35
+ # Convert alphas_bar_sqrt to betas
36
+ alphas_bar = alphas_bar_sqrt.square() # Revert sqrt
37
+
38
+ return alphas_bar
39
+
40
+
41
+ def set_timesteps(self, num_inference_steps, denoising_strength=1.0, **kwargs):
42
+ # The timesteps are aligned to 999...0, which is different from other implementations,
43
+ # but I think this implementation is more reasonable in theory.
44
+ max_timestep = max(round(self.num_train_timesteps * denoising_strength) - 1, 0)
45
+ num_inference_steps = min(num_inference_steps, max_timestep + 1)
46
+ if num_inference_steps == 1:
47
+ self.timesteps = torch.Tensor([max_timestep])
48
+ else:
49
+ step_length = max_timestep / (num_inference_steps - 1)
50
+ self.timesteps = torch.Tensor([round(max_timestep - i*step_length) for i in range(num_inference_steps)])
51
+
52
+
53
+ def denoise(self, model_output, sample, alpha_prod_t, alpha_prod_t_prev):
54
+ if self.prediction_type == "epsilon":
55
+ weight_e = math.sqrt(1 - alpha_prod_t_prev) - math.sqrt(alpha_prod_t_prev * (1 - alpha_prod_t) / alpha_prod_t)
56
+ weight_x = math.sqrt(alpha_prod_t_prev / alpha_prod_t)
57
+ prev_sample = sample * weight_x + model_output * weight_e
58
+ elif self.prediction_type == "v_prediction":
59
+ weight_e = -math.sqrt(alpha_prod_t_prev * (1 - alpha_prod_t)) + math.sqrt(alpha_prod_t * (1 - alpha_prod_t_prev))
60
+ weight_x = math.sqrt(alpha_prod_t * alpha_prod_t_prev) + math.sqrt((1 - alpha_prod_t) * (1 - alpha_prod_t_prev))
61
+ prev_sample = sample * weight_x + model_output * weight_e
62
+ else:
63
+ raise NotImplementedError(f"{self.prediction_type} is not implemented")
64
+ return prev_sample
65
+
66
+
67
+ def step(self, model_output, timestep, sample, to_final=False):
68
+ alpha_prod_t = self.alphas_cumprod[int(timestep.flatten().tolist()[0])]
69
+ if isinstance(timestep, torch.Tensor):
70
+ timestep = timestep.cpu()
71
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
72
+ if to_final or timestep_id + 1 >= len(self.timesteps):
73
+ alpha_prod_t_prev = 1.0
74
+ else:
75
+ timestep_prev = int(self.timesteps[timestep_id + 1])
76
+ alpha_prod_t_prev = self.alphas_cumprod[timestep_prev]
77
+
78
+ return self.denoise(model_output, sample, alpha_prod_t, alpha_prod_t_prev)
79
+
80
+
81
+ def return_to_timestep(self, timestep, sample, sample_stablized):
82
+ alpha_prod_t = self.alphas_cumprod[int(timestep.flatten().tolist()[0])]
83
+ noise_pred = (sample - math.sqrt(alpha_prod_t) * sample_stablized) / math.sqrt(1 - alpha_prod_t)
84
+ return noise_pred
85
+
86
+
87
+ def add_noise(self, original_samples, noise, timestep):
88
+ sqrt_alpha_prod = math.sqrt(self.alphas_cumprod[int(timestep.flatten().tolist()[0])])
89
+ sqrt_one_minus_alpha_prod = math.sqrt(1 - self.alphas_cumprod[int(timestep.flatten().tolist()[0])])
90
+ noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise
91
+ return noisy_samples
92
+
93
+
94
+ def training_target(self, sample, noise, timestep):
95
+ if self.prediction_type == "epsilon":
96
+ return noise
97
+ else:
98
+ sqrt_alpha_prod = math.sqrt(self.alphas_cumprod[int(timestep.flatten().tolist()[0])])
99
+ sqrt_one_minus_alpha_prod = math.sqrt(1 - self.alphas_cumprod[int(timestep.flatten().tolist()[0])])
100
+ target = sqrt_alpha_prod * noise - sqrt_one_minus_alpha_prod * sample
101
+ return target
102
+
103
+
104
+ def training_weight(self, timestep):
105
+ return 1.0
diffsynth/schedulers/flow_match.py ADDED
@@ -0,0 +1,120 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, math
2
+
3
+
4
+
5
+ class FlowMatchScheduler():
6
+
7
+ def __init__(
8
+ self,
9
+ num_inference_steps=100,
10
+ num_train_timesteps=1000,
11
+ shift=3.0,
12
+ sigma_max=1.0,
13
+ sigma_min=0.003/1.002,
14
+ inverse_timesteps=False,
15
+ extra_one_step=False,
16
+ reverse_sigmas=False,
17
+ exponential_shift=False,
18
+ exponential_shift_mu=None,
19
+ shift_terminal=None,
20
+ ):
21
+ self.num_train_timesteps = num_train_timesteps
22
+ self.shift = shift
23
+ self.sigma_max = sigma_max
24
+ self.sigma_min = sigma_min
25
+ self.inverse_timesteps = inverse_timesteps
26
+ self.extra_one_step = extra_one_step
27
+ self.reverse_sigmas = reverse_sigmas
28
+ self.exponential_shift = exponential_shift
29
+ self.exponential_shift_mu = exponential_shift_mu
30
+ self.shift_terminal = shift_terminal
31
+ self.set_timesteps(num_inference_steps)
32
+
33
+
34
+ def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None, dynamic_shift_len=None):
35
+ if shift is not None:
36
+ self.shift = shift
37
+ sigma_start = self.sigma_min + (self.sigma_max - self.sigma_min) * denoising_strength
38
+ if self.extra_one_step:
39
+ self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
40
+ else:
41
+ self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps)
42
+ if self.inverse_timesteps:
43
+ self.sigmas = torch.flip(self.sigmas, dims=[0])
44
+ if self.exponential_shift:
45
+ mu = self.calculate_shift(dynamic_shift_len) if dynamic_shift_len is not None else self.exponential_shift_mu
46
+ self.sigmas = math.exp(mu) / (math.exp(mu) + (1 / self.sigmas - 1))
47
+ else:
48
+ self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
49
+ if self.shift_terminal is not None:
50
+ one_minus_z = 1 - self.sigmas
51
+ scale_factor = one_minus_z[-1] / (1 - self.shift_terminal)
52
+ self.sigmas = 1 - (one_minus_z / scale_factor)
53
+ if self.reverse_sigmas:
54
+ self.sigmas = 1 - self.sigmas
55
+ self.timesteps = self.sigmas * self.num_train_timesteps
56
+ if training:
57
+ x = self.timesteps
58
+ y = torch.exp(-2 * ((x - num_inference_steps / 2) / num_inference_steps) ** 2)
59
+ y_shifted = y - y.min()
60
+ bsmntw_weighing = y_shifted * (num_inference_steps / y_shifted.sum())
61
+ self.linear_timesteps_weights = bsmntw_weighing
62
+ self.training = True
63
+ else:
64
+ self.training = False
65
+
66
+
67
+ def step(self, model_output, timestep, sample, to_final=False, **kwargs):
68
+ if isinstance(timestep, torch.Tensor):
69
+ timestep = timestep.cpu()
70
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
71
+ sigma = self.sigmas[timestep_id]
72
+ if to_final or timestep_id + 1 >= len(self.timesteps):
73
+ sigma_ = 1 if (self.inverse_timesteps or self.reverse_sigmas) else 0
74
+ else:
75
+ sigma_ = self.sigmas[timestep_id + 1]
76
+ prev_sample = sample + model_output * (sigma_ - sigma)
77
+ return prev_sample
78
+
79
+
80
+ def return_to_timestep(self, timestep, sample, sample_stablized):
81
+ if isinstance(timestep, torch.Tensor):
82
+ timestep = timestep.cpu()
83
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
84
+ sigma = self.sigmas[timestep_id]
85
+ model_output = (sample - sample_stablized) / sigma
86
+ return model_output
87
+
88
+
89
+ def add_noise(self, original_samples, noise, timestep):
90
+ if isinstance(timestep, torch.Tensor):
91
+ timestep = timestep.cpu()
92
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
93
+ sigma = self.sigmas[timestep_id]
94
+ sample = (1 - sigma) * original_samples + sigma * noise
95
+ return sample
96
+
97
+
98
+ def training_target(self, sample, noise, timestep):
99
+ target = noise - sample
100
+ return target
101
+
102
+
103
+ def training_weight(self, timestep):
104
+ timestep_id = torch.argmin((self.timesteps - timestep.to(self.timesteps.device)).abs())
105
+ weights = self.linear_timesteps_weights[timestep_id]
106
+ return weights
107
+
108
+
109
+ def calculate_shift(
110
+ self,
111
+ image_seq_len,
112
+ base_seq_len: int = 256,
113
+ max_seq_len: int = 8192,
114
+ base_shift: float = 0.5,
115
+ max_shift: float = 0.9,
116
+ ):
117
+ m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
118
+ b = base_shift - m * base_seq_len
119
+ mu = image_seq_len * m + b
120
+ return mu
diffsynth/schedulers/flow_match_plf.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, math
2
+
3
+
4
+
5
+ class FlowMatchPLFScheduler():
6
+
7
+ def __init__(
8
+ self,
9
+ num_inference_steps=100,
10
+ num_train_timesteps=1000,
11
+ shift=3.0,
12
+ sigma_max=1.0,
13
+ sigma_min=0.003/1.002,
14
+ inverse_timesteps=False,
15
+ extra_one_step=False,
16
+ reverse_sigmas=False,
17
+ exponential_shift=False,
18
+ exponential_shift_mu=None,
19
+ shift_terminal=None,
20
+ ):
21
+ self.num_train_timesteps = num_train_timesteps
22
+ self.shift = shift
23
+ self.sigma_max = sigma_max
24
+ self.sigma_min = sigma_min
25
+ self.inverse_timesteps = inverse_timesteps
26
+ self.extra_one_step = extra_one_step
27
+ self.reverse_sigmas = reverse_sigmas
28
+ self.exponential_shift = exponential_shift
29
+ self.exponential_shift_mu = exponential_shift_mu
30
+ self.shift_terminal = shift_terminal
31
+ self.set_timesteps(num_inference_steps)
32
+
33
+
34
+ def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None, dynamic_shift_len=None):
35
+ if shift is not None:
36
+ self.shift = shift
37
+ sigma_start = self.sigma_min + (self.sigma_max - self.sigma_min) * denoising_strength
38
+ if self.extra_one_step:
39
+ self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
40
+ else:
41
+ self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps)
42
+ if self.inverse_timesteps:
43
+ self.sigmas = torch.flip(self.sigmas, dims=[0])
44
+ if self.exponential_shift:
45
+ mu = self.calculate_shift(dynamic_shift_len) if dynamic_shift_len is not None else self.exponential_shift_mu
46
+ self.sigmas = math.exp(mu) / (math.exp(mu) + (1 / self.sigmas - 1))
47
+ else:
48
+ self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
49
+ if self.shift_terminal is not None:
50
+ one_minus_z = 1 - self.sigmas
51
+ scale_factor = one_minus_z[-1] / (1 - self.shift_terminal)
52
+ self.sigmas = 1 - (one_minus_z / scale_factor)
53
+ if self.reverse_sigmas:
54
+ self.sigmas = 1 - self.sigmas
55
+ self.timesteps = self.sigmas * self.num_train_timesteps
56
+ if training:
57
+ x = self.timesteps
58
+ y = torch.exp(-2 * ((x - num_inference_steps / 2) / num_inference_steps) ** 2)
59
+ y_shifted = y - y.min()
60
+ bsmntw_weighing = y_shifted * (num_inference_steps / y_shifted.sum())
61
+ self.linear_timesteps_weights = bsmntw_weighing
62
+ self.training = True
63
+ else:
64
+ self.training = False
65
+
66
+
67
+ def _get_sigma_bounds(self, timestep, to_final=False):
68
+ if isinstance(timestep, torch.Tensor):
69
+ timestep = timestep.cpu()
70
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
71
+ sigma = self.sigmas[timestep_id]
72
+ if to_final or timestep_id + 1 >= len(self.timesteps):
73
+ sigma_ = 1 if (self.inverse_timesteps or self.reverse_sigmas) else 0
74
+ else:
75
+ sigma_ = self.sigmas[timestep_id + 1]
76
+ return sigma, sigma_, timestep_id
77
+
78
+
79
+ def step(self, model_output, timestep, sample, to_final=False, **kwargs):
80
+ sigma, sigma_, _ = self._get_sigma_bounds(timestep, to_final=to_final)
81
+ prev_sample = sample + model_output * (sigma_ - sigma)
82
+ return prev_sample
83
+
84
+
85
+ def sigma_from_timestep(self, timestep, device=None, dtype=None):
86
+ sigma, _, _ = self._get_sigma_bounds(timestep)
87
+ if isinstance(sigma, torch.Tensor):
88
+ if device is not None or dtype is not None:
89
+ sigma = sigma.to(device=device or sigma.device, dtype=dtype or sigma.dtype)
90
+ return sigma
91
+
92
+
93
+ def step_towards_target(self, model_output, timestep, sample, target_sample, to_final=False, **kwargs):
94
+ sigma, sigma_, _ = self._get_sigma_bounds(timestep, to_final=to_final)
95
+ if isinstance(sigma, torch.Tensor):
96
+ sigma = sigma.to(device=sample.device, dtype=sample.dtype)
97
+ else:
98
+ sigma = torch.tensor(sigma, device=sample.device, dtype=sample.dtype)
99
+ if isinstance(sigma_, torch.Tensor):
100
+ sigma_ = sigma_.to(device=sample.device, dtype=sample.dtype)
101
+ else:
102
+ sigma_ = torch.tensor(sigma_, device=sample.device, dtype=sample.dtype)
103
+ if torch.abs(sigma) < 1e-8:
104
+ return sample
105
+ fusion_vector = (sample - target_sample) / sigma
106
+ prev_sample = sample + fusion_vector * (sigma_ - sigma)
107
+ return prev_sample
108
+
109
+
110
+ def return_to_timestep(self, timestep, sample, sample_stablized):
111
+ if isinstance(timestep, torch.Tensor):
112
+ timestep = timestep.cpu()
113
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
114
+ sigma = self.sigmas[timestep_id]
115
+ model_output = (sample - sample_stablized) / sigma
116
+ return model_output
117
+
118
+
119
+ def add_noise(self, original_samples, noise, timestep):
120
+ if isinstance(timestep, torch.Tensor):
121
+ timestep = timestep.cpu()
122
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
123
+ sigma = self.sigmas[timestep_id]
124
+ sample = (1 - sigma) * original_samples + sigma * noise
125
+ return sample
126
+
127
+
128
+ def training_target(self, sample, noise, timestep):
129
+ target = noise - sample
130
+ return target
131
+
132
+
133
+ def training_weight(self, timestep):
134
+ timestep_id = torch.argmin((self.timesteps - timestep.to(self.timesteps.device)).abs())
135
+ weights = self.linear_timesteps_weights[timestep_id]
136
+ return weights
137
+
138
+
139
+ def calculate_shift(
140
+ self,
141
+ image_seq_len,
142
+ base_seq_len: int = 256,
143
+ max_seq_len: int = 8192,
144
+ base_shift: float = 0.5,
145
+ max_shift: float = 0.9,
146
+ ):
147
+ m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
148
+ b = base_shift - m * base_seq_len
149
+ mu = image_seq_len * m + b
150
+ return mu
diffsynth/tokenizer_configs/__init__.py ADDED
File without changes
diffsynth/tokenizer_configs/cog/tokenizer/added_tokens.json ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "<extra_id_0>": 32099,
3
+ "<extra_id_10>": 32089,
4
+ "<extra_id_11>": 32088,
5
+ "<extra_id_12>": 32087,
6
+ "<extra_id_13>": 32086,
7
+ "<extra_id_14>": 32085,
8
+ "<extra_id_15>": 32084,
9
+ "<extra_id_16>": 32083,
10
+ "<extra_id_17>": 32082,
11
+ "<extra_id_18>": 32081,
12
+ "<extra_id_19>": 32080,
13
+ "<extra_id_1>": 32098,
14
+ "<extra_id_20>": 32079,
15
+ "<extra_id_21>": 32078,
16
+ "<extra_id_22>": 32077,
17
+ "<extra_id_23>": 32076,
18
+ "<extra_id_24>": 32075,
19
+ "<extra_id_25>": 32074,
20
+ "<extra_id_26>": 32073,
21
+ "<extra_id_27>": 32072,
22
+ "<extra_id_28>": 32071,
23
+ "<extra_id_29>": 32070,
24
+ "<extra_id_2>": 32097,
25
+ "<extra_id_30>": 32069,
26
+ "<extra_id_31>": 32068,
27
+ "<extra_id_32>": 32067,
28
+ "<extra_id_33>": 32066,
29
+ "<extra_id_34>": 32065,
30
+ "<extra_id_35>": 32064,
31
+ "<extra_id_36>": 32063,
32
+ "<extra_id_37>": 32062,
33
+ "<extra_id_38>": 32061,
34
+ "<extra_id_39>": 32060,
35
+ "<extra_id_3>": 32096,
36
+ "<extra_id_40>": 32059,
37
+ "<extra_id_41>": 32058,
38
+ "<extra_id_42>": 32057,
39
+ "<extra_id_43>": 32056,
40
+ "<extra_id_44>": 32055,
41
+ "<extra_id_45>": 32054,
42
+ "<extra_id_46>": 32053,
43
+ "<extra_id_47>": 32052,
44
+ "<extra_id_48>": 32051,
45
+ "<extra_id_49>": 32050,
46
+ "<extra_id_4>": 32095,
47
+ "<extra_id_50>": 32049,
48
+ "<extra_id_51>": 32048,
49
+ "<extra_id_52>": 32047,
50
+ "<extra_id_53>": 32046,
51
+ "<extra_id_54>": 32045,
52
+ "<extra_id_55>": 32044,
53
+ "<extra_id_56>": 32043,
54
+ "<extra_id_57>": 32042,
55
+ "<extra_id_58>": 32041,
56
+ "<extra_id_59>": 32040,
57
+ "<extra_id_5>": 32094,
58
+ "<extra_id_60>": 32039,
59
+ "<extra_id_61>": 32038,
60
+ "<extra_id_62>": 32037,
61
+ "<extra_id_63>": 32036,
62
+ "<extra_id_64>": 32035,
63
+ "<extra_id_65>": 32034,
64
+ "<extra_id_66>": 32033,
65
+ "<extra_id_67>": 32032,
66
+ "<extra_id_68>": 32031,
67
+ "<extra_id_69>": 32030,
68
+ "<extra_id_6>": 32093,
69
+ "<extra_id_70>": 32029,
70
+ "<extra_id_71>": 32028,
71
+ "<extra_id_72>": 32027,
72
+ "<extra_id_73>": 32026,
73
+ "<extra_id_74>": 32025,
74
+ "<extra_id_75>": 32024,
75
+ "<extra_id_76>": 32023,
76
+ "<extra_id_77>": 32022,
77
+ "<extra_id_78>": 32021,
78
+ "<extra_id_79>": 32020,
79
+ "<extra_id_7>": 32092,
80
+ "<extra_id_80>": 32019,
81
+ "<extra_id_81>": 32018,
82
+ "<extra_id_82>": 32017,
83
+ "<extra_id_83>": 32016,
84
+ "<extra_id_84>": 32015,
85
+ "<extra_id_85>": 32014,
86
+ "<extra_id_86>": 32013,
87
+ "<extra_id_87>": 32012,
88
+ "<extra_id_88>": 32011,
89
+ "<extra_id_89>": 32010,
90
+ "<extra_id_8>": 32091,
91
+ "<extra_id_90>": 32009,
92
+ "<extra_id_91>": 32008,
93
+ "<extra_id_92>": 32007,
94
+ "<extra_id_93>": 32006,
95
+ "<extra_id_94>": 32005,
96
+ "<extra_id_95>": 32004,
97
+ "<extra_id_96>": 32003,
98
+ "<extra_id_97>": 32002,
99
+ "<extra_id_98>": 32001,
100
+ "<extra_id_99>": 32000,
101
+ "<extra_id_9>": 32090
102
+ }
diffsynth/tokenizer_configs/cog/tokenizer/special_tokens_map.json ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<extra_id_0>",
4
+ "<extra_id_1>",
5
+ "<extra_id_2>",
6
+ "<extra_id_3>",
7
+ "<extra_id_4>",
8
+ "<extra_id_5>",
9
+ "<extra_id_6>",
10
+ "<extra_id_7>",
11
+ "<extra_id_8>",
12
+ "<extra_id_9>",
13
+ "<extra_id_10>",
14
+ "<extra_id_11>",
15
+ "<extra_id_12>",
16
+ "<extra_id_13>",
17
+ "<extra_id_14>",
18
+ "<extra_id_15>",
19
+ "<extra_id_16>",
20
+ "<extra_id_17>",
21
+ "<extra_id_18>",
22
+ "<extra_id_19>",
23
+ "<extra_id_20>",
24
+ "<extra_id_21>",
25
+ "<extra_id_22>",
26
+ "<extra_id_23>",
27
+ "<extra_id_24>",
28
+ "<extra_id_25>",
29
+ "<extra_id_26>",
30
+ "<extra_id_27>",
31
+ "<extra_id_28>",
32
+ "<extra_id_29>",
33
+ "<extra_id_30>",
34
+ "<extra_id_31>",
35
+ "<extra_id_32>",
36
+ "<extra_id_33>",
37
+ "<extra_id_34>",
38
+ "<extra_id_35>",
39
+ "<extra_id_36>",
40
+ "<extra_id_37>",
41
+ "<extra_id_38>",
42
+ "<extra_id_39>",
43
+ "<extra_id_40>",
44
+ "<extra_id_41>",
45
+ "<extra_id_42>",
46
+ "<extra_id_43>",
47
+ "<extra_id_44>",
48
+ "<extra_id_45>",
49
+ "<extra_id_46>",
50
+ "<extra_id_47>",
51
+ "<extra_id_48>",
52
+ "<extra_id_49>",
53
+ "<extra_id_50>",
54
+ "<extra_id_51>",
55
+ "<extra_id_52>",
56
+ "<extra_id_53>",
57
+ "<extra_id_54>",
58
+ "<extra_id_55>",
59
+ "<extra_id_56>",
60
+ "<extra_id_57>",
61
+ "<extra_id_58>",
62
+ "<extra_id_59>",
63
+ "<extra_id_60>",
64
+ "<extra_id_61>",
65
+ "<extra_id_62>",
66
+ "<extra_id_63>",
67
+ "<extra_id_64>",
68
+ "<extra_id_65>",
69
+ "<extra_id_66>",
70
+ "<extra_id_67>",
71
+ "<extra_id_68>",
72
+ "<extra_id_69>",
73
+ "<extra_id_70>",
74
+ "<extra_id_71>",
75
+ "<extra_id_72>",
76
+ "<extra_id_73>",
77
+ "<extra_id_74>",
78
+ "<extra_id_75>",
79
+ "<extra_id_76>",
80
+ "<extra_id_77>",
81
+ "<extra_id_78>",
82
+ "<extra_id_79>",
83
+ "<extra_id_80>",
84
+ "<extra_id_81>",
85
+ "<extra_id_82>",
86
+ "<extra_id_83>",
87
+ "<extra_id_84>",
88
+ "<extra_id_85>",
89
+ "<extra_id_86>",
90
+ "<extra_id_87>",
91
+ "<extra_id_88>",
92
+ "<extra_id_89>",
93
+ "<extra_id_90>",
94
+ "<extra_id_91>",
95
+ "<extra_id_92>",
96
+ "<extra_id_93>",
97
+ "<extra_id_94>",
98
+ "<extra_id_95>",
99
+ "<extra_id_96>",
100
+ "<extra_id_97>",
101
+ "<extra_id_98>",
102
+ "<extra_id_99>"
103
+ ],
104
+ "eos_token": {
105
+ "content": "</s>",
106
+ "lstrip": false,
107
+ "normalized": false,
108
+ "rstrip": false,
109
+ "single_word": false
110
+ },
111
+ "pad_token": {
112
+ "content": "<pad>",
113
+ "lstrip": false,
114
+ "normalized": false,
115
+ "rstrip": false,
116
+ "single_word": false
117
+ },
118
+ "unk_token": {
119
+ "content": "<unk>",
120
+ "lstrip": false,
121
+ "normalized": false,
122
+ "rstrip": false,
123
+ "single_word": false
124
+ }
125
+ }
diffsynth/tokenizer_configs/cog/tokenizer/spiece.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d60acb128cf7b7f2536e8f38a5b18a05535c9e14c7a355904270e15b0945ea86
3
+ size 791656
diffsynth/tokenizer_configs/cog/tokenizer/tokenizer_config.json ADDED
@@ -0,0 +1,940 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": true,
3
+ "added_tokens_decoder": {
4
+ "0": {
5
+ "content": "<pad>",
6
+ "lstrip": false,
7
+ "normalized": false,
8
+ "rstrip": false,
9
+ "single_word": false,
10
+ "special": true
11
+ },
12
+ "1": {
13
+ "content": "</s>",
14
+ "lstrip": false,
15
+ "normalized": false,
16
+ "rstrip": false,
17
+ "single_word": false,
18
+ "special": true
19
+ },
20
+ "2": {
21
+ "content": "<unk>",
22
+ "lstrip": false,
23
+ "normalized": false,
24
+ "rstrip": false,
25
+ "single_word": false,
26
+ "special": true
27
+ },
28
+ "32000": {
29
+ "content": "<extra_id_99>",
30
+ "lstrip": true,
31
+ "normalized": false,
32
+ "rstrip": true,
33
+ "single_word": false,
34
+ "special": true
35
+ },
36
+ "32001": {
37
+ "content": "<extra_id_98>",
38
+ "lstrip": true,
39
+ "normalized": false,
40
+ "rstrip": true,
41
+ "single_word": false,
42
+ "special": true
43
+ },
44
+ "32002": {
45
+ "content": "<extra_id_97>",
46
+ "lstrip": true,
47
+ "normalized": false,
48
+ "rstrip": true,
49
+ "single_word": false,
50
+ "special": true
51
+ },
52
+ "32003": {
53
+ "content": "<extra_id_96>",
54
+ "lstrip": true,
55
+ "normalized": false,
56
+ "rstrip": true,
57
+ "single_word": false,
58
+ "special": true
59
+ },
60
+ "32004": {
61
+ "content": "<extra_id_95>",
62
+ "lstrip": true,
63
+ "normalized": false,
64
+ "rstrip": true,
65
+ "single_word": false,
66
+ "special": true
67
+ },
68
+ "32005": {
69
+ "content": "<extra_id_94>",
70
+ "lstrip": true,
71
+ "normalized": false,
72
+ "rstrip": true,
73
+ "single_word": false,
74
+ "special": true
75
+ },
76
+ "32006": {
77
+ "content": "<extra_id_93>",
78
+ "lstrip": true,
79
+ "normalized": false,
80
+ "rstrip": true,
81
+ "single_word": false,
82
+ "special": true
83
+ },
84
+ "32007": {
85
+ "content": "<extra_id_92>",
86
+ "lstrip": true,
87
+ "normalized": false,
88
+ "rstrip": true,
89
+ "single_word": false,
90
+ "special": true
91
+ },
92
+ "32008": {
93
+ "content": "<extra_id_91>",
94
+ "lstrip": true,
95
+ "normalized": false,
96
+ "rstrip": true,
97
+ "single_word": false,
98
+ "special": true
99
+ },
100
+ "32009": {
101
+ "content": "<extra_id_90>",
102
+ "lstrip": true,
103
+ "normalized": false,
104
+ "rstrip": true,
105
+ "single_word": false,
106
+ "special": true
107
+ },
108
+ "32010": {
109
+ "content": "<extra_id_89>",
110
+ "lstrip": true,
111
+ "normalized": false,
112
+ "rstrip": true,
113
+ "single_word": false,
114
+ "special": true
115
+ },
116
+ "32011": {
117
+ "content": "<extra_id_88>",
118
+ "lstrip": true,
119
+ "normalized": false,
120
+ "rstrip": true,
121
+ "single_word": false,
122
+ "special": true
123
+ },
124
+ "32012": {
125
+ "content": "<extra_id_87>",
126
+ "lstrip": true,
127
+ "normalized": false,
128
+ "rstrip": true,
129
+ "single_word": false,
130
+ "special": true
131
+ },
132
+ "32013": {
133
+ "content": "<extra_id_86>",
134
+ "lstrip": true,
135
+ "normalized": false,
136
+ "rstrip": true,
137
+ "single_word": false,
138
+ "special": true
139
+ },
140
+ "32014": {
141
+ "content": "<extra_id_85>",
142
+ "lstrip": true,
143
+ "normalized": false,
144
+ "rstrip": true,
145
+ "single_word": false,
146
+ "special": true
147
+ },
148
+ "32015": {
149
+ "content": "<extra_id_84>",
150
+ "lstrip": true,
151
+ "normalized": false,
152
+ "rstrip": true,
153
+ "single_word": false,
154
+ "special": true
155
+ },
156
+ "32016": {
157
+ "content": "<extra_id_83>",
158
+ "lstrip": true,
159
+ "normalized": false,
160
+ "rstrip": true,
161
+ "single_word": false,
162
+ "special": true
163
+ },
164
+ "32017": {
165
+ "content": "<extra_id_82>",
166
+ "lstrip": true,
167
+ "normalized": false,
168
+ "rstrip": true,
169
+ "single_word": false,
170
+ "special": true
171
+ },
172
+ "32018": {
173
+ "content": "<extra_id_81>",
174
+ "lstrip": true,
175
+ "normalized": false,
176
+ "rstrip": true,
177
+ "single_word": false,
178
+ "special": true
179
+ },
180
+ "32019": {
181
+ "content": "<extra_id_80>",
182
+ "lstrip": true,
183
+ "normalized": false,
184
+ "rstrip": true,
185
+ "single_word": false,
186
+ "special": true
187
+ },
188
+ "32020": {
189
+ "content": "<extra_id_79>",
190
+ "lstrip": true,
191
+ "normalized": false,
192
+ "rstrip": true,
193
+ "single_word": false,
194
+ "special": true
195
+ },
196
+ "32021": {
197
+ "content": "<extra_id_78>",
198
+ "lstrip": true,
199
+ "normalized": false,
200
+ "rstrip": true,
201
+ "single_word": false,
202
+ "special": true
203
+ },
204
+ "32022": {
205
+ "content": "<extra_id_77>",
206
+ "lstrip": true,
207
+ "normalized": false,
208
+ "rstrip": true,
209
+ "single_word": false,
210
+ "special": true
211
+ },
212
+ "32023": {
213
+ "content": "<extra_id_76>",
214
+ "lstrip": true,
215
+ "normalized": false,
216
+ "rstrip": true,
217
+ "single_word": false,
218
+ "special": true
219
+ },
220
+ "32024": {
221
+ "content": "<extra_id_75>",
222
+ "lstrip": true,
223
+ "normalized": false,
224
+ "rstrip": true,
225
+ "single_word": false,
226
+ "special": true
227
+ },
228
+ "32025": {
229
+ "content": "<extra_id_74>",
230
+ "lstrip": true,
231
+ "normalized": false,
232
+ "rstrip": true,
233
+ "single_word": false,
234
+ "special": true
235
+ },
236
+ "32026": {
237
+ "content": "<extra_id_73>",
238
+ "lstrip": true,
239
+ "normalized": false,
240
+ "rstrip": true,
241
+ "single_word": false,
242
+ "special": true
243
+ },
244
+ "32027": {
245
+ "content": "<extra_id_72>",
246
+ "lstrip": true,
247
+ "normalized": false,
248
+ "rstrip": true,
249
+ "single_word": false,
250
+ "special": true
251
+ },
252
+ "32028": {
253
+ "content": "<extra_id_71>",
254
+ "lstrip": true,
255
+ "normalized": false,
256
+ "rstrip": true,
257
+ "single_word": false,
258
+ "special": true
259
+ },
260
+ "32029": {
261
+ "content": "<extra_id_70>",
262
+ "lstrip": true,
263
+ "normalized": false,
264
+ "rstrip": true,
265
+ "single_word": false,
266
+ "special": true
267
+ },
268
+ "32030": {
269
+ "content": "<extra_id_69>",
270
+ "lstrip": true,
271
+ "normalized": false,
272
+ "rstrip": true,
273
+ "single_word": false,
274
+ "special": true
275
+ },
276
+ "32031": {
277
+ "content": "<extra_id_68>",
278
+ "lstrip": true,
279
+ "normalized": false,
280
+ "rstrip": true,
281
+ "single_word": false,
282
+ "special": true
283
+ },
284
+ "32032": {
285
+ "content": "<extra_id_67>",
286
+ "lstrip": true,
287
+ "normalized": false,
288
+ "rstrip": true,
289
+ "single_word": false,
290
+ "special": true
291
+ },
292
+ "32033": {
293
+ "content": "<extra_id_66>",
294
+ "lstrip": true,
295
+ "normalized": false,
296
+ "rstrip": true,
297
+ "single_word": false,
298
+ "special": true
299
+ },
300
+ "32034": {
301
+ "content": "<extra_id_65>",
302
+ "lstrip": true,
303
+ "normalized": false,
304
+ "rstrip": true,
305
+ "single_word": false,
306
+ "special": true
307
+ },
308
+ "32035": {
309
+ "content": "<extra_id_64>",
310
+ "lstrip": true,
311
+ "normalized": false,
312
+ "rstrip": true,
313
+ "single_word": false,
314
+ "special": true
315
+ },
316
+ "32036": {
317
+ "content": "<extra_id_63>",
318
+ "lstrip": true,
319
+ "normalized": false,
320
+ "rstrip": true,
321
+ "single_word": false,
322
+ "special": true
323
+ },
324
+ "32037": {
325
+ "content": "<extra_id_62>",
326
+ "lstrip": true,
327
+ "normalized": false,
328
+ "rstrip": true,
329
+ "single_word": false,
330
+ "special": true
331
+ },
332
+ "32038": {
333
+ "content": "<extra_id_61>",
334
+ "lstrip": true,
335
+ "normalized": false,
336
+ "rstrip": true,
337
+ "single_word": false,
338
+ "special": true
339
+ },
340
+ "32039": {
341
+ "content": "<extra_id_60>",
342
+ "lstrip": true,
343
+ "normalized": false,
344
+ "rstrip": true,
345
+ "single_word": false,
346
+ "special": true
347
+ },
348
+ "32040": {
349
+ "content": "<extra_id_59>",
350
+ "lstrip": true,
351
+ "normalized": false,
352
+ "rstrip": true,
353
+ "single_word": false,
354
+ "special": true
355
+ },
356
+ "32041": {
357
+ "content": "<extra_id_58>",
358
+ "lstrip": true,
359
+ "normalized": false,
360
+ "rstrip": true,
361
+ "single_word": false,
362
+ "special": true
363
+ },
364
+ "32042": {
365
+ "content": "<extra_id_57>",
366
+ "lstrip": true,
367
+ "normalized": false,
368
+ "rstrip": true,
369
+ "single_word": false,
370
+ "special": true
371
+ },
372
+ "32043": {
373
+ "content": "<extra_id_56>",
374
+ "lstrip": true,
375
+ "normalized": false,
376
+ "rstrip": true,
377
+ "single_word": false,
378
+ "special": true
379
+ },
380
+ "32044": {
381
+ "content": "<extra_id_55>",
382
+ "lstrip": true,
383
+ "normalized": false,
384
+ "rstrip": true,
385
+ "single_word": false,
386
+ "special": true
387
+ },
388
+ "32045": {
389
+ "content": "<extra_id_54>",
390
+ "lstrip": true,
391
+ "normalized": false,
392
+ "rstrip": true,
393
+ "single_word": false,
394
+ "special": true
395
+ },
396
+ "32046": {
397
+ "content": "<extra_id_53>",
398
+ "lstrip": true,
399
+ "normalized": false,
400
+ "rstrip": true,
401
+ "single_word": false,
402
+ "special": true
403
+ },
404
+ "32047": {
405
+ "content": "<extra_id_52>",
406
+ "lstrip": true,
407
+ "normalized": false,
408
+ "rstrip": true,
409
+ "single_word": false,
410
+ "special": true
411
+ },
412
+ "32048": {
413
+ "content": "<extra_id_51>",
414
+ "lstrip": true,
415
+ "normalized": false,
416
+ "rstrip": true,
417
+ "single_word": false,
418
+ "special": true
419
+ },
420
+ "32049": {
421
+ "content": "<extra_id_50>",
422
+ "lstrip": true,
423
+ "normalized": false,
424
+ "rstrip": true,
425
+ "single_word": false,
426
+ "special": true
427
+ },
428
+ "32050": {
429
+ "content": "<extra_id_49>",
430
+ "lstrip": true,
431
+ "normalized": false,
432
+ "rstrip": true,
433
+ "single_word": false,
434
+ "special": true
435
+ },
436
+ "32051": {
437
+ "content": "<extra_id_48>",
438
+ "lstrip": true,
439
+ "normalized": false,
440
+ "rstrip": true,
441
+ "single_word": false,
442
+ "special": true
443
+ },
444
+ "32052": {
445
+ "content": "<extra_id_47>",
446
+ "lstrip": true,
447
+ "normalized": false,
448
+ "rstrip": true,
449
+ "single_word": false,
450
+ "special": true
451
+ },
452
+ "32053": {
453
+ "content": "<extra_id_46>",
454
+ "lstrip": true,
455
+ "normalized": false,
456
+ "rstrip": true,
457
+ "single_word": false,
458
+ "special": true
459
+ },
460
+ "32054": {
461
+ "content": "<extra_id_45>",
462
+ "lstrip": true,
463
+ "normalized": false,
464
+ "rstrip": true,
465
+ "single_word": false,
466
+ "special": true
467
+ },
468
+ "32055": {
469
+ "content": "<extra_id_44>",
470
+ "lstrip": true,
471
+ "normalized": false,
472
+ "rstrip": true,
473
+ "single_word": false,
474
+ "special": true
475
+ },
476
+ "32056": {
477
+ "content": "<extra_id_43>",
478
+ "lstrip": true,
479
+ "normalized": false,
480
+ "rstrip": true,
481
+ "single_word": false,
482
+ "special": true
483
+ },
484
+ "32057": {
485
+ "content": "<extra_id_42>",
486
+ "lstrip": true,
487
+ "normalized": false,
488
+ "rstrip": true,
489
+ "single_word": false,
490
+ "special": true
491
+ },
492
+ "32058": {
493
+ "content": "<extra_id_41>",
494
+ "lstrip": true,
495
+ "normalized": false,
496
+ "rstrip": true,
497
+ "single_word": false,
498
+ "special": true
499
+ },
500
+ "32059": {
501
+ "content": "<extra_id_40>",
502
+ "lstrip": true,
503
+ "normalized": false,
504
+ "rstrip": true,
505
+ "single_word": false,
506
+ "special": true
507
+ },
508
+ "32060": {
509
+ "content": "<extra_id_39>",
510
+ "lstrip": true,
511
+ "normalized": false,
512
+ "rstrip": true,
513
+ "single_word": false,
514
+ "special": true
515
+ },
516
+ "32061": {
517
+ "content": "<extra_id_38>",
518
+ "lstrip": true,
519
+ "normalized": false,
520
+ "rstrip": true,
521
+ "single_word": false,
522
+ "special": true
523
+ },
524
+ "32062": {
525
+ "content": "<extra_id_37>",
526
+ "lstrip": true,
527
+ "normalized": false,
528
+ "rstrip": true,
529
+ "single_word": false,
530
+ "special": true
531
+ },
532
+ "32063": {
533
+ "content": "<extra_id_36>",
534
+ "lstrip": true,
535
+ "normalized": false,
536
+ "rstrip": true,
537
+ "single_word": false,
538
+ "special": true
539
+ },
540
+ "32064": {
541
+ "content": "<extra_id_35>",
542
+ "lstrip": true,
543
+ "normalized": false,
544
+ "rstrip": true,
545
+ "single_word": false,
546
+ "special": true
547
+ },
548
+ "32065": {
549
+ "content": "<extra_id_34>",
550
+ "lstrip": true,
551
+ "normalized": false,
552
+ "rstrip": true,
553
+ "single_word": false,
554
+ "special": true
555
+ },
556
+ "32066": {
557
+ "content": "<extra_id_33>",
558
+ "lstrip": true,
559
+ "normalized": false,
560
+ "rstrip": true,
561
+ "single_word": false,
562
+ "special": true
563
+ },
564
+ "32067": {
565
+ "content": "<extra_id_32>",
566
+ "lstrip": true,
567
+ "normalized": false,
568
+ "rstrip": true,
569
+ "single_word": false,
570
+ "special": true
571
+ },
572
+ "32068": {
573
+ "content": "<extra_id_31>",
574
+ "lstrip": true,
575
+ "normalized": false,
576
+ "rstrip": true,
577
+ "single_word": false,
578
+ "special": true
579
+ },
580
+ "32069": {
581
+ "content": "<extra_id_30>",
582
+ "lstrip": true,
583
+ "normalized": false,
584
+ "rstrip": true,
585
+ "single_word": false,
586
+ "special": true
587
+ },
588
+ "32070": {
589
+ "content": "<extra_id_29>",
590
+ "lstrip": true,
591
+ "normalized": false,
592
+ "rstrip": true,
593
+ "single_word": false,
594
+ "special": true
595
+ },
596
+ "32071": {
597
+ "content": "<extra_id_28>",
598
+ "lstrip": true,
599
+ "normalized": false,
600
+ "rstrip": true,
601
+ "single_word": false,
602
+ "special": true
603
+ },
604
+ "32072": {
605
+ "content": "<extra_id_27>",
606
+ "lstrip": true,
607
+ "normalized": false,
608
+ "rstrip": true,
609
+ "single_word": false,
610
+ "special": true
611
+ },
612
+ "32073": {
613
+ "content": "<extra_id_26>",
614
+ "lstrip": true,
615
+ "normalized": false,
616
+ "rstrip": true,
617
+ "single_word": false,
618
+ "special": true
619
+ },
620
+ "32074": {
621
+ "content": "<extra_id_25>",
622
+ "lstrip": true,
623
+ "normalized": false,
624
+ "rstrip": true,
625
+ "single_word": false,
626
+ "special": true
627
+ },
628
+ "32075": {
629
+ "content": "<extra_id_24>",
630
+ "lstrip": true,
631
+ "normalized": false,
632
+ "rstrip": true,
633
+ "single_word": false,
634
+ "special": true
635
+ },
636
+ "32076": {
637
+ "content": "<extra_id_23>",
638
+ "lstrip": true,
639
+ "normalized": false,
640
+ "rstrip": true,
641
+ "single_word": false,
642
+ "special": true
643
+ },
644
+ "32077": {
645
+ "content": "<extra_id_22>",
646
+ "lstrip": true,
647
+ "normalized": false,
648
+ "rstrip": true,
649
+ "single_word": false,
650
+ "special": true
651
+ },
652
+ "32078": {
653
+ "content": "<extra_id_21>",
654
+ "lstrip": true,
655
+ "normalized": false,
656
+ "rstrip": true,
657
+ "single_word": false,
658
+ "special": true
659
+ },
660
+ "32079": {
661
+ "content": "<extra_id_20>",
662
+ "lstrip": true,
663
+ "normalized": false,
664
+ "rstrip": true,
665
+ "single_word": false,
666
+ "special": true
667
+ },
668
+ "32080": {
669
+ "content": "<extra_id_19>",
670
+ "lstrip": true,
671
+ "normalized": false,
672
+ "rstrip": true,
673
+ "single_word": false,
674
+ "special": true
675
+ },
676
+ "32081": {
677
+ "content": "<extra_id_18>",
678
+ "lstrip": true,
679
+ "normalized": false,
680
+ "rstrip": true,
681
+ "single_word": false,
682
+ "special": true
683
+ },
684
+ "32082": {
685
+ "content": "<extra_id_17>",
686
+ "lstrip": true,
687
+ "normalized": false,
688
+ "rstrip": true,
689
+ "single_word": false,
690
+ "special": true
691
+ },
692
+ "32083": {
693
+ "content": "<extra_id_16>",
694
+ "lstrip": true,
695
+ "normalized": false,
696
+ "rstrip": true,
697
+ "single_word": false,
698
+ "special": true
699
+ },
700
+ "32084": {
701
+ "content": "<extra_id_15>",
702
+ "lstrip": true,
703
+ "normalized": false,
704
+ "rstrip": true,
705
+ "single_word": false,
706
+ "special": true
707
+ },
708
+ "32085": {
709
+ "content": "<extra_id_14>",
710
+ "lstrip": true,
711
+ "normalized": false,
712
+ "rstrip": true,
713
+ "single_word": false,
714
+ "special": true
715
+ },
716
+ "32086": {
717
+ "content": "<extra_id_13>",
718
+ "lstrip": true,
719
+ "normalized": false,
720
+ "rstrip": true,
721
+ "single_word": false,
722
+ "special": true
723
+ },
724
+ "32087": {
725
+ "content": "<extra_id_12>",
726
+ "lstrip": true,
727
+ "normalized": false,
728
+ "rstrip": true,
729
+ "single_word": false,
730
+ "special": true
731
+ },
732
+ "32088": {
733
+ "content": "<extra_id_11>",
734
+ "lstrip": true,
735
+ "normalized": false,
736
+ "rstrip": true,
737
+ "single_word": false,
738
+ "special": true
739
+ },
740
+ "32089": {
741
+ "content": "<extra_id_10>",
742
+ "lstrip": true,
743
+ "normalized": false,
744
+ "rstrip": true,
745
+ "single_word": false,
746
+ "special": true
747
+ },
748
+ "32090": {
749
+ "content": "<extra_id_9>",
750
+ "lstrip": true,
751
+ "normalized": false,
752
+ "rstrip": true,
753
+ "single_word": false,
754
+ "special": true
755
+ },
756
+ "32091": {
757
+ "content": "<extra_id_8>",
758
+ "lstrip": true,
759
+ "normalized": false,
760
+ "rstrip": true,
761
+ "single_word": false,
762
+ "special": true
763
+ },
764
+ "32092": {
765
+ "content": "<extra_id_7>",
766
+ "lstrip": true,
767
+ "normalized": false,
768
+ "rstrip": true,
769
+ "single_word": false,
770
+ "special": true
771
+ },
772
+ "32093": {
773
+ "content": "<extra_id_6>",
774
+ "lstrip": true,
775
+ "normalized": false,
776
+ "rstrip": true,
777
+ "single_word": false,
778
+ "special": true
779
+ },
780
+ "32094": {
781
+ "content": "<extra_id_5>",
782
+ "lstrip": true,
783
+ "normalized": false,
784
+ "rstrip": true,
785
+ "single_word": false,
786
+ "special": true
787
+ },
788
+ "32095": {
789
+ "content": "<extra_id_4>",
790
+ "lstrip": true,
791
+ "normalized": false,
792
+ "rstrip": true,
793
+ "single_word": false,
794
+ "special": true
795
+ },
796
+ "32096": {
797
+ "content": "<extra_id_3>",
798
+ "lstrip": true,
799
+ "normalized": false,
800
+ "rstrip": true,
801
+ "single_word": false,
802
+ "special": true
803
+ },
804
+ "32097": {
805
+ "content": "<extra_id_2>",
806
+ "lstrip": true,
807
+ "normalized": false,
808
+ "rstrip": true,
809
+ "single_word": false,
810
+ "special": true
811
+ },
812
+ "32098": {
813
+ "content": "<extra_id_1>",
814
+ "lstrip": true,
815
+ "normalized": false,
816
+ "rstrip": true,
817
+ "single_word": false,
818
+ "special": true
819
+ },
820
+ "32099": {
821
+ "content": "<extra_id_0>",
822
+ "lstrip": true,
823
+ "normalized": false,
824
+ "rstrip": true,
825
+ "single_word": false,
826
+ "special": true
827
+ }
828
+ },
829
+ "additional_special_tokens": [
830
+ "<extra_id_0>",
831
+ "<extra_id_1>",
832
+ "<extra_id_2>",
833
+ "<extra_id_3>",
834
+ "<extra_id_4>",
835
+ "<extra_id_5>",
836
+ "<extra_id_6>",
837
+ "<extra_id_7>",
838
+ "<extra_id_8>",
839
+ "<extra_id_9>",
840
+ "<extra_id_10>",
841
+ "<extra_id_11>",
842
+ "<extra_id_12>",
843
+ "<extra_id_13>",
844
+ "<extra_id_14>",
845
+ "<extra_id_15>",
846
+ "<extra_id_16>",
847
+ "<extra_id_17>",
848
+ "<extra_id_18>",
849
+ "<extra_id_19>",
850
+ "<extra_id_20>",
851
+ "<extra_id_21>",
852
+ "<extra_id_22>",
853
+ "<extra_id_23>",
854
+ "<extra_id_24>",
855
+ "<extra_id_25>",
856
+ "<extra_id_26>",
857
+ "<extra_id_27>",
858
+ "<extra_id_28>",
859
+ "<extra_id_29>",
860
+ "<extra_id_30>",
861
+ "<extra_id_31>",
862
+ "<extra_id_32>",
863
+ "<extra_id_33>",
864
+ "<extra_id_34>",
865
+ "<extra_id_35>",
866
+ "<extra_id_36>",
867
+ "<extra_id_37>",
868
+ "<extra_id_38>",
869
+ "<extra_id_39>",
870
+ "<extra_id_40>",
871
+ "<extra_id_41>",
872
+ "<extra_id_42>",
873
+ "<extra_id_43>",
874
+ "<extra_id_44>",
875
+ "<extra_id_45>",
876
+ "<extra_id_46>",
877
+ "<extra_id_47>",
878
+ "<extra_id_48>",
879
+ "<extra_id_49>",
880
+ "<extra_id_50>",
881
+ "<extra_id_51>",
882
+ "<extra_id_52>",
883
+ "<extra_id_53>",
884
+ "<extra_id_54>",
885
+ "<extra_id_55>",
886
+ "<extra_id_56>",
887
+ "<extra_id_57>",
888
+ "<extra_id_58>",
889
+ "<extra_id_59>",
890
+ "<extra_id_60>",
891
+ "<extra_id_61>",
892
+ "<extra_id_62>",
893
+ "<extra_id_63>",
894
+ "<extra_id_64>",
895
+ "<extra_id_65>",
896
+ "<extra_id_66>",
897
+ "<extra_id_67>",
898
+ "<extra_id_68>",
899
+ "<extra_id_69>",
900
+ "<extra_id_70>",
901
+ "<extra_id_71>",
902
+ "<extra_id_72>",
903
+ "<extra_id_73>",
904
+ "<extra_id_74>",
905
+ "<extra_id_75>",
906
+ "<extra_id_76>",
907
+ "<extra_id_77>",
908
+ "<extra_id_78>",
909
+ "<extra_id_79>",
910
+ "<extra_id_80>",
911
+ "<extra_id_81>",
912
+ "<extra_id_82>",
913
+ "<extra_id_83>",
914
+ "<extra_id_84>",
915
+ "<extra_id_85>",
916
+ "<extra_id_86>",
917
+ "<extra_id_87>",
918
+ "<extra_id_88>",
919
+ "<extra_id_89>",
920
+ "<extra_id_90>",
921
+ "<extra_id_91>",
922
+ "<extra_id_92>",
923
+ "<extra_id_93>",
924
+ "<extra_id_94>",
925
+ "<extra_id_95>",
926
+ "<extra_id_96>",
927
+ "<extra_id_97>",
928
+ "<extra_id_98>",
929
+ "<extra_id_99>"
930
+ ],
931
+ "clean_up_tokenization_spaces": true,
932
+ "eos_token": "</s>",
933
+ "extra_ids": 100,
934
+ "legacy": true,
935
+ "model_max_length": 226,
936
+ "pad_token": "<pad>",
937
+ "sp_model_kwargs": {},
938
+ "tokenizer_class": "T5Tokenizer",
939
+ "unk_token": "<unk>"
940
+ }
diffsynth/tokenizer_configs/flux/tokenizer_1/merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
diffsynth/tokenizer_configs/flux/tokenizer_1/special_tokens_map.json ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token": {
3
+ "content": "<|startoftext|>",
4
+ "lstrip": false,
5
+ "normalized": true,
6
+ "rstrip": false,
7
+ "single_word": false
8
+ },
9
+ "eos_token": {
10
+ "content": "<|endoftext|>",
11
+ "lstrip": false,
12
+ "normalized": false,
13
+ "rstrip": false,
14
+ "single_word": false
15
+ },
16
+ "pad_token": {
17
+ "content": "<|endoftext|>",
18
+ "lstrip": false,
19
+ "normalized": false,
20
+ "rstrip": false,
21
+ "single_word": false
22
+ },
23
+ "unk_token": {
24
+ "content": "<|endoftext|>",
25
+ "lstrip": false,
26
+ "normalized": false,
27
+ "rstrip": false,
28
+ "single_word": false
29
+ }
30
+ }
diffsynth/tokenizer_configs/flux/tokenizer_1/tokenizer_config.json ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "added_tokens_decoder": {
4
+ "49406": {
5
+ "content": "<|startoftext|>",
6
+ "lstrip": false,
7
+ "normalized": true,
8
+ "rstrip": false,
9
+ "single_word": false,
10
+ "special": true
11
+ },
12
+ "49407": {
13
+ "content": "<|endoftext|>",
14
+ "lstrip": false,
15
+ "normalized": false,
16
+ "rstrip": false,
17
+ "single_word": false,
18
+ "special": true
19
+ }
20
+ },
21
+ "bos_token": "<|startoftext|>",
22
+ "clean_up_tokenization_spaces": true,
23
+ "do_lower_case": true,
24
+ "eos_token": "<|endoftext|>",
25
+ "errors": "replace",
26
+ "model_max_length": 77,
27
+ "pad_token": "<|endoftext|>",
28
+ "tokenizer_class": "CLIPTokenizer",
29
+ "unk_token": "<|endoftext|>"
30
+ }
diffsynth/tokenizer_configs/flux/tokenizer_1/vocab.json ADDED
The diff for this file is too large to render. See raw diff
 
diffsynth/tokenizer_configs/flux/tokenizer_2/special_tokens_map.json ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<extra_id_0>",
4
+ "<extra_id_1>",
5
+ "<extra_id_2>",
6
+ "<extra_id_3>",
7
+ "<extra_id_4>",
8
+ "<extra_id_5>",
9
+ "<extra_id_6>",
10
+ "<extra_id_7>",
11
+ "<extra_id_8>",
12
+ "<extra_id_9>",
13
+ "<extra_id_10>",
14
+ "<extra_id_11>",
15
+ "<extra_id_12>",
16
+ "<extra_id_13>",
17
+ "<extra_id_14>",
18
+ "<extra_id_15>",
19
+ "<extra_id_16>",
20
+ "<extra_id_17>",
21
+ "<extra_id_18>",
22
+ "<extra_id_19>",
23
+ "<extra_id_20>",
24
+ "<extra_id_21>",
25
+ "<extra_id_22>",
26
+ "<extra_id_23>",
27
+ "<extra_id_24>",
28
+ "<extra_id_25>",
29
+ "<extra_id_26>",
30
+ "<extra_id_27>",
31
+ "<extra_id_28>",
32
+ "<extra_id_29>",
33
+ "<extra_id_30>",
34
+ "<extra_id_31>",
35
+ "<extra_id_32>",
36
+ "<extra_id_33>",
37
+ "<extra_id_34>",
38
+ "<extra_id_35>",
39
+ "<extra_id_36>",
40
+ "<extra_id_37>",
41
+ "<extra_id_38>",
42
+ "<extra_id_39>",
43
+ "<extra_id_40>",
44
+ "<extra_id_41>",
45
+ "<extra_id_42>",
46
+ "<extra_id_43>",
47
+ "<extra_id_44>",
48
+ "<extra_id_45>",
49
+ "<extra_id_46>",
50
+ "<extra_id_47>",
51
+ "<extra_id_48>",
52
+ "<extra_id_49>",
53
+ "<extra_id_50>",
54
+ "<extra_id_51>",
55
+ "<extra_id_52>",
56
+ "<extra_id_53>",
57
+ "<extra_id_54>",
58
+ "<extra_id_55>",
59
+ "<extra_id_56>",
60
+ "<extra_id_57>",
61
+ "<extra_id_58>",
62
+ "<extra_id_59>",
63
+ "<extra_id_60>",
64
+ "<extra_id_61>",
65
+ "<extra_id_62>",
66
+ "<extra_id_63>",
67
+ "<extra_id_64>",
68
+ "<extra_id_65>",
69
+ "<extra_id_66>",
70
+ "<extra_id_67>",
71
+ "<extra_id_68>",
72
+ "<extra_id_69>",
73
+ "<extra_id_70>",
74
+ "<extra_id_71>",
75
+ "<extra_id_72>",
76
+ "<extra_id_73>",
77
+ "<extra_id_74>",
78
+ "<extra_id_75>",
79
+ "<extra_id_76>",
80
+ "<extra_id_77>",
81
+ "<extra_id_78>",
82
+ "<extra_id_79>",
83
+ "<extra_id_80>",
84
+ "<extra_id_81>",
85
+ "<extra_id_82>",
86
+ "<extra_id_83>",
87
+ "<extra_id_84>",
88
+ "<extra_id_85>",
89
+ "<extra_id_86>",
90
+ "<extra_id_87>",
91
+ "<extra_id_88>",
92
+ "<extra_id_89>",
93
+ "<extra_id_90>",
94
+ "<extra_id_91>",
95
+ "<extra_id_92>",
96
+ "<extra_id_93>",
97
+ "<extra_id_94>",
98
+ "<extra_id_95>",
99
+ "<extra_id_96>",
100
+ "<extra_id_97>",
101
+ "<extra_id_98>",
102
+ "<extra_id_99>"
103
+ ],
104
+ "eos_token": {
105
+ "content": "</s>",
106
+ "lstrip": false,
107
+ "normalized": false,
108
+ "rstrip": false,
109
+ "single_word": false
110
+ },
111
+ "pad_token": {
112
+ "content": "<pad>",
113
+ "lstrip": false,
114
+ "normalized": false,
115
+ "rstrip": false,
116
+ "single_word": false
117
+ },
118
+ "unk_token": {
119
+ "content": "<unk>",
120
+ "lstrip": false,
121
+ "normalized": false,
122
+ "rstrip": false,
123
+ "single_word": false
124
+ }
125
+ }
diffsynth/tokenizer_configs/flux/tokenizer_2/spiece.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d60acb128cf7b7f2536e8f38a5b18a05535c9e14c7a355904270e15b0945ea86
3
+ size 791656