Update Riflex && Pipeline callback bug fix && Training Code bug fix (#150)
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
FROM nvidia/cuda:12.1.0-cudnn8-devel-ubuntu22.04
|
||||
ENV DEBIAN_FRONTEND noninteractive
|
||||
|
||||
RUN rm -r /etc/apt/sources.list.d/
|
||||
|
||||
RUN apt-get update -y && apt-get install -y \
|
||||
libgl1 libglib2.0-0 google-perftools \
|
||||
sudo wget git git-lfs vim tig pkg-config libcairo2-dev \
|
||||
aria2 telnet curl net-tools iputils-ping jq \
|
||||
python3-pip python-is-python3 python3.10-venv tzdata lsof zip tmux
|
||||
RUN apt-get update && \
|
||||
apt-get install -y software-properties-common && \
|
||||
add-apt-repository ppa:ubuntuhandbook1/ffmpeg6 && \
|
||||
apt-get update && \
|
||||
apt-get install -y ffmpeg
|
||||
|
||||
RUN pip3 install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/
|
||||
|
||||
# add all extensions
|
||||
RUN pip install wandb tqdm GitPython==3.1.32 Pillow==9.5.0 setuptools --upgrade -i https://mirrors.aliyun.com/pypi/simple/
|
||||
|
||||
RUN pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 --index-url https://download.pytorch.org/whl/cu118
|
||||
RUN pip install xformers==0.0.27.post2 --index-url https://download.pytorch.org/whl/cu118
|
||||
|
||||
# install vllm (video-caption)
|
||||
RUN pip install vllm==0.6.3
|
||||
|
||||
# install requirements (video-caption)
|
||||
WORKDIR /root/
|
||||
COPY easyanimate/video_caption/requirements.txt /root/requirements-video_caption.txt
|
||||
RUN pip install -r /root/requirements-video_caption.txt
|
||||
RUN rm /root/requirements-video_caption.txt
|
||||
|
||||
RUN pip install -U http://eas-data.oss-cn-shanghai.aliyuncs.com/sdk/allspark-0.15-py2.py3-none-any.whl
|
||||
RUN pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers
|
||||
RUN pip install came-pytorch deepspeed pytorch_lightning==1.9.4 func_timeout -i https://mirrors.aliyun.com/pypi/simple/
|
||||
|
||||
# install requirements
|
||||
RUN pip install bitsandbytes mamba-ssm causal-conv1d>=1.4.0 -i https://mirrors.aliyun.com/pypi/simple/
|
||||
RUN pip install ipykernel -i https://mirrors.aliyun.com/pypi/simple/
|
||||
COPY ./requirements.txt /root/requirements.txt
|
||||
RUN pip install -r /root/requirements.txt -i https://mirrors.aliyun.com/pypi/simple/
|
||||
RUN rm -rf /root/requirements.txt
|
||||
|
||||
# install package patches (video-caption)
|
||||
COPY easyanimate/video_caption/package_patches/easyocr_detection_patched.py /usr/local/lib/python3.10/dist-packages/easyocr/detection.py
|
||||
COPY easyanimate/video_caption/package_patches/vila_siglip_encoder_patched.py /usr/local/lib/python3.10/dist-packages/llava/model/multimodal_encoder/siglip_encoder.py
|
||||
|
||||
ENV PYTHONUNBUFFERED 1
|
||||
ENV NVIDIA_DISABLE_REQUIRE 1
|
||||
|
||||
WORKDIR /root/
|
||||
Regular → Executable
+20
-1
@@ -22,7 +22,7 @@ class FunTextBox:
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"multiline": True, "default": "",}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING_PROMPT",)
|
||||
@@ -33,9 +33,26 @@ class FunTextBox:
|
||||
def process(self, prompt):
|
||||
return (prompt, )
|
||||
|
||||
class FunRiflex:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"riflex_k": ("INT", {"default": 6, "min": 0, "max": 10086}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("RIFLEXT_ARGS",)
|
||||
RETURN_NAMES = ("riflex_k",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def process(self, riflex_k):
|
||||
return (riflex_k, )
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FunTextBox": FunTextBox,
|
||||
"FunRiflex": FunRiflex,
|
||||
|
||||
"LoadCogVideoXFunModel": LoadCogVideoXFunModel,
|
||||
"LoadCogVideoXFunLora": LoadCogVideoXFunLora,
|
||||
@@ -62,6 +79,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FunTextBox": "FunTextBox",
|
||||
"FunRiflex": "FunRiflex",
|
||||
|
||||
"LoadCogVideoXFunModel": "Load CogVideoX-Fun Model",
|
||||
"LoadCogVideoXFunLora": "Load CogVideoX-Fun Lora",
|
||||
"CogVideoXFunInpaintSampler": "CogVideoX-Fun Sampler for Image to Video",
|
||||
|
||||
Regular → Executable
+15
-5
@@ -263,7 +263,7 @@ class WanT2VSampler:
|
||||
"STRING_PROMPT",
|
||||
),
|
||||
"video_length": (
|
||||
"INT", {"default": 81, "min": 5, "max": 81, "step": 4}
|
||||
"INT", {"default": 81, "min": 5, "max": 161, "step": 4}
|
||||
),
|
||||
"width": (
|
||||
"INT", {"default": 832, "min": 64, "max": 2048, "step": 16}
|
||||
@@ -310,6 +310,9 @@ class WanT2VSampler:
|
||||
[False, True], {"default": True,}
|
||||
),
|
||||
},
|
||||
"optional":{
|
||||
"riflex_k": ("RIFLEXT_ARGS",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -317,7 +320,7 @@ class WanT2VSampler:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, width, height, is_image, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload):
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, width, height, is_image, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, riflex_k=0):
|
||||
global transformer_cpu_cache
|
||||
global lora_path_before
|
||||
device = mm.get_torch_device()
|
||||
@@ -348,6 +351,9 @@ class WanT2VSampler:
|
||||
with torch.no_grad():
|
||||
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
|
||||
if riflex_k > 0:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
# Apply lora
|
||||
if funmodels.get("lora_cache", False):
|
||||
if len(funmodels.get("loras", [])) != 0:
|
||||
@@ -412,7 +418,7 @@ class WanI2VSampler:
|
||||
"STRING_PROMPT",
|
||||
),
|
||||
"video_length": (
|
||||
"INT", {"default": 81, "min": 5, "max": 81, "step": 4}
|
||||
"INT", {"default": 81, "min": 5, "max": 161, "step": 4}
|
||||
),
|
||||
"base_resolution": (
|
||||
[
|
||||
@@ -455,7 +461,8 @@ class WanI2VSampler:
|
||||
),
|
||||
},
|
||||
"optional":{
|
||||
"start_img": ("IMAGE",)
|
||||
"start_img": ("IMAGE",),
|
||||
"riflex_k": ("RIFLEXT_ARGS",),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -464,7 +471,7 @@ class WanI2VSampler:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, start_img=None, end_img=None):
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, start_img=None, end_img=None, riflex_k=0):
|
||||
global transformer_cpu_cache
|
||||
global lora_path_before
|
||||
device = mm.get_torch_device()
|
||||
@@ -502,6 +509,9 @@ class WanI2VSampler:
|
||||
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_img, end_img, video_length=video_length, sample_size=(height, width))
|
||||
|
||||
if riflex_k > 0:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
# Apply lora
|
||||
if funmodels.get("lora_cache", False):
|
||||
if len(funmodels.get("loras", [])) != 0:
|
||||
|
||||
Regular → Executable
+24
-8
@@ -268,7 +268,7 @@ class WanFunT2VSampler:
|
||||
"STRING_PROMPT",
|
||||
),
|
||||
"video_length": (
|
||||
"INT", {"default": 81, "min": 5, "max": 81, "step": 4}
|
||||
"INT", {"default": 81, "min": 5, "max": 161, "step": 4}
|
||||
),
|
||||
"width": (
|
||||
"INT", {"default": 832, "min": 64, "max": 2048, "step": 16}
|
||||
@@ -315,6 +315,9 @@ class WanFunT2VSampler:
|
||||
[False, True], {"default": True,}
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"riflex_k": ("RIFLEXT_ARGS",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -322,7 +325,7 @@ class WanFunT2VSampler:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, width, height, is_image, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload):
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, width, height, is_image, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, riflex_k=0):
|
||||
global transformer_cpu_cache
|
||||
global lora_path_before
|
||||
device = mm.get_torch_device()
|
||||
@@ -353,6 +356,9 @@ class WanFunT2VSampler:
|
||||
with torch.no_grad():
|
||||
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
|
||||
if riflex_k > 0:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
# Apply lora
|
||||
if funmodels.get("lora_cache", False):
|
||||
if len(funmodels.get("loras", [])) != 0:
|
||||
@@ -434,7 +440,7 @@ class WanFunInpaintSampler:
|
||||
"STRING_PROMPT",
|
||||
),
|
||||
"video_length": (
|
||||
"INT", {"default": 81, "min": 5, "max": 81, "step": 4}
|
||||
"INT", {"default": 81, "min": 5, "max": 161, "step": 4}
|
||||
),
|
||||
"base_resolution": (
|
||||
[
|
||||
@@ -476,9 +482,10 @@ class WanFunInpaintSampler:
|
||||
[False, True], {"default": True,}
|
||||
),
|
||||
},
|
||||
"optional":{
|
||||
"optional": {
|
||||
"start_img": ("IMAGE",),
|
||||
"end_img": ("IMAGE",),
|
||||
"riflex_k": ("RIFLEXT_ARGS",),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -487,7 +494,7 @@ class WanFunInpaintSampler:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, start_img=None, end_img=None):
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, start_img=None, end_img=None, riflex_k=0):
|
||||
global transformer_cpu_cache
|
||||
global lora_path_before
|
||||
device = mm.get_torch_device()
|
||||
@@ -523,6 +530,10 @@ class WanFunInpaintSampler:
|
||||
|
||||
with torch.no_grad():
|
||||
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
|
||||
if riflex_k > 0:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_img, end_img, video_length=video_length, sample_size=(height, width))
|
||||
|
||||
# Apply lora
|
||||
@@ -593,7 +604,7 @@ class WanFunV2VSampler:
|
||||
"STRING_PROMPT",
|
||||
),
|
||||
"video_length": (
|
||||
"INT", {"default": 81, "min": 1, "max": 81, "step": 4}
|
||||
"INT", {"default": 81, "min": 1, "max": 161, "step": 4}
|
||||
),
|
||||
"base_resolution": (
|
||||
[
|
||||
@@ -638,10 +649,11 @@ class WanFunV2VSampler:
|
||||
[False, True], {"default": True,}
|
||||
),
|
||||
},
|
||||
"optional":{
|
||||
"optional": {
|
||||
"validation_video": ("IMAGE",),
|
||||
"control_video": ("IMAGE",),
|
||||
"ref_image": ("IMAGE",),
|
||||
"riflex_k": ("RIFLEXT_ARGS",),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -650,7 +662,7 @@ class WanFunV2VSampler:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, denoise_strength, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, validation_video=None, control_video=None, ref_image=None, ):
|
||||
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, denoise_strength, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, validation_video=None, control_video=None, ref_image=None, riflex_k=0):
|
||||
global transformer_cpu_cache
|
||||
global lora_path_before
|
||||
|
||||
@@ -704,6 +716,10 @@ class WanFunV2VSampler:
|
||||
|
||||
with torch.no_grad():
|
||||
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
|
||||
if riflex_k > 0:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
if model_type == "Inpaint":
|
||||
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=(height, width), fps=16, ref_image=ref_image[0] if ref_image is not None else ref_image)
|
||||
else:
|
||||
|
||||
@@ -49,6 +49,11 @@ if __name__ == "__main__":
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Server ip
|
||||
server_name = "0.0.0.0"
|
||||
server_port = 7860
|
||||
|
||||
@@ -53,6 +53,11 @@ num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Config and model path
|
||||
config_path = "config/wan2.1/wan_civitai.yaml"
|
||||
# model path
|
||||
@@ -195,6 +200,9 @@ with torch.no_grad():
|
||||
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
|
||||
|
||||
if enable_riflex:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
sample = pipeline(
|
||||
|
||||
@@ -41,6 +41,7 @@ GPU_memory_mode = "sequential_cpu_offload"
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
|
||||
# TeaCache config
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
|
||||
# but it may cause slight differences between the generated content and the original content.
|
||||
@@ -51,10 +52,15 @@ num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Config and model path
|
||||
config_path = "config/wan2.1/wan_civitai.yaml"
|
||||
# model path
|
||||
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-14B"
|
||||
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
|
||||
|
||||
# Choose the sampler in "Flow"
|
||||
sampler_name = "Flow"
|
||||
@@ -181,6 +187,9 @@ with torch.no_grad():
|
||||
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
|
||||
|
||||
if enable_riflex:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
sample = pipeline(
|
||||
prompt,
|
||||
num_frames = video_length,
|
||||
|
||||
@@ -49,6 +49,11 @@ if __name__ == "__main__":
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Server ip
|
||||
server_name = "0.0.0.0"
|
||||
server_port = 7860
|
||||
@@ -62,11 +67,11 @@ if __name__ == "__main__":
|
||||
model_type = "Inpaint"
|
||||
|
||||
if ui_mode == "host":
|
||||
demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype)
|
||||
demo, controller = ui_host(GPU_memory_mode, flow_scheduler_dict, model_name, model_type, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype)
|
||||
elif ui_mode == "client":
|
||||
demo, controller = ui_client(flow_scheduler_dict, model_name)
|
||||
else:
|
||||
demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype)
|
||||
demo, controller = ui(GPU_memory_mode, flow_scheduler_dict, config_path, 1, 1, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype)
|
||||
|
||||
def gr_launch():
|
||||
# launch gradio
|
||||
|
||||
Regular → Executable
+9
@@ -42,6 +42,7 @@ GPU_memory_mode = "sequential_cpu_offload"
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
|
||||
# Support TeaCache.
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
|
||||
# but it may cause slight differences between the generated content and the original content.
|
||||
@@ -52,6 +53,11 @@ num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Config and model path
|
||||
config_path = "config/wan2.1/wan_civitai.yaml"
|
||||
# model path
|
||||
@@ -195,6 +201,9 @@ with torch.no_grad():
|
||||
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
|
||||
|
||||
if enable_riflex:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, validation_image_end, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
sample = pipeline(
|
||||
|
||||
@@ -42,6 +42,7 @@ GPU_memory_mode = "sequential_cpu_offload"
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
|
||||
# Support TeaCache.
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
|
||||
# but it may cause slight differences between the generated content and the original content.
|
||||
@@ -52,6 +53,11 @@ num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Config and model path
|
||||
config_path = "config/wan2.1/wan_civitai.yaml"
|
||||
# model path
|
||||
@@ -202,6 +208,9 @@ with torch.no_grad():
|
||||
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
|
||||
|
||||
if enable_riflex:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
if transformer.config.in_channels != vae.config.latent_channels:
|
||||
input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
|
||||
Regular → Executable
+9
@@ -44,6 +44,7 @@ GPU_memory_mode = "sequential_cpu_offload"
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
|
||||
# Support TeaCache.
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
|
||||
# but it may cause slight differences between the generated content and the original content.
|
||||
@@ -54,6 +55,11 @@ num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Config and model path
|
||||
config_path = "config/wan2.1/wan_civitai.yaml"
|
||||
# model path
|
||||
@@ -202,6 +208,9 @@ with torch.no_grad():
|
||||
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
|
||||
|
||||
if enable_riflex:
|
||||
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=ref_image)
|
||||
|
||||
sample = pipeline(
|
||||
|
||||
Regular → Executable
+3
-3
@@ -41,7 +41,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train.py \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -89,7 +89,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -142,7 +142,7 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
|
||||
Regular → Executable
+3
-3
@@ -39,7 +39,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_lora.py \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -83,7 +83,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -132,7 +132,7 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
|
||||
Regular → Executable
+23
-14
@@ -1,4 +1,4 @@
|
||||
## Lora Training Code
|
||||
## Training Code
|
||||
|
||||
We can choose whether to use deep speed in Wan, which can save a lot of video memory.
|
||||
|
||||
@@ -31,7 +31,7 @@ export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
@@ -39,7 +39,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -47,7 +47,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
@@ -60,8 +62,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--train_mode="inpaint" \
|
||||
--low_vram
|
||||
--low_vram \
|
||||
--train_mode="normal" \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Wan T2V with deepspeed zero-2:
|
||||
@@ -76,7 +79,7 @@ export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_lora.py \
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
@@ -84,7 +87,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -92,7 +95,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
@@ -105,9 +110,10 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--train_mode="inpaint" \
|
||||
--low_vram
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Wan T2V with deepspeed zero-3:
|
||||
@@ -127,14 +133,14 @@ export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1/train.py \
|
||||
--config_path="config/wan2.1_fun/wan_civitai.yaml" \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -142,7 +148,9 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
@@ -155,7 +163,8 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--train_mode="inpaint" \
|
||||
--low_vram
|
||||
--trainable_modules "."
|
||||
```
|
||||
@@ -1,7 +1,5 @@
|
||||
## Training Code
|
||||
|
||||
The default training commands for the different versions are as follows:
|
||||
|
||||
We can choose whether to use deep speed in Wan-Fun, which can save a lot of video memory.
|
||||
|
||||
The metadata_control.json is a little different from normal json in Wan-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file.
|
||||
@@ -41,7 +39,7 @@ Some parameters in the sh file can be confusing, and they are explained in this
|
||||
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
|
||||
Wan-Fun without deepspeed:
|
||||
Wan-Fun-Control without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
@@ -83,7 +81,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--train_mode="control_object" \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--trainable_modules "."
|
||||
```
|
||||
@@ -131,7 +129,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_object" \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--trainable_modules "."
|
||||
```
|
||||
@@ -153,14 +151,14 @@ export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1_fun/wan_civitai.yaml" \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -168,7 +166,9 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
@@ -181,7 +181,9 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--train_mode="inpaint" \
|
||||
--low_vram
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--trainable_modules "."
|
||||
```
|
||||
@@ -0,0 +1,180 @@
|
||||
## Training Code
|
||||
|
||||
We can choose whether to use deep speed in Wan-Fun, which can save a lot of video memory.
|
||||
|
||||
The metadata_control.json is a little different from normal json in Wan-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file.
|
||||
|
||||
```json
|
||||
[
|
||||
{
|
||||
"file_path": "train/00000001.mp4",
|
||||
"control_file_path": "control/00000001.mp4",
|
||||
"text": "A group of young men in suits and sunglasses are walking down a city street.",
|
||||
"type": "video"
|
||||
},
|
||||
{
|
||||
"file_path": "train/00000002.jpg",
|
||||
"control_file_path": "control/00000002.jpg",
|
||||
"text": "A group of young men in suits and sunglasses are walking down a city street.",
|
||||
"type": "image"
|
||||
},
|
||||
.....
|
||||
]
|
||||
```
|
||||
|
||||
Some parameters in the sh file can be confusing, and they are explained in this document:
|
||||
|
||||
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution.
|
||||
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
|
||||
- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. For training videos, the height and width will be set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum.
|
||||
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`, and the resolution of video inputs for training is `512x512x49` to `1024x1024x49`.
|
||||
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49`.
|
||||
- `training_with_video_token_length` specifies training the model according to token length. For training images and videos, the height and width will be set to `image_sample_size` as the maximum and `video_sample_size` as the minimum.
|
||||
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=1024`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x49`.
|
||||
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x9`.
|
||||
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
|
||||
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
|
||||
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
|
||||
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
|
||||
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
|
||||
Wan-Fun-Control without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--low_vram
|
||||
```
|
||||
|
||||
Wan-Fun with deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--low_vram
|
||||
```
|
||||
|
||||
Wan T2V with deepspeed zero-3:
|
||||
|
||||
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model:
|
||||
```sh
|
||||
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
|
||||
```
|
||||
|
||||
Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--save_state \
|
||||
--use_deepspeed \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--low_vram
|
||||
```
|
||||
Regular → Executable
+3
-3
@@ -39,7 +39,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -84,7 +84,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
@@ -134,7 +134,7 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=3 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
|
||||
Regular → Executable
+26
-25
@@ -33,13 +33,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000003.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/597d786b-66e1-4610-8ba0-01bd334dccb3" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000003.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/f0b14ac6-4b11-44df-a060-bc86bbd91d1e" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000003.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/411b329f-6d63-4557-820d-80d56ded2005" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -51,13 +51,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000004.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/cb4ce6d1-1b70-480a-b3c1-e7d1c05cc93d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000004.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/017203b3-1091-4ba9-95db-71ab62804b7a" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000004.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/d4c7dc8e-57dc-4d08-aef6-d11a7bdb4972" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -69,13 +69,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000007.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/65bdd0dd-717d-4300-b566-a615ab1f81c2" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000007.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/07dca036-5b72-4d1e-9ed3-dc725a18f654" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000007.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/e97b1158-e772-40f3-a28f-538bca5584e3" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -87,13 +87,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000008.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/6af00186-2ec4-43d7-9360-3b316e5240a8" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000008.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/3f9a5fb2-bfea-469b-8d07-c4289990ee66" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000008.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/db851428-0df7-4ae3-82eb-c85fe902fb96" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
@@ -122,13 +122,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000001.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/7bfded69-54e5-4654-91e9-37aa61d5b5f3" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000001.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/07bb0403-bcd4-439d-a6a8-d97c266305a8" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000001.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/15e13963-3020-4ffe-9b7b-35603435806f" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -140,13 +140,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000002.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/050f3bff-05c9-4931-9112-e9499b136435" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000002.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/9d1e09bd-0acb-4646-a2ef-2c2d582ba150" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000002.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/f2f64228-dd4b-4c8a-9844-acaf44b6c83c" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -158,13 +158,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000005.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/f6b467db-5e0a-4be6-87a0-a1acce692b1b" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000005.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/487b0405-8e99-44bb-ae6d-009442849a94" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000005.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/d581a5a2-1393-4d19-a13e-b669850a0754" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -176,17 +176,18 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
||||
</details>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000008.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/b3228456-1ee9-4f4c-bfd2-03e13df980d7" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000008.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/4d85e13d-4741-437b-991a-195802cb9485" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000008.mp4" width="100%" controls autoplay loop></video>
|
||||
<video src="https://github.com/user-attachments/assets/6dd31da7-4f1a-44b4-bdbb-14fe62e2796b" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
|
||||
> [!NOTE]
|
||||
> The above test prompts are from <a href="https://github.com/KaiyueSun98/T2V-CompBench">T2V-CompBench</a> and expanded into detailed prompts by Llama-3.3.
|
||||
> Videos are generated with HPSv2.1 Reward LoRA weight 0.7 and MPS Reward LoRA weight 0.7.
|
||||
@@ -205,4 +206,4 @@ Set `lora_path` and `lora_weight` in [examples/wan2.1_fun/predict_t2v.py](https
|
||||
<ol>
|
||||
<li id="ref1">Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.</li>
|
||||
<li id="ref2">Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).</li>
|
||||
</ol>
|
||||
</ol>
|
||||
|
||||
+89
-86
@@ -21,6 +21,7 @@ import logging
|
||||
import math
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
|
||||
@@ -48,7 +49,8 @@ from omegaconf import OmegaConf
|
||||
from packaging import version
|
||||
from PIL import Image
|
||||
from torch.distributed.fsdp.fully_sharded_data_parallel import (
|
||||
FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig)
|
||||
FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig,
|
||||
ShardedStateDictConfig)
|
||||
from torch.utils.data import RandomSampler
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from torchvision import transforms
|
||||
@@ -64,16 +66,16 @@ for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
|
||||
from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
|
||||
ASPECT_RATIO_RANDOM_CROP_512,
|
||||
ASPECT_RATIO_RANDOM_CROP_PROB,
|
||||
AspectRatioBatchImageVideoSampler,
|
||||
RandomSampler, get_closest_ratio)
|
||||
ASPECT_RATIO_RANDOM_CROP_512,
|
||||
ASPECT_RATIO_RANDOM_CROP_PROB,
|
||||
AspectRatioBatchImageVideoSampler,
|
||||
RandomSampler, get_closest_ratio)
|
||||
from videox_fun.data.dataset_image_video import (ImageVideoDataset,
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
|
||||
WanTransformer3DModel)
|
||||
from videox_fun.pipeline import WanFunPipeline, WanFunInpaintPipeline
|
||||
WanTransformer3DModel)
|
||||
from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
|
||||
@@ -88,38 +90,6 @@ def filter_kwargs(cls, kwargs):
|
||||
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
|
||||
return filtered_kwargs
|
||||
|
||||
def get_random_downsample_ratio(sample_size, image_ratio=[],
|
||||
all_choices=False, rng=None):
|
||||
def _create_special_list(length):
|
||||
if length == 1:
|
||||
return [1.0]
|
||||
if length >= 2:
|
||||
first_element = 0.75
|
||||
remaining_sum = 1.0 - first_element
|
||||
other_elements_value = remaining_sum / (length - 1)
|
||||
special_list = [first_element] + [other_elements_value] * (length - 1)
|
||||
return special_list
|
||||
|
||||
if sample_size >= 1536:
|
||||
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
|
||||
elif sample_size >= 1024:
|
||||
number_list = [1, 1.25, 1.5, 2] + image_ratio
|
||||
elif sample_size >= 768:
|
||||
number_list = [1, 1.25, 1.5] + image_ratio
|
||||
elif sample_size >= 512:
|
||||
number_list = [1] + image_ratio
|
||||
else:
|
||||
number_list = [1]
|
||||
|
||||
if all_choices:
|
||||
return number_list
|
||||
|
||||
number_list_prob = np.array(_create_special_list(len(number_list)))
|
||||
if rng is None:
|
||||
return np.random.choice(number_list, p = number_list_prob)
|
||||
else:
|
||||
return rng.choice(number_list, p = number_list_prob)
|
||||
|
||||
def resize_mask(mask, latent, process_first_frame_only=True):
|
||||
latent_size = latent.size()
|
||||
batch_size, channels, num_frames, height, width = mask.shape
|
||||
@@ -156,6 +126,19 @@ def resize_mask(mask, latent, process_first_frame_only=True):
|
||||
)
|
||||
return resized_mask
|
||||
|
||||
def linear_decay(initial_value, final_value, total_steps, current_step):
|
||||
if current_step >= total_steps:
|
||||
return final_value
|
||||
current_step = max(0, current_step)
|
||||
step_size = (final_value - initial_value) / total_steps
|
||||
current_value = initial_value + step_size * current_step
|
||||
return current_value
|
||||
|
||||
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.18.0.dev0")
|
||||
|
||||
@@ -274,19 +257,6 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
|
||||
print(f"Eval error with info {e}")
|
||||
return None
|
||||
|
||||
def linear_decay(initial_value, final_value, total_steps, current_step):
|
||||
if current_step >= total_steps:
|
||||
return final_value
|
||||
current_step = max(0, current_step)
|
||||
step_size = (final_value - initial_value) / total_steps
|
||||
current_value = initial_value + step_size * current_step
|
||||
return current_value
|
||||
|
||||
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Simple example of a training script.")
|
||||
parser.add_argument(
|
||||
@@ -1082,6 +1052,14 @@ def main():
|
||||
image_sample_size=args.image_sample_size,
|
||||
enable_bucket=args.enable_bucket, enable_inpaint=True if args.train_mode != "normal" else False,
|
||||
)
|
||||
|
||||
def worker_init_fn(_seed):
|
||||
_seed = _seed * 256
|
||||
def _worker_init_fn(worker_id):
|
||||
print(f"worker_init_fn with {_seed + worker_id}")
|
||||
np.random.seed(_seed + worker_id)
|
||||
random.seed(_seed + worker_id)
|
||||
return _worker_init_fn
|
||||
|
||||
if args.enable_bucket:
|
||||
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
|
||||
@@ -1092,22 +1070,54 @@ def main():
|
||||
aspect_ratios=aspect_ratio_sample_size,
|
||||
)
|
||||
|
||||
def get_length_to_frame_num(token_length):
|
||||
if args.image_sample_size > args.video_sample_size:
|
||||
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
|
||||
|
||||
if sample_sizes[-1] != args.image_sample_size:
|
||||
sample_sizes.append(args.image_sample_size)
|
||||
else:
|
||||
sample_sizes = [args.image_sample_size]
|
||||
|
||||
length_to_frame_num = {
|
||||
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
|
||||
}
|
||||
|
||||
return length_to_frame_num
|
||||
|
||||
def collate_fn(examples):
|
||||
def get_length_to_frame_num(token_length):
|
||||
if args.image_sample_size > args.video_sample_size:
|
||||
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
|
||||
|
||||
if sample_sizes[-1] != args.image_sample_size:
|
||||
sample_sizes.append(args.image_sample_size)
|
||||
else:
|
||||
sample_sizes = [args.image_sample_size]
|
||||
|
||||
length_to_frame_num = {
|
||||
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
|
||||
}
|
||||
|
||||
return length_to_frame_num
|
||||
|
||||
def get_random_downsample_ratio(sample_size, image_ratio=[],
|
||||
all_choices=False, rng=None):
|
||||
def _create_special_list(length):
|
||||
if length == 1:
|
||||
return [1.0]
|
||||
if length >= 2:
|
||||
first_element = 0.90
|
||||
remaining_sum = 1.0 - first_element
|
||||
other_elements_value = remaining_sum / (length - 1)
|
||||
special_list = [first_element] + [other_elements_value] * (length - 1)
|
||||
return special_list
|
||||
|
||||
if sample_size >= 1536:
|
||||
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
|
||||
elif sample_size >= 1024:
|
||||
number_list = [1, 1.25, 1.5, 2] + image_ratio
|
||||
elif sample_size >= 768:
|
||||
number_list = [1, 1.25, 1.5] + image_ratio
|
||||
elif sample_size >= 512:
|
||||
number_list = [1] + image_ratio
|
||||
else:
|
||||
number_list = [1]
|
||||
|
||||
if all_choices:
|
||||
return number_list
|
||||
|
||||
number_list_prob = np.array(_create_special_list(len(number_list)))
|
||||
if rng is None:
|
||||
return np.random.choice(number_list, p = number_list_prob)
|
||||
else:
|
||||
return rng.choice(number_list, p = number_list_prob)
|
||||
|
||||
# Get token length
|
||||
target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size
|
||||
length_to_frame_num = get_length_to_frame_num(target_token_length)
|
||||
@@ -1128,7 +1138,7 @@ def main():
|
||||
data_type = examples[0]["data_type"]
|
||||
f, h, w, c = np.shape(pixel_value)
|
||||
if data_type == 'image':
|
||||
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size], rng=rng)
|
||||
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size])
|
||||
|
||||
aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
|
||||
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
|
||||
@@ -1142,15 +1152,11 @@ def main():
|
||||
choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25]
|
||||
if len(choice_list) == 0:
|
||||
choice_list = list(length_to_frame_num.keys())
|
||||
if rng is None:
|
||||
local_video_sample_size = np.random.choice(choice_list)
|
||||
else:
|
||||
local_video_sample_size = rng.choice(choice_list)
|
||||
local_video_sample_size = np.random.choice(choice_list)
|
||||
batch_video_length = length_to_frame_num[local_video_sample_size]
|
||||
random_downsample_ratio = args.video_sample_size / local_video_sample_size
|
||||
else:
|
||||
random_downsample_ratio = get_random_downsample_ratio(
|
||||
args.video_sample_size, rng=rng)
|
||||
random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size)
|
||||
batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
|
||||
else:
|
||||
random_downsample_ratio = 1
|
||||
@@ -1162,14 +1168,9 @@ def main():
|
||||
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
|
||||
closest_size = [int(x / 16) * 16 for x in closest_size]
|
||||
if args.random_ratio_crop:
|
||||
if rng is None:
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
else:
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
|
||||
|
||||
for example in examples:
|
||||
@@ -1265,6 +1266,7 @@ def main():
|
||||
collate_fn=collate_fn,
|
||||
persistent_workers=True if args.dataloader_num_workers != 0 else False,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
|
||||
)
|
||||
else:
|
||||
# DataLoaders creation:
|
||||
@@ -1275,6 +1277,7 @@ def main():
|
||||
batch_sampler=batch_sampler,
|
||||
persistent_workers=True if args.dataloader_num_workers != 0 else False,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
|
||||
)
|
||||
|
||||
# Scheduler and math around the number of training steps.
|
||||
@@ -1420,7 +1423,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and not zero_stage == 3:
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1441,7 +1444,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and not zero_stage == 3:
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
@@ -21,6 +21,7 @@ import logging
|
||||
import math
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
|
||||
@@ -47,8 +48,6 @@ from einops import rearrange
|
||||
from omegaconf import OmegaConf
|
||||
from packaging import version
|
||||
from PIL import Image
|
||||
from torch.distributed.fsdp.fully_sharded_data_parallel import (
|
||||
FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig)
|
||||
from torch.utils.data import RandomSampler
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from torchvision import transforms
|
||||
@@ -72,10 +71,11 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset,
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
|
||||
WanTransformer3DModel)
|
||||
from videox_fun.pipeline import WanFunControlPipeline, WanFunInpaintPipeline
|
||||
WanTransformer3DModel)
|
||||
from videox_fun.pipeline import WanFunControlPipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.utils import get_video_to_video_latent, save_videos_grid
|
||||
from videox_fun.utils.utils import (get_video_to_video_latent,
|
||||
save_videos_grid)
|
||||
|
||||
if is_wandb_available():
|
||||
import wandb
|
||||
@@ -88,73 +88,18 @@ def filter_kwargs(cls, kwargs):
|
||||
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
|
||||
return filtered_kwargs
|
||||
|
||||
def get_random_downsample_ratio(sample_size, image_ratio=[],
|
||||
all_choices=False, rng=None):
|
||||
def _create_special_list(length):
|
||||
if length == 1:
|
||||
return [1.0]
|
||||
if length >= 2:
|
||||
first_element = 0.75
|
||||
remaining_sum = 1.0 - first_element
|
||||
other_elements_value = remaining_sum / (length - 1)
|
||||
special_list = [first_element] + [other_elements_value] * (length - 1)
|
||||
return special_list
|
||||
|
||||
if sample_size >= 1536:
|
||||
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
|
||||
elif sample_size >= 1024:
|
||||
number_list = [1, 1.25, 1.5, 2] + image_ratio
|
||||
elif sample_size >= 768:
|
||||
number_list = [1, 1.25, 1.5] + image_ratio
|
||||
elif sample_size >= 512:
|
||||
number_list = [1] + image_ratio
|
||||
else:
|
||||
number_list = [1]
|
||||
def linear_decay(initial_value, final_value, total_steps, current_step):
|
||||
if current_step >= total_steps:
|
||||
return final_value
|
||||
current_step = max(0, current_step)
|
||||
step_size = (final_value - initial_value) / total_steps
|
||||
current_value = initial_value + step_size * current_step
|
||||
return current_value
|
||||
|
||||
if all_choices:
|
||||
return number_list
|
||||
|
||||
number_list_prob = np.array(_create_special_list(len(number_list)))
|
||||
if rng is None:
|
||||
return np.random.choice(number_list, p = number_list_prob)
|
||||
else:
|
||||
return rng.choice(number_list, p = number_list_prob)
|
||||
|
||||
def resize_mask(mask, latent, process_first_frame_only=True):
|
||||
latent_size = latent.size()
|
||||
batch_size, channels, num_frames, height, width = mask.shape
|
||||
|
||||
if process_first_frame_only:
|
||||
target_size = list(latent_size[2:])
|
||||
target_size[0] = 1
|
||||
first_frame_resized = F.interpolate(
|
||||
mask[:, :, 0:1, :, :],
|
||||
size=target_size,
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
target_size = list(latent_size[2:])
|
||||
target_size[0] = target_size[0] - 1
|
||||
if target_size[0] != 0:
|
||||
remaining_frames_resized = F.interpolate(
|
||||
mask[:, :, 1:, :, :],
|
||||
size=target_size,
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
)
|
||||
resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2)
|
||||
else:
|
||||
resized_mask = first_frame_resized
|
||||
else:
|
||||
target_size = list(latent_size[2:])
|
||||
resized_mask = F.interpolate(
|
||||
mask,
|
||||
size=target_size,
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
)
|
||||
return resized_mask
|
||||
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.18.0.dev0")
|
||||
@@ -222,19 +167,6 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
|
||||
print(f"Eval error with info {e}")
|
||||
return None
|
||||
|
||||
def linear_decay(initial_value, final_value, total_steps, current_step):
|
||||
if current_step >= total_steps:
|
||||
return final_value
|
||||
current_step = max(0, current_step)
|
||||
step_size = (final_value - initial_value) / total_steps
|
||||
current_value = initial_value + step_size * current_step
|
||||
return current_value
|
||||
|
||||
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Simple example of a training script.")
|
||||
parser.add_argument(
|
||||
@@ -1040,6 +972,14 @@ def main():
|
||||
image_sample_size=args.image_sample_size,
|
||||
enable_bucket=args.enable_bucket, enable_inpaint=False,
|
||||
)
|
||||
|
||||
def worker_init_fn(_seed):
|
||||
_seed = _seed * 256
|
||||
def _worker_init_fn(worker_id):
|
||||
print(f"worker_init_fn with {_seed + worker_id}")
|
||||
np.random.seed(_seed + worker_id)
|
||||
random.seed(_seed + worker_id)
|
||||
return _worker_init_fn
|
||||
|
||||
if args.enable_bucket:
|
||||
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
|
||||
@@ -1050,22 +990,54 @@ def main():
|
||||
aspect_ratios=aspect_ratio_sample_size,
|
||||
)
|
||||
|
||||
def get_length_to_frame_num(token_length):
|
||||
if args.image_sample_size > args.video_sample_size:
|
||||
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
|
||||
|
||||
if sample_sizes[-1] != args.image_sample_size:
|
||||
sample_sizes.append(args.image_sample_size)
|
||||
else:
|
||||
sample_sizes = [args.image_sample_size]
|
||||
|
||||
length_to_frame_num = {
|
||||
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
|
||||
}
|
||||
|
||||
return length_to_frame_num
|
||||
|
||||
def collate_fn(examples):
|
||||
def get_length_to_frame_num(token_length):
|
||||
if args.image_sample_size > args.video_sample_size:
|
||||
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
|
||||
|
||||
if sample_sizes[-1] != args.image_sample_size:
|
||||
sample_sizes.append(args.image_sample_size)
|
||||
else:
|
||||
sample_sizes = [args.image_sample_size]
|
||||
|
||||
length_to_frame_num = {
|
||||
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
|
||||
}
|
||||
|
||||
return length_to_frame_num
|
||||
|
||||
def get_random_downsample_ratio(sample_size, image_ratio=[],
|
||||
all_choices=False, rng=None):
|
||||
def _create_special_list(length):
|
||||
if length == 1:
|
||||
return [1.0]
|
||||
if length >= 2:
|
||||
first_element = 0.90
|
||||
remaining_sum = 1.0 - first_element
|
||||
other_elements_value = remaining_sum / (length - 1)
|
||||
special_list = [first_element] + [other_elements_value] * (length - 1)
|
||||
return special_list
|
||||
|
||||
if sample_size >= 1536:
|
||||
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
|
||||
elif sample_size >= 1024:
|
||||
number_list = [1, 1.25, 1.5, 2] + image_ratio
|
||||
elif sample_size >= 768:
|
||||
number_list = [1, 1.25, 1.5] + image_ratio
|
||||
elif sample_size >= 512:
|
||||
number_list = [1] + image_ratio
|
||||
else:
|
||||
number_list = [1]
|
||||
|
||||
if all_choices:
|
||||
return number_list
|
||||
|
||||
number_list_prob = np.array(_create_special_list(len(number_list)))
|
||||
if rng is None:
|
||||
return np.random.choice(number_list, p = number_list_prob)
|
||||
else:
|
||||
return rng.choice(number_list, p = number_list_prob)
|
||||
|
||||
# Get token length
|
||||
target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size
|
||||
length_to_frame_num = get_length_to_frame_num(target_token_length)
|
||||
@@ -1087,7 +1059,7 @@ def main():
|
||||
data_type = examples[0]["data_type"]
|
||||
f, h, w, c = np.shape(pixel_value)
|
||||
if data_type == 'image':
|
||||
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size], rng=rng)
|
||||
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size])
|
||||
|
||||
aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
|
||||
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
|
||||
@@ -1101,15 +1073,11 @@ def main():
|
||||
choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25]
|
||||
if len(choice_list) == 0:
|
||||
choice_list = list(length_to_frame_num.keys())
|
||||
if rng is None:
|
||||
local_video_sample_size = np.random.choice(choice_list)
|
||||
else:
|
||||
local_video_sample_size = rng.choice(choice_list)
|
||||
local_video_sample_size = np.random.choice(choice_list)
|
||||
batch_video_length = length_to_frame_num[local_video_sample_size]
|
||||
random_downsample_ratio = args.video_sample_size / local_video_sample_size
|
||||
else:
|
||||
random_downsample_ratio = get_random_downsample_ratio(
|
||||
args.video_sample_size, rng=rng)
|
||||
random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size)
|
||||
batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
|
||||
else:
|
||||
random_downsample_ratio = 1
|
||||
@@ -1121,14 +1089,9 @@ def main():
|
||||
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
|
||||
closest_size = [int(x / 16) * 16 for x in closest_size]
|
||||
if args.random_ratio_crop:
|
||||
if rng is None:
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
else:
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
|
||||
|
||||
for example in examples:
|
||||
@@ -1239,6 +1202,7 @@ def main():
|
||||
collate_fn=collate_fn,
|
||||
persistent_workers=True if args.dataloader_num_workers != 0 else False,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
|
||||
)
|
||||
else:
|
||||
# DataLoaders creation:
|
||||
@@ -1249,6 +1213,7 @@ def main():
|
||||
batch_sampler=batch_sampler,
|
||||
persistent_workers=True if args.dataloader_num_workers != 0 else False,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
|
||||
)
|
||||
|
||||
# Scheduler and math around the number of training steps.
|
||||
@@ -1400,7 +1365,7 @@ def main():
|
||||
control_pixel_values = batch["control_pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and not zero_stage == 3:
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1))
|
||||
@@ -1422,7 +1387,7 @@ def main():
|
||||
ref_pixel_values = batch["ref_pixel_values"].to(weight_dtype)
|
||||
clip_pixel_values = batch["clip_pixel_values"]
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length and not zero_stage == 3:
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,39 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
export NCCL_IB_DISABLE=1
|
||||
export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--image_sample_size=1024 \
|
||||
--video_sample_size=256 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--video_repeat=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--train_mode="control_ref" \
|
||||
--control_ref_image="first_frame" \
|
||||
--low_vram
|
||||
+117
-109
@@ -21,6 +21,7 @@ import logging
|
||||
import math
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
|
||||
@@ -61,18 +62,19 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
|
||||
for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
|
||||
ASPECT_RATIO_RANDOM_CROP_512,
|
||||
ASPECT_RATIO_RANDOM_CROP_PROB,
|
||||
AspectRatioBatchImageVideoSampler,
|
||||
RandomSampler, get_closest_ratio)
|
||||
ASPECT_RATIO_RANDOM_CROP_512,
|
||||
ASPECT_RATIO_RANDOM_CROP_PROB,
|
||||
AspectRatioBatchImageVideoSampler,
|
||||
RandomSampler, get_closest_ratio)
|
||||
from videox_fun.data.dataset_image_video import (ImageVideoDataset,
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
|
||||
WanTransformer3DModel)
|
||||
WanTransformer3DModel)
|
||||
from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.lora_utils import create_network, merge_lora, unmerge_lora
|
||||
from videox_fun.utils.lora_utils import (create_network, merge_lora,
|
||||
unmerge_lora)
|
||||
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
|
||||
if is_wandb_available():
|
||||
@@ -85,38 +87,6 @@ def filter_kwargs(cls, kwargs):
|
||||
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
|
||||
return filtered_kwargs
|
||||
|
||||
def get_random_downsample_ratio(sample_size, image_ratio=[],
|
||||
all_choices=False, rng=None):
|
||||
def _create_special_list(length):
|
||||
if length == 1:
|
||||
return [1.0]
|
||||
if length >= 2:
|
||||
first_element = 0.75
|
||||
remaining_sum = 1.0 - first_element
|
||||
other_elements_value = remaining_sum / (length - 1)
|
||||
special_list = [first_element] + [other_elements_value] * (length - 1)
|
||||
return special_list
|
||||
|
||||
if sample_size >= 1536:
|
||||
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
|
||||
elif sample_size >= 1024:
|
||||
number_list = [1, 1.25, 1.5, 2] + image_ratio
|
||||
elif sample_size >= 768:
|
||||
number_list = [1, 1.25, 1.5] + image_ratio
|
||||
elif sample_size >= 512:
|
||||
number_list = [1] + image_ratio
|
||||
else:
|
||||
number_list = [1]
|
||||
|
||||
if all_choices:
|
||||
return number_list
|
||||
|
||||
number_list_prob = np.array(_create_special_list(len(number_list)))
|
||||
if rng is None:
|
||||
return np.random.choice(number_list, p = number_list_prob)
|
||||
else:
|
||||
return rng.choice(number_list, p = number_list_prob)
|
||||
|
||||
def resize_mask(mask, latent, process_first_frame_only=True):
|
||||
latent_size = latent.size()
|
||||
batch_size, channels, num_frames, height, width = mask.shape
|
||||
@@ -153,6 +123,19 @@ def resize_mask(mask, latent, process_first_frame_only=True):
|
||||
)
|
||||
return resized_mask
|
||||
|
||||
def linear_decay(initial_value, final_value, total_steps, current_step):
|
||||
if current_step >= total_steps:
|
||||
return final_value
|
||||
current_step = max(0, current_step)
|
||||
step_size = (final_value - initial_value) / total_steps
|
||||
current_value = initial_value + step_size * current_step
|
||||
return current_value
|
||||
|
||||
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.18.0.dev0")
|
||||
|
||||
@@ -272,19 +255,6 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
|
||||
print(f"Eval error with info {e}")
|
||||
return None
|
||||
|
||||
def linear_decay(initial_value, final_value, total_steps, current_step):
|
||||
if current_step >= total_steps:
|
||||
return final_value
|
||||
current_step = max(0, current_step)
|
||||
step_size = (final_value - initial_value) / total_steps
|
||||
current_value = initial_value + step_size * current_step
|
||||
return current_value
|
||||
|
||||
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Simple example of a training script.")
|
||||
parser.add_argument(
|
||||
@@ -585,29 +555,23 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--noise_share_in_frames_ratio", type=float, default=0.5, help="Noise share ratio.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss", action="store_true", help="Whether enable motion sub loss."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--motion_sub_loss_ratio", type=float, default=0.25, help="The ratio of motion sub loss."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--keep_all_node_same_token_length",
|
||||
action="store_true",
|
||||
help="Reference of the length token.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_sampling_steps",
|
||||
type=int,
|
||||
default=1000,
|
||||
help="Run train_sampling_steps.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--keep_all_node_same_token_length",
|
||||
action="store_true",
|
||||
help="Reference of the length token.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--token_sample_size",
|
||||
type=int,
|
||||
@@ -652,12 +616,6 @@ def parse_args():
|
||||
"The config of the model in training."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image_repeat_in_forward",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Num of repeat image in forward.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--transformer_path",
|
||||
type=str,
|
||||
@@ -675,7 +633,7 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
'--tokenizer_max_length',
|
||||
type=int,
|
||||
default=226,
|
||||
default=512,
|
||||
help='Max length of tokenizer'
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -712,6 +670,12 @@ def parse_args():
|
||||
default=1.29,
|
||||
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora_skip_name",
|
||||
type=str,
|
||||
default=None,
|
||||
help=("The module is not trained in loras. "),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
|
||||
@@ -761,6 +725,12 @@ def main():
|
||||
else:
|
||||
zero_stage = 0
|
||||
print("DeepSpeed is not enabled.")
|
||||
if zero_stage == 3:
|
||||
accelerator_transformer3d = Accelerator(
|
||||
gradient_accumulation_steps=args.gradient_accumulation_steps,
|
||||
mixed_precision=args.mixed_precision,
|
||||
project_config=accelerator_project_config,
|
||||
)
|
||||
if accelerator.is_main_process:
|
||||
writer = SummaryWriter(log_dir=logging_dir)
|
||||
|
||||
@@ -843,6 +813,7 @@ def main():
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=weight_dtype,
|
||||
)
|
||||
text_encoder = text_encoder.eval()
|
||||
# Get Vae
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
|
||||
@@ -877,7 +848,7 @@ def main():
|
||||
text_encoder,
|
||||
transformer3d,
|
||||
neuron_dropout=None,
|
||||
add_lora_in_attn_temporal=True,
|
||||
skip_name=args.lora_skip_name,
|
||||
)
|
||||
network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True)
|
||||
|
||||
@@ -1012,6 +983,14 @@ def main():
|
||||
image_sample_size=args.image_sample_size,
|
||||
enable_bucket=args.enable_bucket, enable_inpaint=True if args.train_mode != "normal" else False,
|
||||
)
|
||||
|
||||
def worker_init_fn(_seed):
|
||||
_seed = _seed * 256
|
||||
def _worker_init_fn(worker_id):
|
||||
print(f"worker_init_fn with {_seed + worker_id}")
|
||||
np.random.seed(_seed + worker_id)
|
||||
random.seed(_seed + worker_id)
|
||||
return _worker_init_fn
|
||||
|
||||
if args.enable_bucket:
|
||||
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
|
||||
@@ -1022,22 +1001,54 @@ def main():
|
||||
aspect_ratios=aspect_ratio_sample_size,
|
||||
)
|
||||
|
||||
def get_length_to_frame_num(token_length):
|
||||
if args.image_sample_size > args.video_sample_size:
|
||||
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
|
||||
|
||||
if sample_sizes[-1] != args.image_sample_size:
|
||||
sample_sizes.append(args.image_sample_size)
|
||||
else:
|
||||
sample_sizes = [args.image_sample_size]
|
||||
|
||||
length_to_frame_num = {
|
||||
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
|
||||
}
|
||||
|
||||
return length_to_frame_num
|
||||
|
||||
def collate_fn(examples):
|
||||
def get_length_to_frame_num(token_length):
|
||||
if args.image_sample_size > args.video_sample_size:
|
||||
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
|
||||
|
||||
if sample_sizes[-1] != args.image_sample_size:
|
||||
sample_sizes.append(args.image_sample_size)
|
||||
else:
|
||||
sample_sizes = [args.image_sample_size]
|
||||
|
||||
length_to_frame_num = {
|
||||
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
|
||||
}
|
||||
|
||||
return length_to_frame_num
|
||||
|
||||
def get_random_downsample_ratio(sample_size, image_ratio=[],
|
||||
all_choices=False, rng=None):
|
||||
def _create_special_list(length):
|
||||
if length == 1:
|
||||
return [1.0]
|
||||
if length >= 2:
|
||||
first_element = 0.90
|
||||
remaining_sum = 1.0 - first_element
|
||||
other_elements_value = remaining_sum / (length - 1)
|
||||
special_list = [first_element] + [other_elements_value] * (length - 1)
|
||||
return special_list
|
||||
|
||||
if sample_size >= 1536:
|
||||
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
|
||||
elif sample_size >= 1024:
|
||||
number_list = [1, 1.25, 1.5, 2] + image_ratio
|
||||
elif sample_size >= 768:
|
||||
number_list = [1, 1.25, 1.5] + image_ratio
|
||||
elif sample_size >= 512:
|
||||
number_list = [1] + image_ratio
|
||||
else:
|
||||
number_list = [1]
|
||||
|
||||
if all_choices:
|
||||
return number_list
|
||||
|
||||
number_list_prob = np.array(_create_special_list(len(number_list)))
|
||||
if rng is None:
|
||||
return np.random.choice(number_list, p = number_list_prob)
|
||||
else:
|
||||
return rng.choice(number_list, p = number_list_prob)
|
||||
|
||||
# Get token length
|
||||
target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size
|
||||
length_to_frame_num = get_length_to_frame_num(target_token_length)
|
||||
@@ -1058,7 +1069,7 @@ def main():
|
||||
data_type = examples[0]["data_type"]
|
||||
f, h, w, c = np.shape(pixel_value)
|
||||
if data_type == 'image':
|
||||
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size], rng=rng)
|
||||
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size])
|
||||
|
||||
aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
|
||||
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
|
||||
@@ -1072,15 +1083,11 @@ def main():
|
||||
choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25]
|
||||
if len(choice_list) == 0:
|
||||
choice_list = list(length_to_frame_num.keys())
|
||||
if rng is None:
|
||||
local_video_sample_size = np.random.choice(choice_list)
|
||||
else:
|
||||
local_video_sample_size = rng.choice(choice_list)
|
||||
local_video_sample_size = np.random.choice(choice_list)
|
||||
batch_video_length = length_to_frame_num[local_video_sample_size]
|
||||
random_downsample_ratio = args.video_sample_size / local_video_sample_size
|
||||
else:
|
||||
random_downsample_ratio = get_random_downsample_ratio(
|
||||
args.video_sample_size, rng=rng)
|
||||
random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size)
|
||||
batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
|
||||
else:
|
||||
random_downsample_ratio = 1
|
||||
@@ -1092,14 +1099,9 @@ def main():
|
||||
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
|
||||
closest_size = [int(x / 16) * 16 for x in closest_size]
|
||||
if args.random_ratio_crop:
|
||||
if rng is None:
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
else:
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
random_sample_size = aspect_ratio_random_crop_sample_size[
|
||||
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
|
||||
]
|
||||
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
|
||||
|
||||
for example in examples:
|
||||
@@ -1195,6 +1197,7 @@ def main():
|
||||
collate_fn=collate_fn,
|
||||
persistent_workers=True if args.dataloader_num_workers != 0 else False,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
|
||||
)
|
||||
else:
|
||||
# DataLoaders creation:
|
||||
@@ -1205,6 +1208,7 @@ def main():
|
||||
batch_sampler=batch_sampler,
|
||||
persistent_workers=True if args.dataloader_num_workers != 0 else False,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
|
||||
)
|
||||
|
||||
# Scheduler and math around the number of training steps.
|
||||
@@ -1225,7 +1229,10 @@ def main():
|
||||
network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||
network, optimizer, train_dataloader, lr_scheduler
|
||||
)
|
||||
|
||||
if zero_stage == 3:
|
||||
transformer3d = accelerator_transformer3d.prepare(
|
||||
transformer3d
|
||||
)
|
||||
# Move text_encode and vae to gpu and cast to weight_dtype
|
||||
vae.to(accelerator.device, dtype=weight_dtype)
|
||||
transformer3d.to(accelerator.device, dtype=weight_dtype)
|
||||
@@ -1297,10 +1304,11 @@ def main():
|
||||
first_epoch = global_step // num_update_steps_per_epoch
|
||||
print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
|
||||
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
|
||||
m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
if zero_stage != 3:
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
|
||||
m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
|
||||
accelerator.print(f"Resuming from checkpoint {path}")
|
||||
accelerator.load_state(os.path.join(args.output_dir, path))
|
||||
@@ -1356,7 +1364,7 @@ def main():
|
||||
pixel_values = batch["pixel_values"].to(weight_dtype)
|
||||
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length:
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
|
||||
if args.enable_text_encoder_in_dataloader:
|
||||
@@ -1377,7 +1385,7 @@ def main():
|
||||
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
|
||||
mask = batch["mask"].to(weight_dtype)
|
||||
# Increase the batch size when the length of the latent sequence of the current sample is small
|
||||
if args.training_with_video_token_length:
|
||||
if args.training_with_video_token_length and zero_stage != 3:
|
||||
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
|
||||
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
|
||||
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
|
||||
|
||||
Regular → Executable
+1
-1
@@ -620,7 +620,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
|
||||
# 3. Transformer blocks
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
|
||||
@@ -5,10 +5,11 @@ import glob
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import warnings
|
||||
from typing import Any, Dict
|
||||
import types
|
||||
import warnings
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
import torch.nn as nn
|
||||
@@ -18,12 +19,11 @@ from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.utils import is_torch_version, logging
|
||||
from torch import nn
|
||||
|
||||
from .cache_utils import TeaCache
|
||||
from ..dist import (get_sequence_parallel_rank,
|
||||
get_sequence_parallel_world_size,
|
||||
get_sp_group,
|
||||
get_sequence_parallel_world_size, get_sp_group,
|
||||
xFuserLongContextAttention)
|
||||
from ..dist.wan_xfuser import usp_attn_forward
|
||||
from .cache_utils import TeaCache
|
||||
|
||||
try:
|
||||
import flash_attn_interface
|
||||
@@ -219,6 +219,65 @@ def rope_params(max_seq_len, dim, theta=10000):
|
||||
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
||||
return freqs
|
||||
|
||||
# modified from https://github.com/thu-ml/RIFLEx/blob/main/riflex_utils.py
|
||||
@amp.autocast(enabled=False)
|
||||
def get_1d_rotary_pos_embed_riflex(
|
||||
pos: Union[np.ndarray, int],
|
||||
dim: int,
|
||||
theta: float = 10000.0,
|
||||
use_real=False,
|
||||
k: Optional[int] = None,
|
||||
L_test: Optional[int] = None,
|
||||
L_test_scale: Optional[int] = None,
|
||||
):
|
||||
"""
|
||||
RIFLEx: Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
|
||||
|
||||
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
|
||||
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
|
||||
data type.
|
||||
|
||||
Args:
|
||||
dim (`int`): Dimension of the frequency tensor.
|
||||
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
|
||||
theta (`float`, *optional*, defaults to 10000.0):
|
||||
Scaling factor for frequency computation. Defaults to 10000.0.
|
||||
use_real (`bool`, *optional*):
|
||||
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
|
||||
k (`int`, *optional*, defaults to None): the index for the intrinsic frequency in RoPE
|
||||
L_test (`int`, *optional*, defaults to None): the number of frames for inference
|
||||
Returns:
|
||||
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
|
||||
"""
|
||||
assert dim % 2 == 0
|
||||
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos)
|
||||
if isinstance(pos, np.ndarray):
|
||||
pos = torch.from_numpy(pos) # type: ignore # [S]
|
||||
|
||||
freqs = 1.0 / torch.pow(theta,
|
||||
torch.arange(0, dim, 2).to(torch.float64).div(dim))
|
||||
|
||||
# === Riflex modification start ===
|
||||
# Reduce the intrinsic frequency to stay within a single period after extrapolation (see Eq. (8)).
|
||||
# Empirical observations show that a few videos may exhibit repetition in the tail frames.
|
||||
# To be conservative, we multiply by 0.9 to keep the extrapolated length below 90% of a single period.
|
||||
if k is not None:
|
||||
freqs[k-1] = 0.9 * 2 * torch.pi / L_test
|
||||
# === Riflex modification end ===
|
||||
if L_test_scale is not None:
|
||||
freqs[k-1] = freqs[k-1] / L_test_scale
|
||||
|
||||
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
|
||||
if use_real:
|
||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
|
||||
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
|
||||
return freqs_cos, freqs_sin
|
||||
else:
|
||||
# lumina
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
|
||||
return freqs_cis
|
||||
|
||||
@amp.autocast(enabled=False)
|
||||
def rope_apply(x, grid_sizes, freqs):
|
||||
@@ -320,9 +379,9 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
q = self.norm_q(self.q(x.to(dtype))).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x.to(dtype))).view(b, s, n, d)
|
||||
v = self.v(x.to(dtype)).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
@@ -343,7 +402,7 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, context_lens):
|
||||
def forward(self, x, context, context_lens, dtype):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -353,12 +412,18 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
q = self.norm_q(self.q(x.to(dtype))).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context.to(dtype))).view(b, -1, n, d)
|
||||
v = self.v(context.to(dtype)).view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
x = attention(q, k, v, k_lens=context_lens)
|
||||
x = attention(
|
||||
q.to(dtype),
|
||||
k.to(dtype),
|
||||
v.to(dtype),
|
||||
k_lens=context_lens
|
||||
)
|
||||
x = x.to(dtype)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
@@ -381,7 +446,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
# self.alpha = nn.Parameter(torch.zeros((1, )))
|
||||
self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, context, context_lens):
|
||||
def forward(self, x, context, context_lens, dtype):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -393,14 +458,27 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
k_img = self.norm_k_img(self.k_img(context_img)).view(b, -1, n, d)
|
||||
v_img = self.v_img(context_img).view(b, -1, n, d)
|
||||
img_x = attention(q, k_img, v_img, k_lens=None)
|
||||
q = self.norm_q(self.q(x.to(dtype))).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(context.to(dtype))).view(b, -1, n, d)
|
||||
v = self.v(context.to(dtype)).view(b, -1, n, d)
|
||||
k_img = self.norm_k_img(self.k_img(context_img.to(dtype))).view(b, -1, n, d)
|
||||
v_img = self.v_img(context_img.to(dtype)).view(b, -1, n, d)
|
||||
|
||||
img_x = attention(
|
||||
q.to(dtype),
|
||||
k_img.to(dtype),
|
||||
v_img.to(dtype),
|
||||
k_lens=None
|
||||
)
|
||||
img_x = img_x.to(dtype)
|
||||
# compute attention
|
||||
x = attention(q, k, v, k_lens=context_lens)
|
||||
x = attention(
|
||||
q.to(dtype),
|
||||
k.to(dtype),
|
||||
v.to(dtype),
|
||||
k_lens=context_lens
|
||||
)
|
||||
x = x.to(dtype)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
@@ -486,7 +564,10 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
# cross-attention & ffn function
|
||||
def cross_attn_ffn(x, context, context_lens, e):
|
||||
x = x + self.cross_attn(self.norm3(x), context, context_lens)
|
||||
# cross-attention
|
||||
x = x + self.cross_attn(self.norm3(x), context, context_lens, dtype)
|
||||
|
||||
# ffn function
|
||||
temp_x = self.norm2(x) * (1 + e[4]) + e[3]
|
||||
temp_x = temp_x.to(dtype)
|
||||
|
||||
@@ -655,6 +736,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
d = dim // num_heads
|
||||
self.d = d
|
||||
self.freqs = torch.cat(
|
||||
[
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
@@ -684,8 +766,35 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
coefficients, num_steps, rel_l1_thresh=rel_l1_thresh, num_skip_start_steps=num_skip_start_steps, offload=offload
|
||||
)
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
self.gradient_checkpointing = value
|
||||
def disable_teacache(self):
|
||||
self.teacache = None
|
||||
|
||||
def enable_riflex(
|
||||
self,
|
||||
k = 6,
|
||||
L_test = 66,
|
||||
L_test_scale = 4.886,
|
||||
):
|
||||
device = self.freqs.device
|
||||
self.freqs = torch.cat(
|
||||
[
|
||||
get_1d_rotary_pos_embed_riflex(1024, self.d - 4 * (self.d // 6), use_real=False, k=k, L_test=L_test, L_test_scale=L_test_scale),
|
||||
rope_params(1024, 2 * (self.d // 6)),
|
||||
rope_params(1024, 2 * (self.d // 6))
|
||||
],
|
||||
dim=1
|
||||
).to(device)
|
||||
|
||||
def disable_riflex(self):
|
||||
device = self.freqs.device
|
||||
self.freqs = torch.cat(
|
||||
[
|
||||
rope_params(1024, self.d - 4 * (self.d // 6)),
|
||||
rope_params(1024, 2 * (self.d // 6)),
|
||||
rope_params(1024, 2 * (self.d // 6))
|
||||
],
|
||||
dim=1
|
||||
).to(device)
|
||||
|
||||
def enable_multi_gpus_inference(self,):
|
||||
self.sp_world_size = get_sequence_parallel_world_size()
|
||||
@@ -693,6 +802,9 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
for block in self.blocks:
|
||||
block.self_attn.forward = types.MethodType(
|
||||
usp_attn_forward, block.self_attn)
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
self.gradient_checkpointing = value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -757,8 +869,8 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t).float())
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
||||
assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
# to bfloat16 for saving memeory
|
||||
# assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
e0 = e0.to(dtype)
|
||||
e = e.to(dtype)
|
||||
|
||||
@@ -813,7 +925,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
ori_x = x.clone().cpu() if self.teacache.offload else x.clone()
|
||||
|
||||
for block in self.blocks:
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
@@ -852,7 +964,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
self.teacache.previous_residual_uncond = x.cpu() - ori_x if self.teacache.offload else x - ori_x
|
||||
else:
|
||||
for block in self.blocks:
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
|
||||
Regular → Executable
-4
@@ -535,10 +535,6 @@ class WanFunPipeline(DiffusionPipeline):
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
prompt_embeds_2 = callback_outputs.pop("prompt_embeds_2", prompt_embeds_2)
|
||||
negative_prompt_embeds_2 = callback_outputs.pop(
|
||||
"negative_prompt_embeds_2", negative_prompt_embeds_2
|
||||
)
|
||||
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
Regular → Executable
-4
@@ -700,10 +700,6 @@ class WanFunControlPipeline(DiffusionPipeline):
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
prompt_embeds_2 = callback_outputs.pop("prompt_embeds_2", prompt_embeds_2)
|
||||
negative_prompt_embeds_2 = callback_outputs.pop(
|
||||
"negative_prompt_embeds_2", negative_prompt_embeds_2
|
||||
)
|
||||
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
Regular → Executable
+1
-5
@@ -702,11 +702,7 @@ class WanFunInpaintPipeline(DiffusionPipeline):
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
prompt_embeds_2 = callback_outputs.pop("prompt_embeds_2", prompt_embeds_2)
|
||||
negative_prompt_embeds_2 = callback_outputs.pop(
|
||||
"negative_prompt_embeds_2", negative_prompt_embeds_2
|
||||
)
|
||||
|
||||
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if comfyui_progressbar:
|
||||
|
||||
@@ -57,7 +57,8 @@ class Fun_Controller:
|
||||
self, GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
|
||||
config_path=None, ulysses_degree=1, ring_degree=1,
|
||||
enable_teacache=None, teacache_threshold=None,
|
||||
num_skip_start_steps=None, teacache_offload=None, weight_dtype=None,
|
||||
num_skip_start_steps=None, teacache_offload=None,
|
||||
enable_riflex=None, riflex_k=None, weight_dtype=None,
|
||||
):
|
||||
# config dirs
|
||||
self.basedir = os.getcwd()
|
||||
@@ -81,6 +82,8 @@ class Fun_Controller:
|
||||
self.teacache_threshold = teacache_threshold
|
||||
self.num_skip_start_steps = num_skip_start_steps
|
||||
self.teacache_offload = teacache_offload
|
||||
self.enable_riflex = enable_riflex
|
||||
self.riflex_k = riflex_k
|
||||
self.weight_dtype = weight_dtype
|
||||
self.device = set_multi_gpus_devices(self.ulysses_degree, self.ring_degree)
|
||||
|
||||
|
||||
Regular → Executable
+25
-83
@@ -186,91 +186,31 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
else: seed_textbox = np.random.randint(0, 1e10)
|
||||
generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox))
|
||||
|
||||
if self.enable_riflex:
|
||||
self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = length_slider if not is_image else 1)
|
||||
|
||||
try:
|
||||
if self.model_type == "Inpaint":
|
||||
if self.transformer.config.in_channels != self.vae.config.latent_channels:
|
||||
if generation_method == "Long Video Generation":
|
||||
if validation_video is not None:
|
||||
raise gr.Error(f"Video to Video is not Support Long Video Generation now.")
|
||||
init_frames = 0
|
||||
last_frames = init_frames + partial_video_length
|
||||
while init_frames < length_slider:
|
||||
if last_frames >= length_slider:
|
||||
_partial_video_length = length_slider - init_frames
|
||||
_partial_video_length = int((_partial_video_length - 1) // self.vae.config.temporal_compression_ratio * self.vae.config.temporal_compression_ratio) + 1
|
||||
|
||||
if _partial_video_length <= 0:
|
||||
break
|
||||
else:
|
||||
_partial_video_length = partial_video_length
|
||||
|
||||
if last_frames >= length_slider:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=_partial_video_length, sample_size=(height_slider, width_slider))
|
||||
else:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, None, video_length=_partial_video_length, sample_size=(height_slider, width_slider))
|
||||
|
||||
with torch.no_grad():
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = _partial_video_length,
|
||||
generator = generator,
|
||||
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image
|
||||
).videos
|
||||
|
||||
if init_frames != 0:
|
||||
mix_ratio = torch.from_numpy(
|
||||
np.array([float(_index) / float(overlap_video_length) for _index in range(overlap_video_length)], np.float32)
|
||||
).unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
|
||||
new_sample[:, :, -overlap_video_length:] = new_sample[:, :, -overlap_video_length:] * (1 - mix_ratio) + \
|
||||
sample[:, :, :overlap_video_length] * mix_ratio
|
||||
new_sample = torch.cat([new_sample, sample[:, :, overlap_video_length:]], dim = 2)
|
||||
|
||||
sample = new_sample
|
||||
else:
|
||||
new_sample = sample
|
||||
|
||||
if last_frames >= length_slider:
|
||||
break
|
||||
|
||||
start_image = [
|
||||
Image.fromarray(
|
||||
(sample[0, :, _index].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
|
||||
) for _index in range(-overlap_video_length, 0)
|
||||
]
|
||||
|
||||
init_frames = init_frames + _partial_video_length - overlap_video_length
|
||||
last_frames = init_frames + _partial_video_length
|
||||
if validation_video is not None:
|
||||
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=16)
|
||||
else:
|
||||
if validation_video is not None:
|
||||
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=16)
|
||||
strength = denoise_strength
|
||||
else:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
|
||||
strength = 1
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
|
||||
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image
|
||||
).videos
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image
|
||||
).videos
|
||||
else:
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
@@ -335,12 +275,13 @@ class Wan_Fun_Controller(Fun_Controller):
|
||||
Wan_Fun_Controller_Host = Wan_Fun_Controller
|
||||
Wan_Fun_Controller_Client = Fun_Controller_Client
|
||||
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype):
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype):
|
||||
controller = Wan_Fun_Controller(
|
||||
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
|
||||
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
|
||||
enable_teacache=enable_teacache, teacache_threshold=teacache_threshold,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, weight_dtype=weight_dtype,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload,
|
||||
enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype,
|
||||
)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
@@ -460,12 +401,13 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree
|
||||
)
|
||||
return demo, controller
|
||||
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype):
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype):
|
||||
controller = Wan_Fun_Controller_Host(
|
||||
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
|
||||
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
|
||||
enable_teacache=enable_teacache, teacache_threshold=teacache_threshold,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, weight_dtype=weight_dtype,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload,
|
||||
enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype,
|
||||
)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
|
||||
Regular → Executable
+25
-83
@@ -186,91 +186,31 @@ class Wan_Controller(Fun_Controller):
|
||||
else: seed_textbox = np.random.randint(0, 1e10)
|
||||
generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox))
|
||||
|
||||
if self.enable_riflex:
|
||||
self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = length_slider if not is_image else 1)
|
||||
|
||||
try:
|
||||
if self.model_type == "Inpaint":
|
||||
if self.transformer.config.in_channels != self.vae.config.latent_channels:
|
||||
if generation_method == "Long Video Generation":
|
||||
if validation_video is not None:
|
||||
raise gr.Error(f"Video to Video is not Support Long Video Generation now.")
|
||||
init_frames = 0
|
||||
last_frames = init_frames + partial_video_length
|
||||
while init_frames < length_slider:
|
||||
if last_frames >= length_slider:
|
||||
_partial_video_length = length_slider - init_frames
|
||||
_partial_video_length = int((_partial_video_length - 1) // self.vae.config.temporal_compression_ratio * self.vae.config.temporal_compression_ratio) + 1
|
||||
|
||||
if _partial_video_length <= 0:
|
||||
break
|
||||
else:
|
||||
_partial_video_length = partial_video_length
|
||||
|
||||
if last_frames >= length_slider:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=_partial_video_length, sample_size=(height_slider, width_slider))
|
||||
else:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, None, video_length=_partial_video_length, sample_size=(height_slider, width_slider))
|
||||
|
||||
with torch.no_grad():
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = _partial_video_length,
|
||||
generator = generator,
|
||||
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image
|
||||
).videos
|
||||
|
||||
if init_frames != 0:
|
||||
mix_ratio = torch.from_numpy(
|
||||
np.array([float(_index) / float(overlap_video_length) for _index in range(overlap_video_length)], np.float32)
|
||||
).unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
|
||||
new_sample[:, :, -overlap_video_length:] = new_sample[:, :, -overlap_video_length:] * (1 - mix_ratio) + \
|
||||
sample[:, :, :overlap_video_length] * mix_ratio
|
||||
new_sample = torch.cat([new_sample, sample[:, :, overlap_video_length:]], dim = 2)
|
||||
|
||||
sample = new_sample
|
||||
else:
|
||||
new_sample = sample
|
||||
|
||||
if last_frames >= length_slider:
|
||||
break
|
||||
|
||||
start_image = [
|
||||
Image.fromarray(
|
||||
(sample[0, :, _index].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
|
||||
) for _index in range(-overlap_video_length, 0)
|
||||
]
|
||||
|
||||
init_frames = init_frames + _partial_video_length - overlap_video_length
|
||||
last_frames = init_frames + _partial_video_length
|
||||
if validation_video is not None:
|
||||
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=16)
|
||||
else:
|
||||
if validation_video is not None:
|
||||
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=16)
|
||||
strength = denoise_strength
|
||||
else:
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
|
||||
strength = 1
|
||||
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
|
||||
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
negative_prompt = negative_prompt_textbox,
|
||||
num_inference_steps = sample_step_slider,
|
||||
guidance_scale = cfg_scale_slider,
|
||||
width = width_slider,
|
||||
height = height_slider,
|
||||
num_frames = length_slider if not is_image else 1,
|
||||
generator = generator,
|
||||
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image
|
||||
).videos
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
clip_image = clip_image
|
||||
).videos
|
||||
else:
|
||||
sample = self.pipeline(
|
||||
prompt_textbox,
|
||||
@@ -335,12 +275,13 @@ class Wan_Controller(Fun_Controller):
|
||||
Wan_Controller_Host = Wan_Controller
|
||||
Wan_Controller_Client = Fun_Controller_Client
|
||||
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype):
|
||||
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype):
|
||||
controller = Wan_Controller(
|
||||
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
|
||||
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
|
||||
enable_teacache=enable_teacache, teacache_threshold=teacache_threshold,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, weight_dtype=weight_dtype,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload,
|
||||
enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype,
|
||||
)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
@@ -456,12 +397,13 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree
|
||||
)
|
||||
return demo, controller
|
||||
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype):
|
||||
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype):
|
||||
controller = Wan_Controller_Host(
|
||||
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
|
||||
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
|
||||
enable_teacache=enable_teacache, teacache_threshold=teacache_threshold,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, weight_dtype=weight_dtype,
|
||||
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload,
|
||||
enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype,
|
||||
)
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
|
||||
@@ -169,7 +169,7 @@ class LoRANetwork(torch.nn.Module):
|
||||
alpha: float = 1,
|
||||
dropout: Optional[float] = None,
|
||||
module_class: Type[object] = LoRAModule,
|
||||
add_lora_in_attn_temporal: bool = False,
|
||||
skip_name: str = None,
|
||||
varbose: Optional[bool] = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -202,9 +202,8 @@ class LoRANetwork(torch.nn.Module):
|
||||
is_conv2d = child_module.__class__.__name__ == "Conv2d" or child_module.__class__.__name__ == "LoRACompatibleConv"
|
||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
||||
|
||||
if not add_lora_in_attn_temporal:
|
||||
if "attn_temporal" in child_name:
|
||||
continue
|
||||
if skip_name is not None and skip_name in child_name:
|
||||
continue
|
||||
|
||||
if is_linear or is_conv2d:
|
||||
lora_name = prefix + "." + name + "." + child_name
|
||||
@@ -346,7 +345,7 @@ def create_network(
|
||||
text_encoder: Union[T5EncoderModel, List[T5EncoderModel]],
|
||||
transformer,
|
||||
neuron_dropout: Optional[float] = None,
|
||||
add_lora_in_attn_temporal: bool = False,
|
||||
skip_name: str = None,
|
||||
**kwargs,
|
||||
):
|
||||
if network_dim is None:
|
||||
@@ -361,7 +360,7 @@ def create_network(
|
||||
lora_dim=network_dim,
|
||||
alpha=network_alpha,
|
||||
dropout=neuron_dropout,
|
||||
add_lora_in_attn_temporal=add_lora_in_attn_temporal,
|
||||
skip_name=skip_name,
|
||||
varbose=True,
|
||||
)
|
||||
return network
|
||||
|
||||
Reference in New Issue
Block a user