From f149dba611d878941f88a80b78ae85dc0dee829e Mon Sep 17 00:00:00 2001
From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com>
Date: Wed, 9 Apr 2025 15:23:38 +0800
Subject: [PATCH] Update Riflex && Pipeline callback bug fix && Training Code
bug fix (#150)
---
Dockerfile.ds | 52 +
comfyui/comfyui_nodes.py | 21 +-
comfyui/wan2_1/nodes.py | 20 +-
comfyui/wan2_1_fun/nodes.py | 32 +-
examples/wan2.1/app.py | 5 +
examples/wan2.1/predict_i2v.py | 8 +
examples/wan2.1/predict_t2v.py | 11 +-
examples/wan2.1_fun/app.py | 9 +-
examples/wan2.1_fun/predict_i2v.py | 9 +
examples/wan2.1_fun/predict_t2v.py | 9 +
examples/wan2.1_fun/predict_v2v_control.py | 9 +
scripts/wan2.1/README_TRAIN.md | 6 +-
scripts/wan2.1/README_TRAIN_LORA.md | 6 +-
scripts/wan2.1_fun/README_TRAIN.md | 37 +-
scripts/wan2.1_fun/README_TRAIN_CONTROL.md | 22 +-
.../wan2.1_fun/README_TRAIN_CONTROL_LORA.md | 180 ++
scripts/wan2.1_fun/README_TRAIN_LORA.md | 6 +-
scripts/wan2.1_fun/README_TRAIN_REWARD.md | 51 +-
scripts/wan2.1_fun/train.py | 175 +-
scripts/wan2.1_fun/train_control.py | 197 +-
scripts/wan2.1_fun/train_control_lora.py | 1682 +++++++++++++++++
scripts/wan2.1_fun/train_control_lora.sh | 39 +
scripts/wan2.1_fun/train_lora.py | 226 +--
videox_fun/models/cogvideox_transformer3d.py | 2 +-
videox_fun/models/wan_transformer3d.py | 166 +-
videox_fun/pipeline/pipeline_wan_fun.py | 4 -
.../pipeline/pipeline_wan_fun_control.py | 4 -
.../pipeline/pipeline_wan_fun_inpaint.py | 6 +-
videox_fun/ui/controller.py | 5 +-
videox_fun/ui/wan_fun_ui.py | 108 +-
videox_fun/ui/wan_ui.py | 108 +-
videox_fun/utils/lora_utils.py | 11 +-
32 files changed, 2626 insertions(+), 600 deletions(-)
create mode 100644 Dockerfile.ds
mode change 100644 => 100755 comfyui/comfyui_nodes.py
mode change 100644 => 100755 comfyui/wan2_1/nodes.py
mode change 100644 => 100755 comfyui/wan2_1_fun/nodes.py
mode change 100644 => 100755 examples/wan2.1_fun/predict_i2v.py
mode change 100644 => 100755 examples/wan2.1_fun/predict_v2v_control.py
mode change 100644 => 100755 scripts/wan2.1/README_TRAIN.md
mode change 100644 => 100755 scripts/wan2.1/README_TRAIN_LORA.md
mode change 100644 => 100755 scripts/wan2.1_fun/README_TRAIN.md
create mode 100644 scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
mode change 100644 => 100755 scripts/wan2.1_fun/README_TRAIN_LORA.md
mode change 100644 => 100755 scripts/wan2.1_fun/README_TRAIN_REWARD.md
create mode 100644 scripts/wan2.1_fun/train_control_lora.py
create mode 100644 scripts/wan2.1_fun/train_control_lora.sh
mode change 100644 => 100755 videox_fun/models/cogvideox_transformer3d.py
mode change 100644 => 100755 videox_fun/pipeline/pipeline_wan_fun.py
mode change 100644 => 100755 videox_fun/pipeline/pipeline_wan_fun_control.py
mode change 100644 => 100755 videox_fun/pipeline/pipeline_wan_fun_inpaint.py
mode change 100644 => 100755 videox_fun/ui/wan_fun_ui.py
mode change 100644 => 100755 videox_fun/ui/wan_ui.py
diff --git a/Dockerfile.ds b/Dockerfile.ds
new file mode 100644
index 0000000..a7e186d
--- /dev/null
+++ b/Dockerfile.ds
@@ -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/
\ No newline at end of file
diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py
old mode 100644
new mode 100755
index dfd096b..16567d7
--- a/comfyui/comfyui_nodes.py
+++ b/comfyui/comfyui_nodes.py
@@ -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",
diff --git a/comfyui/wan2_1/nodes.py b/comfyui/wan2_1/nodes.py
old mode 100644
new mode 100755
index f550356..52bb724
--- a/comfyui/wan2_1/nodes.py
+++ b/comfyui/wan2_1/nodes.py
@@ -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:
diff --git a/comfyui/wan2_1_fun/nodes.py b/comfyui/wan2_1_fun/nodes.py
old mode 100644
new mode 100755
index 27935e0..6c04eb5
--- a/comfyui/wan2_1_fun/nodes.py
+++ b/comfyui/wan2_1_fun/nodes.py
@@ -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:
diff --git a/examples/wan2.1/app.py b/examples/wan2.1/app.py
index 229d741..a613edb 100755
--- a/examples/wan2.1/app.py
+++ b/examples/wan2.1/app.py
@@ -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
diff --git a/examples/wan2.1/predict_i2v.py b/examples/wan2.1/predict_i2v.py
index 3f8aa00..c65ac9f 100755
--- a/examples/wan2.1/predict_i2v.py
+++ b/examples/wan2.1/predict_i2v.py
@@ -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(
diff --git a/examples/wan2.1/predict_t2v.py b/examples/wan2.1/predict_t2v.py
index f641c21..6fd9040 100755
--- a/examples/wan2.1/predict_t2v.py
+++ b/examples/wan2.1/predict_t2v.py
@@ -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,
diff --git a/examples/wan2.1_fun/app.py b/examples/wan2.1_fun/app.py
index 5f671f5..a84b69f 100755
--- a/examples/wan2.1_fun/app.py
+++ b/examples/wan2.1_fun/app.py
@@ -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
diff --git a/examples/wan2.1_fun/predict_i2v.py b/examples/wan2.1_fun/predict_i2v.py
old mode 100644
new mode 100755
index d9ab88e..0e758a2
--- a/examples/wan2.1_fun/predict_i2v.py
+++ b/examples/wan2.1_fun/predict_i2v.py
@@ -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(
diff --git a/examples/wan2.1_fun/predict_t2v.py b/examples/wan2.1_fun/predict_t2v.py
index cee8857..1bf8867 100755
--- a/examples/wan2.1_fun/predict_t2v.py
+++ b/examples/wan2.1_fun/predict_t2v.py
@@ -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)
diff --git a/examples/wan2.1_fun/predict_v2v_control.py b/examples/wan2.1_fun/predict_v2v_control.py
old mode 100644
new mode 100755
index 9015331..33b1c8c
--- a/examples/wan2.1_fun/predict_v2v_control.py
+++ b/examples/wan2.1_fun/predict_v2v_control.py
@@ -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(
diff --git a/scripts/wan2.1/README_TRAIN.md b/scripts/wan2.1/README_TRAIN.md
old mode 100644
new mode 100755
index 7f70d00..4fa49c3
--- a/scripts/wan2.1/README_TRAIN.md
+++ b/scripts/wan2.1/README_TRAIN.md
@@ -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 \
diff --git a/scripts/wan2.1/README_TRAIN_LORA.md b/scripts/wan2.1/README_TRAIN_LORA.md
old mode 100644
new mode 100755
index a6ac4b7..2e507c8
--- a/scripts/wan2.1/README_TRAIN_LORA.md
+++ b/scripts/wan2.1/README_TRAIN_LORA.md
@@ -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 \
diff --git a/scripts/wan2.1_fun/README_TRAIN.md b/scripts/wan2.1_fun/README_TRAIN.md
old mode 100644
new mode 100755
index acd333a..37108e6
--- a/scripts/wan2.1_fun/README_TRAIN.md
+++ b/scripts/wan2.1_fun/README_TRAIN.md
@@ -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 "."
```
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md
index edd14c3..36feed9 100644
--- a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md
+++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md
@@ -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 "."
```
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
new file mode 100644
index 0000000..494035c
--- /dev/null
+++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md
@@ -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
+```
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/README_TRAIN_LORA.md b/scripts/wan2.1_fun/README_TRAIN_LORA.md
old mode 100644
new mode 100755
index fc3d0ec..c4b9de5
--- a/scripts/wan2.1_fun/README_TRAIN_LORA.md
+++ b/scripts/wan2.1_fun/README_TRAIN_LORA.md
@@ -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 \
diff --git a/scripts/wan2.1_fun/README_TRAIN_REWARD.md b/scripts/wan2.1_fun/README_TRAIN_REWARD.md
old mode 100644
new mode 100755
index d2f58e6..d7a18a1
--- a/scripts/wan2.1_fun/README_TRAIN_REWARD.md
+++ b/scripts/wan2.1_fun/README_TRAIN_REWARD.md
@@ -33,13 +33,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
-
+
|
-
+
|
-
+
|
@@ -51,13 +51,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
-
+
|
-
+
|
-
+
|
@@ -69,13 +69,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
-
+
|
-
+
|
-
+
|
@@ -87,13 +87,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
-
+
|
-
+
|
-
+
|
@@ -122,13 +122,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
-
+
|
-
+
|
-
+
|
@@ -140,13 +140,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
-
+
|
-
+
|
-
+
|
@@ -158,13 +158,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
-
+
|
-
+
|
-
+
|
@@ -176,17 +176,18 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
|
-
+
|
-
+
|
-
+
|
+
> [!NOTE]
> The above test prompts are from T2V-CompBench 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
- Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.
- Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).
-
\ No newline at end of file
+
diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py
index a28c929..019981d 100755
--- a/scripts/wan2.1_fun/train.py
+++ b/scripts/wan2.1_fun/train.py
@@ -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))
diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py
index 95e7615..378a9ec 100755
--- a/scripts/wan2.1_fun/train_control.py
+++ b/scripts/wan2.1_fun/train_control.py
@@ -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))
diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py
new file mode 100644
index 0000000..2dc56d5
--- /dev/null
+++ b/scripts/wan2.1_fun/train_control_lora.py
@@ -0,0 +1,1682 @@
+"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py
+"""
+#!/usr/bin/env python
+# coding=utf-8
+# Copyright 2024 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+
+import argparse
+import gc
+import logging
+import math
+import os
+import pickle
+import random
+import shutil
+import sys
+
+import accelerate
+import diffusers
+import numpy as np
+import torch
+import torch.nn.functional as F
+import torch.utils.checkpoint
+import torchvision.transforms.functional as TF
+import transformers
+from accelerate import Accelerator, FullyShardedDataParallelPlugin
+from accelerate.logging import get_logger
+from accelerate.state import AcceleratorState
+from accelerate.utils import ProjectConfiguration, set_seed
+from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler
+from diffusers.optimization import get_scheduler
+from diffusers.training_utils import (EMAModel,
+ compute_density_for_timestep_sampling,
+ compute_loss_weighting_for_sd3)
+from diffusers.utils import check_min_version, deprecate, is_wandb_available
+from diffusers.utils.torch_utils import is_compiled_module
+from einops import rearrange
+from omegaconf import OmegaConf
+from packaging import version
+from PIL import Image
+from torch.utils.data import RandomSampler
+from torch.utils.tensorboard import SummaryWriter
+from torchvision import transforms
+from tqdm.auto import tqdm
+from transformers import AutoTokenizer
+from transformers.utils import ContextManagers
+
+import datasets
+
+current_file_path = os.path.abspath(__file__)
+project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
+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)
+from videox_fun.data.dataset_image_video import (ImageVideoControlDataset,
+ ImageVideoDataset,
+ ImageVideoSampler,
+ get_random_mask)
+from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
+ WanTransformer3DModel)
+from videox_fun.pipeline import WanFunControlPipeline
+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.utils import (get_image_to_video_latent,
+ get_video_to_video_latent,
+ save_videos_grid)
+
+if is_wandb_available():
+ import wandb
+
+
+def filter_kwargs(cls, kwargs):
+ import inspect
+ sig = inspect.signature(cls.__init__)
+ valid_params = set(sig.parameters.keys()) - {'self', 'cls'}
+ filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
+ return filtered_kwargs
+
+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")
+
+logger = get_logger(__name__, log_level="INFO")
+
+def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, config, args, accelerator, weight_dtype, global_step):
+ try:
+ logger.info("Running validation... ")
+
+ transformer3d_val = WanTransformer3DModel.from_pretrained(
+ os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
+ ).to(weight_dtype)
+ transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
+ scheduler = FlowMatchEulerDiscreteScheduler(
+ **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
+ )
+
+ pipeline = WanFunControlPipeline(
+ vae=accelerator.unwrap_model(vae).to(weight_dtype),
+ text_encoder=accelerator.unwrap_model(text_encoder),
+ tokenizer=tokenizer,
+ transformer=transformer3d_val,
+ scheduler=scheduler,
+ clip_image_encoder=clip_image_encoder,
+ )
+ pipeline = pipeline.to(accelerator.device)
+
+ pipeline = merge_lora(
+ pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True
+ )
+
+ if args.seed is None:
+ generator = None
+ else:
+ generator = torch.Generator(device=accelerator.device).manual_seed(args.seed)
+
+ images = []
+ for i in range(len(args.validation_prompts)):
+ with torch.no_grad():
+ with torch.autocast("cuda", dtype=weight_dtype):
+ video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
+ input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
+ sample = pipeline(
+ args.validation_prompts[i],
+ num_frames = video_length,
+ negative_prompt = "bad detailed",
+ height = args.video_sample_size,
+ width = args.video_sample_size,
+ generator = generator,
+
+ control_video = input_video,
+ ).videos
+ os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
+ save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif"))
+
+ del pipeline
+ del transformer3d_val
+ gc.collect()
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+
+ return images
+ except Exception as e:
+ gc.collect()
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+ print(f"Eval error with info {e}")
+ return None
+
+def parse_args():
+ parser = argparse.ArgumentParser(description="Simple example of a training script.")
+ parser.add_argument(
+ "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1."
+ )
+ parser.add_argument(
+ "--pretrained_model_name_or_path",
+ type=str,
+ default=None,
+ required=True,
+ help="Path to pretrained model or model identifier from huggingface.co/models.",
+ )
+ parser.add_argument(
+ "--revision",
+ type=str,
+ default=None,
+ required=False,
+ help="Revision of pretrained model identifier from huggingface.co/models.",
+ )
+ parser.add_argument(
+ "--variant",
+ type=str,
+ default=None,
+ help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16",
+ )
+ parser.add_argument(
+ "--train_data_dir",
+ type=str,
+ default=None,
+ help=(
+ "A folder containing the training data. "
+ ),
+ )
+ parser.add_argument(
+ "--train_data_meta",
+ type=str,
+ default=None,
+ help=(
+ "A csv containing the training data. "
+ ),
+ )
+ parser.add_argument(
+ "--max_train_samples",
+ type=int,
+ default=None,
+ help=(
+ "For debugging purposes or quicker training, truncate the number of training examples to this "
+ "value if set."
+ ),
+ )
+ parser.add_argument(
+ "--validation_prompts",
+ type=str,
+ default=None,
+ nargs="+",
+ help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
+ )
+ parser.add_argument(
+ "--output_dir",
+ type=str,
+ default="sd-model-finetuned",
+ help="The output directory where the model predictions and checkpoints will be written.",
+ )
+ parser.add_argument(
+ "--cache_dir",
+ type=str,
+ default=None,
+ help="The directory where the downloaded models and datasets will be stored.",
+ )
+ parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
+ parser.add_argument(
+ "--random_flip",
+ action="store_true",
+ help="whether to randomly flip images horizontally",
+ )
+ parser.add_argument(
+ "--use_came",
+ action="store_true",
+ help="whether to use came",
+ )
+ parser.add_argument(
+ "--multi_stream",
+ action="store_true",
+ help="whether to use cuda multi-stream",
+ )
+ parser.add_argument(
+ "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader."
+ )
+ parser.add_argument(
+ "--vae_mini_batch", type=int, default=32, help="mini batch size for vae."
+ )
+ parser.add_argument("--num_train_epochs", type=int, default=100)
+ parser.add_argument(
+ "--max_train_steps",
+ type=int,
+ default=None,
+ help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
+ )
+ parser.add_argument(
+ "--gradient_accumulation_steps",
+ type=int,
+ default=1,
+ help="Number of updates steps to accumulate before performing a backward/update pass.",
+ )
+ parser.add_argument(
+ "--gradient_checkpointing",
+ action="store_true",
+ help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
+ )
+ parser.add_argument(
+ "--learning_rate",
+ type=float,
+ default=1e-4,
+ help="Initial learning rate (after the potential warmup period) to use.",
+ )
+ parser.add_argument(
+ "--scale_lr",
+ action="store_true",
+ default=False,
+ help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
+ )
+ parser.add_argument(
+ "--lr_scheduler",
+ type=str,
+ default="constant",
+ help=(
+ 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
+ ' "constant", "constant_with_warmup"]'
+ ),
+ )
+ parser.add_argument(
+ "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler."
+ )
+ parser.add_argument(
+ "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes."
+ )
+ parser.add_argument(
+ "--allow_tf32",
+ action="store_true",
+ help=(
+ "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
+ " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
+ ),
+ )
+ parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.")
+ parser.add_argument(
+ "--non_ema_revision",
+ type=str,
+ default=None,
+ required=False,
+ help=(
+ "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or"
+ " remote repository specified with --pretrained_model_name_or_path."
+ ),
+ )
+ parser.add_argument(
+ "--dataloader_num_workers",
+ type=int,
+ default=0,
+ help=(
+ "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process."
+ ),
+ )
+ parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.")
+ parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.")
+ parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.")
+ parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer")
+ parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
+ parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.")
+ parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.")
+ parser.add_argument(
+ "--prediction_type",
+ type=str,
+ default=None,
+ help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.",
+ )
+ parser.add_argument(
+ "--hub_model_id",
+ type=str,
+ default=None,
+ help="The name of the repository to keep in sync with the local `output_dir`.",
+ )
+ parser.add_argument(
+ "--logging_dir",
+ type=str,
+ default="logs",
+ help=(
+ "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
+ " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
+ ),
+ )
+ parser.add_argument(
+ "--mixed_precision",
+ type=str,
+ default=None,
+ choices=["no", "fp16", "bf16"],
+ help=(
+ "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
+ " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
+ " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
+ ),
+ )
+ parser.add_argument(
+ "--report_to",
+ type=str,
+ default="tensorboard",
+ help=(
+ 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`'
+ ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.'
+ ),
+ )
+ parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank")
+ parser.add_argument(
+ "--checkpointing_steps",
+ type=int,
+ default=500,
+ help=(
+ "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming"
+ " training using `--resume_from_checkpoint`."
+ ),
+ )
+ parser.add_argument(
+ "--checkpoints_total_limit",
+ type=int,
+ default=None,
+ help=("Max number of checkpoints to store."),
+ )
+ parser.add_argument(
+ "--resume_from_checkpoint",
+ type=str,
+ default=None,
+ help=(
+ "Whether training should be resumed from a previous checkpoint. Use a path saved by"
+ ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
+ ),
+ )
+ parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.")
+ parser.add_argument(
+ "--validation_epochs",
+ type=int,
+ default=5,
+ help="Run validation every X epochs.",
+ )
+ parser.add_argument(
+ "--validation_steps",
+ type=int,
+ default=2000,
+ help="Run validation every X steps.",
+ )
+ parser.add_argument(
+ "--tracker_project_name",
+ type=str,
+ default="text2image-fine-tune",
+ help=(
+ "The `project_name` argument passed to Accelerator.init_trackers for"
+ " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator"
+ ),
+ )
+
+ parser.add_argument(
+ "--rank",
+ type=int,
+ default=128,
+ help=("The dimension of the LoRA update matrices."),
+ )
+ parser.add_argument(
+ "--network_alpha",
+ type=int,
+ default=64,
+ help=("The dimension of the LoRA update matrices."),
+ )
+ parser.add_argument(
+ "--train_text_encoder",
+ action="store_true",
+ help="Whether to train the text encoder. If set, the text encoder should be float32 precision.",
+ )
+ parser.add_argument(
+ "--snr_loss", action="store_true", help="Whether or not to use snr_loss."
+ )
+ parser.add_argument(
+ "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling."
+ )
+ parser.add_argument(
+ "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader."
+ )
+ parser.add_argument(
+ "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets."
+ )
+ parser.add_argument(
+ "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets."
+ )
+ parser.add_argument(
+ "--random_frame_crop", action="store_true", help="Whether enable random frame crop sample in datasets."
+ )
+ parser.add_argument(
+ "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets."
+ )
+ parser.add_argument(
+ "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.",
+ )
+ 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(
+ "--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,
+ default=512,
+ help="Sample size of the token.",
+ )
+ parser.add_argument(
+ "--video_sample_size",
+ type=int,
+ default=512,
+ help="Sample size of the video.",
+ )
+ parser.add_argument(
+ "--image_sample_size",
+ type=int,
+ default=512,
+ help="Sample size of the video.",
+ )
+ parser.add_argument(
+ "--video_sample_stride",
+ type=int,
+ default=4,
+ help="Sample stride of the video.",
+ )
+ parser.add_argument(
+ "--video_sample_n_frames",
+ type=int,
+ default=17,
+ help="Num frame of video.",
+ )
+ parser.add_argument(
+ "--video_repeat",
+ type=int,
+ default=0,
+ help="Num of repeat video.",
+ )
+ parser.add_argument(
+ "--config_path",
+ type=str,
+ default=None,
+ help=(
+ "The config of the model in training."
+ ),
+ )
+ parser.add_argument(
+ "--transformer_path",
+ type=str,
+ default=None,
+ help=("If you want to load the weight from other transformers, input its path."),
+ )
+ parser.add_argument(
+ "--vae_path",
+ type=str,
+ default=None,
+ help=("If you want to load the weight from other vaes, input its path."),
+ )
+ parser.add_argument("--save_state", action="store_true", help="Whether or not to save state.")
+
+ parser.add_argument(
+ '--tokenizer_max_length',
+ type=int,
+ default=512,
+ help='Max length of tokenizer'
+ )
+ parser.add_argument(
+ "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed."
+ )
+ parser.add_argument(
+ "--low_vram", action="store_true", help="Whether enable low_vram mode."
+ )
+ parser.add_argument(
+ "--train_mode",
+ type=str,
+ default="control",
+ help=(
+ 'The format of training data. Support `"control"`'
+ ' (default), `"control_ref"`.'
+ ),
+ )
+ parser.add_argument(
+ "--control_ref_image",
+ type=str,
+ default="first_frame",
+ help=(
+ 'The format of training data. Support `"first_frame"`'
+ ' (default), `"random"`.'
+ ),
+ )
+ parser.add_argument(
+ "--weighting_scheme",
+ type=str,
+ default="none",
+ choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"],
+ help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'),
+ )
+ parser.add_argument(
+ "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme."
+ )
+ parser.add_argument(
+ "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme."
+ )
+ parser.add_argument(
+ "--mode_scale",
+ type=float,
+ 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))
+ if env_local_rank != -1 and env_local_rank != args.local_rank:
+ args.local_rank = env_local_rank
+
+ # default to using the same revision for the non-ema model if not specified
+ if args.non_ema_revision is None:
+ args.non_ema_revision = args.revision
+
+ return args
+
+
+def main():
+ args = parse_args()
+
+ if args.report_to == "wandb" and args.hub_token is not None:
+ raise ValueError(
+ "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token."
+ " Please use `huggingface-cli login` to authenticate with the Hub."
+ )
+
+ if args.non_ema_revision is not None:
+ deprecate(
+ "non_ema_revision!=None",
+ "0.15.0",
+ message=(
+ "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to"
+ " use `--variant=non_ema` instead."
+ ),
+ )
+ logging_dir = os.path.join(args.output_dir, args.logging_dir)
+
+ config = OmegaConf.load(args.config_path)
+ accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir)
+
+ accelerator = Accelerator(
+ gradient_accumulation_steps=args.gradient_accumulation_steps,
+ mixed_precision=args.mixed_precision,
+ log_with=args.report_to,
+ project_config=accelerator_project_config,
+ )
+ deepspeed_plugin = accelerator.state.deepspeed_plugin
+ if deepspeed_plugin is not None:
+ zero_stage = int(deepspeed_plugin.zero_stage)
+ print(f"Using DeepSpeed Zero stage: {zero_stage}")
+ 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)
+
+ # Make one log on every process with the configuration for debugging.
+ logging.basicConfig(
+ format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
+ datefmt="%m/%d/%Y %H:%M:%S",
+ level=logging.INFO,
+ )
+ logger.info(accelerator.state, main_process_only=False)
+ if accelerator.is_local_main_process:
+ datasets.utils.logging.set_verbosity_warning()
+ transformers.utils.logging.set_verbosity_warning()
+ diffusers.utils.logging.set_verbosity_info()
+ else:
+ datasets.utils.logging.set_verbosity_error()
+ transformers.utils.logging.set_verbosity_error()
+ diffusers.utils.logging.set_verbosity_error()
+
+ # If passed along, set the training seed now.
+ if args.seed is not None:
+ set_seed(args.seed)
+ rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index))
+ torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index)
+ else:
+ rng = None
+ torch_rng = None
+ index_rng = np.random.default_rng(np.random.PCG64(43))
+ print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}")
+
+ # Handle the repository creation
+ if accelerator.is_main_process:
+ if args.output_dir is not None:
+ os.makedirs(args.output_dir, exist_ok=True)
+
+ # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision
+ # as these weights are only used for inference, keeping weights in full precision is not required.
+ weight_dtype = torch.float32
+ if accelerator.mixed_precision == "fp16":
+ weight_dtype = torch.float16
+ args.mixed_precision = accelerator.mixed_precision
+ elif accelerator.mixed_precision == "bf16":
+ weight_dtype = torch.bfloat16
+ args.mixed_precision = accelerator.mixed_precision
+
+ # Load scheduler, tokenizer and models.
+ noise_scheduler = FlowMatchEulerDiscreteScheduler(
+ **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
+ )
+
+ # Get Tokenizer
+ tokenizer = AutoTokenizer.from_pretrained(
+ os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
+ )
+
+ def deepspeed_zero_init_disabled_context_manager():
+ """
+ returns either a context list that includes one that will disable zero.Init or an empty context list
+ """
+ deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None
+ if deepspeed_plugin is None:
+ return []
+
+ return [deepspeed_plugin.zero3_init_context_manager(enable=False)]
+
+ # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3.
+ # For this to work properly all models must be run through `accelerate.prepare`. But accelerate
+ # will try to assign the same optimizer with the same weights to all models during
+ # `deepspeed.initialize`, which of course doesn't work.
+ #
+ # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2
+ # frozen models from being partitioned during `zero.Init` which gets called during
+ # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding
+ # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded.
+ with ContextManagers(deepspeed_zero_init_disabled_context_manager()):
+ # Get Text encoder
+ text_encoder = WanT5EncoderModel.from_pretrained(
+ os.path.join(args.pretrained_model_name_or_path, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
+ additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
+ 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')),
+ additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
+ )
+
+ # Get Transformer
+ transformer3d = WanTransformer3DModel.from_pretrained(
+ os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
+ ).to(weight_dtype)
+
+ if args.train_mode != "normal":
+ # Get Clip Image Encoder
+ clip_image_encoder = CLIPModel.from_pretrained(
+ os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
+ )
+ clip_image_encoder = clip_image_encoder.eval()
+
+ # Freeze vae and text_encoder and set transformer3d to trainable
+ vae.requires_grad_(False)
+ text_encoder.requires_grad_(False)
+ transformer3d.requires_grad_(False)
+ if args.train_mode != "normal":
+ clip_image_encoder.requires_grad_(False)
+
+ # Lora will work with this...
+ network = create_network(
+ 1.0,
+ args.rank,
+ args.network_alpha,
+ text_encoder,
+ transformer3d,
+ neuron_dropout=None,
+ 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)
+
+ if args.transformer_path is not None:
+ print(f"From checkpoint: {args.transformer_path}")
+ if args.transformer_path.endswith("safetensors"):
+ from safetensors.torch import load_file, safe_open
+ state_dict = load_file(args.transformer_path)
+ else:
+ state_dict = torch.load(args.transformer_path, map_location="cpu")
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
+
+ m, u = transformer3d.load_state_dict(state_dict, strict=False)
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
+ assert len(u) == 0
+
+ if args.vae_path is not None:
+ print(f"From checkpoint: {args.vae_path}")
+ if args.vae_path.endswith("safetensors"):
+ from safetensors.torch import load_file, safe_open
+ state_dict = load_file(args.vae_path)
+ else:
+ state_dict = torch.load(args.vae_path, map_location="cpu")
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
+
+ m, u = vae.load_state_dict(state_dict, strict=False)
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
+ assert len(u) == 0
+
+ # `accelerate` 0.16.0 will have better support for customized saving
+ if version.parse(accelerate.__version__) >= version.parse("0.16.0"):
+ # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
+ if zero_stage != 3:
+ def save_model_hook(models, weights, output_dir):
+ if accelerator.is_main_process:
+ safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
+ save_model(safetensor_save_path, accelerator.unwrap_model(models[-1]))
+ if not args.use_deepspeed:
+ for _ in range(len(weights)):
+ weights.pop()
+
+ with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
+ pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file)
+
+ def load_model_hook(models, input_dir):
+ pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl")
+ if os.path.exists(pkl_path):
+ with open(pkl_path, 'rb') as file:
+ loaded_number, _ = pickle.load(file)
+ batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0)
+ print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.")
+ else:
+ def save_model_hook(models, weights, output_dir):
+ if accelerator.is_main_process:
+ with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
+ pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file)
+
+ def load_model_hook(models, input_dir):
+ pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl")
+ if os.path.exists(pkl_path):
+ with open(pkl_path, 'rb') as file:
+ loaded_number, _ = pickle.load(file)
+ batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0)
+ print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.")
+
+
+ accelerator.register_save_state_pre_hook(save_model_hook)
+ accelerator.register_load_state_pre_hook(load_model_hook)
+
+ if args.gradient_checkpointing:
+ transformer3d.enable_gradient_checkpointing()
+
+ # Enable TF32 for faster training on Ampere GPUs,
+ # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
+ if args.allow_tf32:
+ torch.backends.cuda.matmul.allow_tf32 = True
+
+ if args.scale_lr:
+ args.learning_rate = (
+ args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes
+ )
+
+ # Initialize the optimizer
+ if args.use_8bit_adam:
+ try:
+ import bitsandbytes as bnb
+ except ImportError:
+ raise ImportError(
+ "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`"
+ )
+
+ optimizer_cls = bnb.optim.AdamW8bit
+ elif args.use_came:
+ try:
+ from came_pytorch import CAME
+ except:
+ raise ImportError(
+ "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`"
+ )
+
+ optimizer_cls = CAME
+ else:
+ optimizer_cls = torch.optim.AdamW
+
+ logging.info("Add network parameters")
+ trainable_params = list(filter(lambda p: p.requires_grad, network.parameters()))
+ trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate)
+
+ if args.use_came:
+ optimizer = optimizer_cls(
+ trainable_params_optim,
+ lr=args.learning_rate,
+ # weight_decay=args.adam_weight_decay,
+ betas=(0.9, 0.999, 0.9999),
+ eps=(1e-30, 1e-16)
+ )
+ else:
+ optimizer = optimizer_cls(
+ trainable_params_optim,
+ lr=args.learning_rate,
+ betas=(args.adam_beta1, args.adam_beta2),
+ weight_decay=args.adam_weight_decay,
+ eps=args.adam_epsilon,
+ )
+
+ # Get the training dataset
+ sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
+
+ # Get the dataset
+ train_dataset = ImageVideoControlDataset(
+ args.train_data_meta, args.train_data_dir,
+ video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
+ video_repeat=args.video_repeat,
+ 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()}
+ batch_sampler_generator = torch.Generator().manual_seed(args.seed)
+ batch_sampler = AspectRatioBatchImageVideoSampler(
+ sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
+ batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True,
+ aspect_ratios=aspect_ratio_sample_size,
+ )
+
+ 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)
+
+ # Create new output
+ new_examples = {}
+ new_examples["target_token_length"] = target_token_length
+ new_examples["pixel_values"] = []
+ new_examples["text"] = []
+ # Used in Control Mode
+ new_examples["control_pixel_values"] = []
+ # Used in Control Ref Mode
+ if args.train_mode != "control":
+ new_examples["ref_pixel_values"] = []
+ new_examples["clip_pixel_values"] = []
+
+ # Get downsample ratio in image and videos
+ pixel_value = examples[0]["pixel_values"]
+ 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])
+
+ 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()}
+
+ batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
+ else:
+ if args.random_hw_adapt:
+ if args.training_with_video_token_length:
+ local_min_size = np.min(np.array([np.mean(np.array([np.shape(example["pixel_values"])[1], np.shape(example["pixel_values"])[2]])) for example in examples]))
+ # The video will be resized to a lower resolution than its own.
+ 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())
+ 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)
+ batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
+ else:
+ random_downsample_ratio = 1
+ batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
+
+ aspect_ratio_sample_size = {key : [x / 512 * args.video_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.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
+
+ 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:
+ 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:
+ # To 0~1
+ pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
+ pixel_values = pixel_values / 255.
+
+ control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous()
+ control_pixel_values = control_pixel_values / 255.
+
+ if args.random_ratio_crop:
+ # Get adapt hw for resize
+ b, c, h, w = pixel_values.size()
+ th, tw = random_sample_size
+ if th / tw > h / w:
+ nh = int(th)
+ nw = int(w / h * nh)
+ else:
+ nw = int(tw)
+ nh = int(h / w * nw)
+
+ transform = transforms.Compose([
+ transforms.Resize([nh, nw]),
+ transforms.CenterCrop([int(x) for x in random_sample_size]),
+ transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
+ ])
+
+ else:
+ # Get adapt hw for resize
+ closest_size = list(map(lambda x: int(x), closest_size))
+ if closest_size[0] / h > closest_size[1] / w:
+ resize_size = closest_size[0], int(w * closest_size[0] / h)
+ else:
+ resize_size = int(h * closest_size[1] / w), closest_size[1]
+
+ transform = transforms.Compose([
+ transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
+ transforms.CenterCrop(closest_size),
+ transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
+ ])
+
+ new_examples["pixel_values"].append(transform(pixel_values))
+ new_examples["control_pixel_values"].append(transform(control_pixel_values))
+
+ new_examples["text"].append(example["text"])
+ # Magvae needs the number of frames to be 4n + 1.
+ batch_video_length = int(
+ min(
+ batch_video_length,
+ (len(pixel_values) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1,
+ )
+ )
+ if batch_video_length == 0:
+ batch_video_length = 1
+
+ if args.train_mode != "control":
+ if args.control_ref_image == "first_frame":
+ clip_index = 0
+ else:
+ def _create_special_list(length):
+ if length == 1:
+ return [1.0]
+ if length >= 2:
+ first_element = 0.40
+ 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
+ number_list_prob = np.array(_create_special_list(len(new_examples["pixel_values"][-1])))
+ clip_index = np.random.choice(list(range(len(new_examples["pixel_values"][-1]))), p = number_list_prob)
+
+ ref_pixel_values = new_examples["pixel_values"][-1][clip_index].unsqueeze(0)
+ new_examples["ref_pixel_values"].append(ref_pixel_values)
+
+ clip_pixel_values = new_examples["pixel_values"][-1][clip_index].permute(1, 2, 0).contiguous()
+ clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255
+ new_examples["clip_pixel_values"].append(clip_pixel_values)
+
+ # Limit the number of frames to the same
+ new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]])
+ new_examples["control_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_pixel_values"]])
+ if args.train_mode != "control":
+ new_examples["ref_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["ref_pixel_values"]])
+ new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]])
+
+ # Encode prompts when enable_text_encoder_in_dataloader=True
+ if args.enable_text_encoder_in_dataloader:
+ prompt_ids = tokenizer(
+ new_examples['text'],
+ max_length=args.tokenizer_max_length,
+ padding="max_length",
+ add_special_tokens=True,
+ truncation=True,
+ return_tensors="pt"
+ )
+ encoder_hidden_states = text_encoder(
+ prompt_ids.input_ids
+ )[0]
+ new_examples['encoder_attention_mask'] = prompt_ids.attention_mask
+ new_examples['encoder_hidden_states'] = encoder_hidden_states
+
+ return new_examples
+
+ # DataLoaders creation:
+ train_dataloader = torch.utils.data.DataLoader(
+ train_dataset,
+ batch_sampler=batch_sampler,
+ 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:
+ batch_sampler_generator = torch.Generator().manual_seed(args.seed)
+ batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
+ train_dataloader = torch.utils.data.DataLoader(
+ train_dataset,
+ 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.
+ overrode_max_train_steps = False
+ num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+ if args.max_train_steps is None:
+ args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
+ overrode_max_train_steps = True
+
+ lr_scheduler = get_scheduler(
+ args.lr_scheduler,
+ optimizer=optimizer,
+ num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
+ num_training_steps=args.max_train_steps * accelerator.num_processes,
+ )
+
+ # Prepare everything with our `accelerator`.
+ 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)
+ if not args.enable_text_encoder_in_dataloader:
+ text_encoder.to(accelerator.device)
+ if args.train_mode != "normal":
+ clip_image_encoder.to(accelerator.device, dtype=weight_dtype)
+
+ # We need to recalculate our total training steps as the size of the training dataloader may have changed.
+ num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+ if overrode_max_train_steps:
+ args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
+ # Afterwards we recalculate our number of training epochs
+ args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
+
+ # We need to initialize the trackers we use, and also store our configuration.
+ # The trackers initializes automatically on the main process.
+ if accelerator.is_main_process:
+ tracker_config = dict(vars(args))
+ tracker_config.pop("validation_prompts")
+ accelerator.init_trackers(args.tracker_project_name, tracker_config)
+
+ # Function for unwrapping if model was compiled with `torch.compile`.
+ def unwrap_model(model):
+ model = accelerator.unwrap_model(model)
+ model = model._orig_mod if is_compiled_module(model) else model
+ return model
+
+ # Train!
+ total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
+
+ logger.info("***** Running training *****")
+ logger.info(f" Num examples = {len(train_dataset)}")
+ logger.info(f" Num Epochs = {args.num_train_epochs}")
+ logger.info(f" Instantaneous batch size per device = {args.train_batch_size}")
+ logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
+ logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
+ logger.info(f" Total optimization steps = {args.max_train_steps}")
+ global_step = 0
+ first_epoch = 0
+
+ # Potentially load in the weights and states from a previous save
+ if args.resume_from_checkpoint:
+ if args.resume_from_checkpoint != "latest":
+ path = os.path.basename(args.resume_from_checkpoint)
+ else:
+ # Get the most recent checkpoint
+ dirs = os.listdir(args.output_dir)
+ dirs = [d for d in dirs if d.startswith("checkpoint")]
+ dirs = sorted(dirs, key=lambda x: int(x.split("-")[1]))
+ path = dirs[-1] if len(dirs) > 0 else None
+
+ if path is None:
+ accelerator.print(
+ f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run."
+ )
+ args.resume_from_checkpoint = None
+ initial_global_step = 0
+ else:
+ global_step = int(path.split("-")[1])
+
+ initial_global_step = global_step
+
+ pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl")
+ if os.path.exists(pkl_path):
+ with open(pkl_path, 'rb') as file:
+ _, first_epoch = pickle.load(file)
+ else:
+ first_epoch = global_step // num_update_steps_per_epoch
+ print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
+
+ 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))
+ else:
+ initial_global_step = 0
+
+ # function for saving/removing
+ def save_model(ckpt_file, unwrapped_nw):
+ os.makedirs(args.output_dir, exist_ok=True)
+ accelerator.print(f"\nsaving checkpoint: {ckpt_file}")
+ unwrapped_nw.save_weights(ckpt_file, weight_dtype, None)
+
+ progress_bar = tqdm(
+ range(0, args.max_train_steps),
+ initial=initial_global_step,
+ desc="Steps",
+ # Only show the progress bar once on each machine.
+ disable=not accelerator.is_local_main_process,
+ )
+
+ if args.multi_stream and args.train_mode != "normal":
+ # create extra cuda streams to speedup inpaint vae computation
+ vae_stream_1 = torch.cuda.Stream()
+ else:
+ vae_stream_1 = None
+
+ idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling)
+
+ for epoch in range(first_epoch, args.num_train_epochs):
+ train_loss = 0.0
+ batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch)
+ for step, batch in enumerate(train_dataloader):
+ # Data batch sanity check
+ if epoch == first_epoch and step == 0:
+ pixel_values, texts = batch['pixel_values'].cpu(), batch['text']
+ control_pixel_values = batch["control_pixel_values"].cpu()
+ pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w")
+ control_pixel_values = rearrange(control_pixel_values, "b f c h w -> b c f h w")
+ os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True)
+ for idx, (pixel_value, control_pixel_value, text) in enumerate(zip(pixel_values, control_pixel_values, texts)):
+ pixel_value = pixel_value[None, ...]
+ control_pixel_value = control_pixel_value[None, ...]
+ gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}'
+ save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True)
+ save_videos_grid(control_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_control.gif", rescale=True)
+
+ if args.train_mode != "control":
+ ref_pixel_values = batch["ref_pixel_values"].cpu()
+ ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> b c f h w")
+ for idx, (ref_pixel_value, text) in enumerate(zip(ref_pixel_values, texts)):
+ ref_pixel_value = ref_pixel_value[None, ...]
+ gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}'
+ save_videos_grid(ref_pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}_ref.gif", rescale=True)
+
+ with accelerator.accumulate(transformer3d):
+ # Convert images to latent space
+ pixel_values = batch["pixel_values"].to(weight_dtype)
+ 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 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))
+ if args.enable_text_encoder_in_dataloader:
+ batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (4, 1, 1))
+ batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (4, 1))
+ else:
+ batch['text'] = batch['text'] * 4
+ elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
+ pixel_values = torch.tile(pixel_values, (2, 1, 1, 1, 1))
+ control_pixel_values = torch.tile(control_pixel_values, (2, 1, 1, 1, 1))
+ if args.enable_text_encoder_in_dataloader:
+ batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (2, 1, 1))
+ batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (2, 1))
+ else:
+ batch['text'] = batch['text'] * 2
+
+ if args.train_mode != "control":
+ 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 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))
+ elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
+ clip_pixel_values = torch.tile(clip_pixel_values, (2, 1, 1, 1))
+ ref_pixel_values = torch.tile(ref_pixel_values, (2, 1, 1, 1, 1))
+
+ if args.random_frame_crop:
+ def _create_special_list(length):
+ if length == 1:
+ return [1.0]
+ if length >= 2:
+ last_element = 0.90
+ remaining_sum = 1.0 - last_element
+ other_elements_value = remaining_sum / (length - 1)
+ special_list = [other_elements_value] * (length - 1) + [last_element]
+ return special_list
+ select_frames = [_tmp for _tmp in list(range(sample_n_frames_bucket_interval + 1, args.video_sample_n_frames + sample_n_frames_bucket_interval, sample_n_frames_bucket_interval))]
+ select_frames_prob = np.array(_create_special_list(len(select_frames)))
+
+ if len(select_frames) != 0:
+ if rng is None:
+ temp_n_frames = np.random.choice(select_frames, p = select_frames_prob)
+ else:
+ temp_n_frames = rng.choice(select_frames, p = select_frames_prob)
+ else:
+ temp_n_frames = 1
+
+ # Magvae needs the number of frames to be 4n + 1.
+ temp_n_frames = (temp_n_frames - 1) // sample_n_frames_bucket_interval + 1
+
+ pixel_values = pixel_values[:, :temp_n_frames, :, :]
+ control_pixel_values = control_pixel_values[:, :temp_n_frames, :, :]
+
+ # Keep all node same token length to accelerate the traning when resolution grows.
+ if args.keep_all_node_same_token_length:
+ if args.token_sample_size > 256:
+ numbers_list = list(range(256, args.token_sample_size + 1, 128))
+
+ if numbers_list[-1] != args.token_sample_size:
+ numbers_list.append(args.token_sample_size)
+ else:
+ numbers_list = [256]
+ numbers_list = [_number * _number * args.video_sample_n_frames for _number in numbers_list]
+
+ actual_token_length = index_rng.choice(numbers_list)
+ actual_video_length = (min(
+ actual_token_length / pixel_values.size()[-1] / pixel_values.size()[-2], args.video_sample_n_frames
+ ) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1
+ actual_video_length = int(max(actual_video_length, 1))
+
+ # Magvae needs the number of frames to be 4n + 1.
+ actual_video_length = (actual_video_length - 1) // sample_n_frames_bucket_interval + 1
+
+ pixel_values = pixel_values[:, :actual_video_length, :, :]
+ control_pixel_values = control_pixel_values[:, :actual_video_length, :, :]
+
+ if args.low_vram:
+ torch.cuda.empty_cache()
+ vae.to(accelerator.device)
+ if args.train_mode != "normal":
+ clip_image_encoder.to(accelerator.device)
+ if not args.enable_text_encoder_in_dataloader:
+ text_encoder.to("cpu")
+
+ with torch.no_grad():
+ # This way is quicker when batch grows up
+ def _batch_encode_vae(pixel_values):
+ pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w")
+ bs = args.vae_mini_batch
+ new_pixel_values = []
+ for i in range(0, pixel_values.shape[0], bs):
+ pixel_values_bs = pixel_values[i : i + bs]
+ pixel_values_bs = vae.encode(pixel_values_bs)[0]
+ pixel_values_bs = pixel_values_bs.sample()
+ new_pixel_values.append(pixel_values_bs)
+ return torch.cat(new_pixel_values, dim = 0)
+ if vae_stream_1 is not None:
+ vae_stream_1.wait_stream(torch.cuda.current_stream())
+ with torch.cuda.stream(vae_stream_1):
+ latents = _batch_encode_vae(pixel_values)
+ else:
+ latents = _batch_encode_vae(pixel_values)
+
+ control_latents = _batch_encode_vae(control_pixel_values)
+ # Make control latents to zero
+ for bs_index in range(control_latents.size()[0]):
+ if rng is None:
+ zero_init_control_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10])
+ else:
+ zero_init_control_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10])
+
+ if zero_init_control_latents_conv_in:
+ control_latents[bs_index] = control_latents[bs_index] * 0
+
+ if args.train_mode != "control":
+ ref_latents = _batch_encode_vae(ref_pixel_values)
+
+ ref_latents_conv_in = torch.zeros_like(latents).to(ref_latents.device, ref_latents.dtype)
+ ref_latents_conv_in[:, :, :1] = ref_latents
+ for bs_index in range(ref_latents.size()[0]):
+ if rng is None:
+ zero_init_ref_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10])
+ else:
+ zero_init_ref_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10])
+
+ if zero_init_ref_latents_conv_in and control_latents.size()[1] != 1:
+ ref_latents_conv_in[bs_index, :, :1] = ref_latents_conv_in[bs_index, :, :1] * 0
+
+ control_latents = torch.cat([control_latents, ref_latents_conv_in], dim = 1)
+
+ clip_context = []
+ for clip_pixel_value in clip_pixel_values:
+ clip_image = Image.fromarray(np.uint8(clip_pixel_value.float().cpu().numpy()))
+ clip_image = TF.to_tensor(clip_image).sub_(0.5).div_(0.5).to(clip_image_encoder.device, weight_dtype)
+ _clip_context = clip_image_encoder([clip_image[:, None, :, :]])
+
+ if rng is None:
+ zero_init_clip_in = np.random.choice([True, False], p=[0.1, 0.9])
+ else:
+ zero_init_clip_in = rng.choice([True, False], p=[0.1, 0.9])
+ clip_context.append(_clip_context if not zero_init_clip_in else torch.zeros_like(_clip_context))
+
+ clip_context = torch.cat(clip_context)
+
+ # wait for latents = vae.encode(pixel_values) to complete
+ if vae_stream_1 is not None:
+ torch.cuda.current_stream().wait_stream(vae_stream_1)
+
+ if args.low_vram:
+ vae.to('cpu')
+ if args.train_mode != "normal":
+ clip_image_encoder.to('cpu')
+ torch.cuda.empty_cache()
+ if not args.enable_text_encoder_in_dataloader:
+ text_encoder.to(accelerator.device)
+
+ if args.enable_text_encoder_in_dataloader:
+ prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
+ else:
+ with torch.no_grad():
+ prompt_ids = tokenizer(
+ batch['text'],
+ padding="max_length",
+ max_length=args.tokenizer_max_length,
+ truncation=True,
+ add_special_tokens=True,
+ return_tensors="pt"
+ )
+ text_input_ids = prompt_ids.input_ids
+ prompt_attention_mask = prompt_ids.attention_mask
+
+ seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
+ prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
+ prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
+
+ if args.low_vram and not args.enable_text_encoder_in_dataloader:
+ text_encoder.to('cpu')
+ torch.cuda.empty_cache()
+
+ bsz, channel, num_frames, height, width = latents.size()
+ noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype)
+
+ if not args.uniform_sampling:
+ u = compute_density_for_timestep_sampling(
+ weighting_scheme=args.weighting_scheme,
+ batch_size=bsz,
+ logit_mean=args.logit_mean,
+ logit_std=args.logit_std,
+ mode_scale=args.mode_scale,
+ )
+ indices = (u * noise_scheduler.config.num_train_timesteps).long()
+ else:
+ # Sample a random timestep for each image
+ # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng)
+ # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng)
+ indices = idx_sampling(bsz, generator=torch_rng, device=latents.device)
+ indices = indices.long().cpu()
+ timesteps = noise_scheduler.timesteps[indices].to(device=latents.device)
+
+ def get_sigmas(timesteps, n_dim=4, dtype=torch.float32):
+ sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype)
+ schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device)
+ timesteps = timesteps.to(accelerator.device)
+ step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
+
+ sigma = sigmas[step_indices].flatten()
+ while len(sigma.shape) < n_dim:
+ sigma = sigma.unsqueeze(-1)
+ return sigma
+
+ # Add noise according to flow matching.
+ # zt = (1 - texp) * x + texp * z1
+ sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype)
+ noisy_latents = (1.0 - sigmas) * latents + sigmas * noise
+
+ # Add noise
+ target = noise - latents
+
+ target_shape = (vae.latent_channels, num_frames, width, height)
+ seq_len = math.ceil(
+ (target_shape[2] * target_shape[3]) /
+ (accelerator.unwrap_model(transformer3d).config.patch_size[1] * accelerator.unwrap_model(transformer3d).config.patch_size[2]) *
+ target_shape[1]
+ )
+
+ # Predict the noise residual
+ with torch.cuda.amp.autocast(dtype=weight_dtype):
+ noise_pred = transformer3d(
+ x=noisy_latents,
+ context=prompt_embeds,
+ t=timesteps,
+ seq_len=seq_len,
+ y=control_latents if args.train_mode != "normal" else None,
+ clip_fea=clip_context if args.train_mode != "normal" else None,
+ )
+
+ def custom_mse_loss(noise_pred, target, weighting=None, threshold=50):
+ noise_pred = noise_pred.float()
+ target = target.float()
+ diff = noise_pred - target
+ mse_loss = F.mse_loss(noise_pred, target, reduction='none')
+ mask = (diff.abs() <= threshold).float()
+ masked_loss = mse_loss * mask
+ if weighting is not None:
+ masked_loss = masked_loss * weighting
+ final_loss = masked_loss.mean()
+ return final_loss
+
+ weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas)
+ loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float())
+ loss = loss.mean()
+
+ if args.motion_sub_loss and noise_pred.size()[1] > 2:
+ gt_sub_noise = noise_pred[:, 1:, :].float() - noise_pred[:, :-1, :].float()
+ pre_sub_noise = target[:, 1:, :].float() - target[:, :-1, :].float()
+ sub_loss = F.mse_loss(gt_sub_noise, pre_sub_noise, reduction="mean")
+ loss = loss * (1 - args.motion_sub_loss_ratio) + sub_loss * args.motion_sub_loss_ratio
+
+ # Gather the losses across all processes for logging (if we use distributed training).
+ avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean()
+ train_loss += avg_loss.item() / args.gradient_accumulation_steps
+
+ # Backpropagate
+ accelerator.backward(loss)
+ if accelerator.sync_gradients:
+ accelerator.clip_grad_norm_(trainable_params, args.max_grad_norm)
+ optimizer.step()
+ lr_scheduler.step()
+ optimizer.zero_grad()
+
+ # Checks if the accelerator has performed an optimization step behind the scenes
+ if accelerator.sync_gradients:
+ progress_bar.update(1)
+ global_step += 1
+ accelerator.log({"train_loss": train_loss}, step=global_step)
+ train_loss = 0.0
+
+ if global_step % args.checkpointing_steps == 0:
+ if args.use_deepspeed or accelerator.is_main_process:
+ # _before_ saving state, check if this save would set us over the `checkpoints_total_limit`
+ if args.checkpoints_total_limit is not None:
+ checkpoints = os.listdir(args.output_dir)
+ checkpoints = [d for d in checkpoints if d.startswith("checkpoint")]
+ checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))
+
+ # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints
+ if len(checkpoints) >= args.checkpoints_total_limit:
+ num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1
+ removing_checkpoints = checkpoints[0:num_to_remove]
+
+ logger.info(
+ f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints"
+ )
+ logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}")
+
+ for removing_checkpoint in removing_checkpoints:
+ removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint)
+ shutil.rmtree(removing_checkpoint)
+ if not args.save_state:
+ safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
+ save_model(safetensor_save_path, accelerator.unwrap_model(network))
+ logger.info(f"Saved safetensor to {safetensor_save_path}")
+ else:
+ accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
+ accelerator.save_state(accelerator_save_path)
+ logger.info(f"Saved state to {accelerator_save_path}")
+
+ if accelerator.is_main_process:
+ if args.validation_prompts is not None and global_step % args.validation_steps == 0:
+ log_validation(
+ vae,
+ text_encoder,
+ tokenizer,
+ clip_image_encoder,
+ transformer3d,
+ network,
+ config,
+ args,
+ accelerator,
+ weight_dtype,
+ global_step,
+ )
+
+ logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]}
+ progress_bar.set_postfix(**logs)
+
+ if global_step >= args.max_train_steps:
+ break
+
+ if accelerator.is_main_process:
+ if args.validation_prompts is not None and epoch % args.validation_epochs == 0:
+ log_validation(
+ vae,
+ text_encoder,
+ tokenizer,
+ clip_image_encoder,
+ transformer3d,
+ network,
+ config,
+ args,
+ accelerator,
+ weight_dtype,
+ global_step,
+ )
+
+ # Create the pipeline using the trained modules and save it.
+ accelerator.wait_for_everyone()
+ if accelerator.is_main_process:
+ safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
+ accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
+ save_model(safetensor_save_path, accelerator.unwrap_model(network))
+ if args.save_state:
+ accelerator.save_state(accelerator_save_path)
+ logger.info(f"Saved state to {accelerator_save_path}")
+
+ accelerator.end_training()
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/wan2.1_fun/train_control_lora.sh b/scripts/wan2.1_fun/train_control_lora.sh
new file mode 100644
index 0000000..b6392e3
--- /dev/null
+++ b/scripts/wan2.1_fun/train_control_lora.sh
@@ -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
\ No newline at end of file
diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py
index 4bd5b08..16ace98 100755
--- a/scripts/wan2.1_fun/train_lora.py
+++ b/scripts/wan2.1_fun/train_lora.py
@@ -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))
diff --git a/videox_fun/models/cogvideox_transformer3d.py b/videox_fun/models/cogvideox_transformer3d.py
old mode 100644
new mode 100755
index 0887b42..e12d479
--- a/videox_fun/models/cogvideox_transformer3d.py
+++ b/videox_fun/models/cogvideox_transformer3d.py
@@ -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):
diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py
index e9f7445..b67091a 100755
--- a/videox_fun/models/wan_transformer3d.py
+++ b/videox_fun/models/wan_transformer3d.py
@@ -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):
diff --git a/videox_fun/pipeline/pipeline_wan_fun.py b/videox_fun/pipeline/pipeline_wan_fun.py
old mode 100644
new mode 100755
index 225aca2..4a6317d
--- a/videox_fun/pipeline/pipeline_wan_fun.py
+++ b/videox_fun/pipeline/pipeline_wan_fun.py
@@ -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()
diff --git a/videox_fun/pipeline/pipeline_wan_fun_control.py b/videox_fun/pipeline/pipeline_wan_fun_control.py
old mode 100644
new mode 100755
index bbe76ca..6146bea
--- a/videox_fun/pipeline/pipeline_wan_fun_control.py
+++ b/videox_fun/pipeline/pipeline_wan_fun_control.py
@@ -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()
diff --git a/videox_fun/pipeline/pipeline_wan_fun_inpaint.py b/videox_fun/pipeline/pipeline_wan_fun_inpaint.py
old mode 100644
new mode 100755
index 9f21f80..bf0e95b
--- a/videox_fun/pipeline/pipeline_wan_fun_inpaint.py
+++ b/videox_fun/pipeline/pipeline_wan_fun_inpaint.py
@@ -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:
diff --git a/videox_fun/ui/controller.py b/videox_fun/ui/controller.py
index ff77d31..d798b47 100755
--- a/videox_fun/ui/controller.py
+++ b/videox_fun/ui/controller.py
@@ -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)
diff --git a/videox_fun/ui/wan_fun_ui.py b/videox_fun/ui/wan_fun_ui.py
old mode 100644
new mode 100755
index 6f7b629..61acb0e
--- a/videox_fun/ui/wan_fun_ui.py
+++ b/videox_fun/ui/wan_fun_ui.py
@@ -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:
diff --git a/videox_fun/ui/wan_ui.py b/videox_fun/ui/wan_ui.py
old mode 100644
new mode 100755
index 621726e..c59750a
--- a/videox_fun/ui/wan_ui.py
+++ b/videox_fun/ui/wan_ui.py
@@ -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:
diff --git a/videox_fun/utils/lora_utils.py b/videox_fun/utils/lora_utils.py
index ecd37d1..f801fa8 100755
--- a/videox_fun/utils/lora_utils.py
+++ b/videox_fun/utils/lora_utils.py
@@ -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