Spaces:
Runtime error
Runtime error
Migrated files batch 31
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +2 -0
- diffsynth/pipelines/pipeline_runner.py +105 -0
- diffsynth/pipelines/qwen_image.py +364 -0
- diffsynth/pipelines/sd3_image.py +147 -0
- diffsynth/pipelines/sd_image.py +191 -0
- diffsynth/pipelines/sd_video.py +269 -0
- diffsynth/pipelines/sdxl_image.py +226 -0
- diffsynth/pipelines/sdxl_video.py +226 -0
- diffsynth/pipelines/step_video.py +209 -0
- diffsynth/pipelines/svd_video.py +300 -0
- diffsynth/pipelines/wan_video.py +626 -0
- diffsynth/pipelines/wan_video_new.py +1125 -0
- diffsynth/pipelines/wan_video_relit_live.py +525 -0
- diffsynth/processors/FastBlend.py +142 -0
- diffsynth/processors/PILEditor.py +28 -0
- diffsynth/processors/RIFE.py +77 -0
- diffsynth/processors/__init__.py +0 -0
- diffsynth/processors/base.py +6 -0
- diffsynth/processors/sequencial_processor.py +41 -0
- diffsynth/prompters/__init__.py +12 -0
- diffsynth/prompters/base_prompter.py +70 -0
- diffsynth/prompters/cog_prompter.py +46 -0
- diffsynth/prompters/flux_prompter.py +74 -0
- diffsynth/prompters/hunyuan_dit_prompter.py +69 -0
- diffsynth/prompters/hunyuan_video_prompter.py +275 -0
- diffsynth/prompters/kolors_prompter.py +354 -0
- diffsynth/prompters/omnigen_prompter.py +356 -0
- diffsynth/prompters/omost.py +323 -0
- diffsynth/prompters/prompt_refiners.py +130 -0
- diffsynth/prompters/sd3_prompter.py +93 -0
- diffsynth/prompters/sd_prompter.py +73 -0
- diffsynth/prompters/sdxl_prompter.py +61 -0
- diffsynth/prompters/stepvideo_prompter.py +56 -0
- diffsynth/prompters/wan_prompter.py +109 -0
- diffsynth/schedulers/__init__.py +4 -0
- diffsynth/schedulers/continuous_ode.py +59 -0
- diffsynth/schedulers/ddim.py +105 -0
- diffsynth/schedulers/flow_match.py +120 -0
- diffsynth/schedulers/flow_match_plf.py +150 -0
- diffsynth/tokenizer_configs/__init__.py +0 -0
- diffsynth/tokenizer_configs/cog/tokenizer/added_tokens.json +102 -0
- diffsynth/tokenizer_configs/cog/tokenizer/special_tokens_map.json +125 -0
- diffsynth/tokenizer_configs/cog/tokenizer/spiece.model +3 -0
- diffsynth/tokenizer_configs/cog/tokenizer/tokenizer_config.json +940 -0
- diffsynth/tokenizer_configs/flux/tokenizer_1/merges.txt +0 -0
- diffsynth/tokenizer_configs/flux/tokenizer_1/special_tokens_map.json +30 -0
- diffsynth/tokenizer_configs/flux/tokenizer_1/tokenizer_config.json +30 -0
- diffsynth/tokenizer_configs/flux/tokenizer_1/vocab.json +0 -0
- diffsynth/tokenizer_configs/flux/tokenizer_2/special_tokens_map.json +125 -0
- 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
|