Update Riflex && Pipeline callback bug fix && Training Code bug fix (#150)

This commit is contained in:
Bubbliiiing
2025-04-09 15:23:38 +08:00
committed by GitHub
parent 8fd7a40698
commit f149dba611
32 changed files with 2626 additions and 600 deletions
+52
View File
@@ -0,0 +1,52 @@
FROM nvidia/cuda:12.1.0-cudnn8-devel-ubuntu22.04
ENV DEBIAN_FRONTEND noninteractive
RUN rm -r /etc/apt/sources.list.d/
RUN apt-get update -y && apt-get install -y \
libgl1 libglib2.0-0 google-perftools \
sudo wget git git-lfs vim tig pkg-config libcairo2-dev \
aria2 telnet curl net-tools iputils-ping jq \
python3-pip python-is-python3 python3.10-venv tzdata lsof zip tmux
RUN apt-get update && \
apt-get install -y software-properties-common && \
add-apt-repository ppa:ubuntuhandbook1/ffmpeg6 && \
apt-get update && \
apt-get install -y ffmpeg
RUN pip3 install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/
# add all extensions
RUN pip install wandb tqdm GitPython==3.1.32 Pillow==9.5.0 setuptools --upgrade -i https://mirrors.aliyun.com/pypi/simple/
RUN pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 --index-url https://download.pytorch.org/whl/cu118
RUN pip install xformers==0.0.27.post2 --index-url https://download.pytorch.org/whl/cu118
# install vllm (video-caption)
RUN pip install vllm==0.6.3
# install requirements (video-caption)
WORKDIR /root/
COPY easyanimate/video_caption/requirements.txt /root/requirements-video_caption.txt
RUN pip install -r /root/requirements-video_caption.txt
RUN rm /root/requirements-video_caption.txt
RUN pip install -U http://eas-data.oss-cn-shanghai.aliyuncs.com/sdk/allspark-0.15-py2.py3-none-any.whl
RUN pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers
RUN pip install came-pytorch deepspeed pytorch_lightning==1.9.4 func_timeout -i https://mirrors.aliyun.com/pypi/simple/
# install requirements
RUN pip install bitsandbytes mamba-ssm causal-conv1d>=1.4.0 -i https://mirrors.aliyun.com/pypi/simple/
RUN pip install ipykernel -i https://mirrors.aliyun.com/pypi/simple/
COPY ./requirements.txt /root/requirements.txt
RUN pip install -r /root/requirements.txt -i https://mirrors.aliyun.com/pypi/simple/
RUN rm -rf /root/requirements.txt
# install package patches (video-caption)
COPY easyanimate/video_caption/package_patches/easyocr_detection_patched.py /usr/local/lib/python3.10/dist-packages/easyocr/detection.py
COPY easyanimate/video_caption/package_patches/vila_siglip_encoder_patched.py /usr/local/lib/python3.10/dist-packages/llava/model/multimodal_encoder/siglip_encoder.py
ENV PYTHONUNBUFFERED 1
ENV NVIDIA_DISABLE_REQUIRE 1
WORKDIR /root/
Regular → Executable
+20 -1
View File
@@ -22,7 +22,7 @@ class FunTextBox:
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "",}),
}
},
}
RETURN_TYPES = ("STRING_PROMPT",)
@@ -33,9 +33,26 @@ class FunTextBox:
def process(self, prompt):
return (prompt, )
class FunRiflex:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"riflex_k": ("INT", {"default": 6, "min": 0, "max": 10086}),
},
}
RETURN_TYPES = ("RIFLEXT_ARGS",)
RETURN_NAMES = ("riflex_k",)
FUNCTION = "process"
CATEGORY = "CogVideoXFUNWrapper"
def process(self, riflex_k):
return (riflex_k, )
NODE_CLASS_MAPPINGS = {
"FunTextBox": FunTextBox,
"FunRiflex": FunRiflex,
"LoadCogVideoXFunModel": LoadCogVideoXFunModel,
"LoadCogVideoXFunLora": LoadCogVideoXFunLora,
@@ -62,6 +79,8 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"FunTextBox": "FunTextBox",
"FunRiflex": "FunRiflex",
"LoadCogVideoXFunModel": "Load CogVideoX-Fun Model",
"LoadCogVideoXFunLora": "Load CogVideoX-Fun Lora",
"CogVideoXFunInpaintSampler": "CogVideoX-Fun Sampler for Image to Video",
Regular → Executable
+15 -5
View File
@@ -263,7 +263,7 @@ class WanT2VSampler:
"STRING_PROMPT",
),
"video_length": (
"INT", {"default": 81, "min": 5, "max": 81, "step": 4}
"INT", {"default": 81, "min": 5, "max": 161, "step": 4}
),
"width": (
"INT", {"default": 832, "min": 64, "max": 2048, "step": 16}
@@ -310,6 +310,9 @@ class WanT2VSampler:
[False, True], {"default": True,}
),
},
"optional":{
"riflex_k": ("RIFLEXT_ARGS",),
},
}
RETURN_TYPES = ("IMAGE",)
@@ -317,7 +320,7 @@ class WanT2VSampler:
FUNCTION = "process"
CATEGORY = "CogVideoXFUNWrapper"
def process(self, funmodels, prompt, negative_prompt, video_length, width, height, is_image, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload):
def process(self, funmodels, prompt, negative_prompt, video_length, width, height, is_image, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, riflex_k=0):
global transformer_cpu_cache
global lora_path_before
device = mm.get_torch_device()
@@ -348,6 +351,9 @@ class WanT2VSampler:
with torch.no_grad():
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
if riflex_k > 0:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
# Apply lora
if funmodels.get("lora_cache", False):
if len(funmodels.get("loras", [])) != 0:
@@ -412,7 +418,7 @@ class WanI2VSampler:
"STRING_PROMPT",
),
"video_length": (
"INT", {"default": 81, "min": 5, "max": 81, "step": 4}
"INT", {"default": 81, "min": 5, "max": 161, "step": 4}
),
"base_resolution": (
[
@@ -455,7 +461,8 @@ class WanI2VSampler:
),
},
"optional":{
"start_img": ("IMAGE",)
"start_img": ("IMAGE",),
"riflex_k": ("RIFLEXT_ARGS",),
},
}
@@ -464,7 +471,7 @@ class WanI2VSampler:
FUNCTION = "process"
CATEGORY = "CogVideoXFUNWrapper"
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, start_img=None, end_img=None):
def process(self, funmodels, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, scheduler, teacache_threshold, enable_teacache, num_skip_start_steps, teacache_offload, start_img=None, end_img=None, riflex_k=0):
global transformer_cpu_cache
global lora_path_before
device = mm.get_torch_device()
@@ -502,6 +509,9 @@ class WanI2VSampler:
video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_img, end_img, video_length=video_length, sample_size=(height, width))
if riflex_k > 0:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
# Apply lora
if funmodels.get("lora_cache", False):
if len(funmodels.get("loras", [])) != 0:
Regular → Executable
+24 -8
View File
@@ -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:
+5
View File
@@ -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
+8
View File
@@ -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(
+10 -1
View File
@@ -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,
+7 -2
View File
@@ -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
+9
View File
@@ -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(
+9
View File
@@ -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)
+9
View File
@@ -44,6 +44,7 @@ GPU_memory_mode = "sequential_cpu_offload"
ulysses_degree = 1
ring_degree = 1
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
@@ -54,6 +55,11 @@ num_skip_start_steps = 5
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
teacache_offload = False
# Riflex config
enable_riflex = False
# Index of intrinsic frequency
riflex_k = 6
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
@@ -202,6 +208,9 @@ with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=ref_image)
sample = pipeline(
Regular → Executable
+3 -3
View File
@@ -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 \
+3 -3
View File
@@ -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 \
+23 -14
View File
@@ -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 "."
```
+12 -10
View File
@@ -1,7 +1,5 @@
## Training Code
The default training commands for the different versions are as follows:
We can choose whether to use deep speed in Wan-Fun, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in Wan-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file.
@@ -41,7 +39,7 @@ Some parameters in the sh file can be confusing, and they are explained in this
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
Wan-Fun without deepspeed:
Wan-Fun-Control without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
@@ -83,7 +81,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \
--enable_bucket \
--uniform_sampling \
--low_vram \
--train_mode="control_object" \
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--trainable_modules "."
```
@@ -131,7 +129,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
--uniform_sampling \
--low_vram \
--use_deepspeed \
--train_mode="control_object" \
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--trainable_modules "."
```
@@ -153,14 +151,14 @@ export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
--config_path="config/wan2.1_fun/wan_civitai.yaml" \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=3 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
@@ -168,7 +166,9 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
@@ -181,7 +181,9 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--use_deepspeed \
--train_mode="inpaint" \
--low_vram
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--trainable_modules "."
```
@@ -0,0 +1,180 @@
## Training Code
We can choose whether to use deep speed in Wan-Fun, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in Wan-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file.
```json
[
{
"file_path": "train/00000001.mp4",
"control_file_path": "control/00000001.mp4",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
{
"file_path": "train/00000002.jpg",
"control_file_path": "control/00000002.jpg",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "image"
},
.....
]
```
Some parameters in the sh file can be confusing, and they are explained in this document:
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution.
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum. For training videos, the height and width will be set to `image_sample_size` as the maximum and `min(video_sample_size, 512)` as the minimum.
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`, and the resolution of video inputs for training is `512x512x49` to `1024x1024x49`.
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49`.
- `training_with_video_token_length` specifies training the model according to token length. For training images and videos, the height and width will be set to `image_sample_size` as the maximum and `video_sample_size` as the minimum.
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=1024`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x49`.
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=1024`, and `image_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `1024x1024x9`.
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
Wan-Fun-Control without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--low_vram
```
Wan-Fun with deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--use_deepspeed \
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--low_vram
```
Wan T2V with deepspeed zero-3:
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model:
```sh
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
```
Training shell command is as follows:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage2.1_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--save_state \
--use_deepspeed \
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--low_vram
```
+3 -3
View File
@@ -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 \
+26 -25
View File
@@ -33,13 +33,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000003.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/597d786b-66e1-4610-8ba0-01bd334dccb3" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000003.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/f0b14ac6-4b11-44df-a060-bc86bbd91d1e" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000003.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/411b329f-6d63-4557-820d-80d56ded2005" width="100%" controls autoplay loop></video>
</td>
</tr>
<tr>
@@ -51,13 +51,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000004.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/cb4ce6d1-1b70-480a-b3c1-e7d1c05cc93d" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000004.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/017203b3-1091-4ba9-95db-71ab62804b7a" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000004.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/d4c7dc8e-57dc-4d08-aef6-d11a7bdb4972" width="100%" controls autoplay loop></video>
</td>
</tr>
<tr>
@@ -69,13 +69,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000007.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/65bdd0dd-717d-4300-b566-a615ab1f81c2" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000007.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/07dca036-5b72-4d1e-9ed3-dc725a18f654" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000007.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/e97b1158-e772-40f3-a28f-538bca5584e3" width="100%" controls autoplay loop></video>
</td>
</tr>
<tr>
@@ -87,13 +87,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_baseline_00000008.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/6af00186-2ec4-43d7-9360-3b316e5240a8" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_hpsv2.1_00000008.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/3f9a5fb2-bfea-469b-8d07-c4289990ee66" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/1.3B_mps_00000008.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/db851428-0df7-4ae3-82eb-c85fe902fb96" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
@@ -122,13 +122,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000001.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/7bfded69-54e5-4654-91e9-37aa61d5b5f3" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000001.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/07bb0403-bcd4-439d-a6a8-d97c266305a8" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000001.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/15e13963-3020-4ffe-9b7b-35603435806f" width="100%" controls autoplay loop></video>
</td>
</tr>
<tr>
@@ -140,13 +140,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000002.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/050f3bff-05c9-4931-9112-e9499b136435" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000002.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/9d1e09bd-0acb-4646-a2ef-2c2d582ba150" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000002.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/f2f64228-dd4b-4c8a-9844-acaf44b6c83c" width="100%" controls autoplay loop></video>
</td>
</tr>
<tr>
@@ -158,13 +158,13 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000005.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/f6b467db-5e0a-4be6-87a0-a1acce692b1b" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000005.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/487b0405-8e99-44bb-ae6d-009442849a94" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000005.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/d581a5a2-1393-4d19-a13e-b669850a0754" width="100%" controls autoplay loop></video>
</td>
</tr>
<tr>
@@ -176,17 +176,18 @@ For more details, please refer to our [GitHub repo](https://github.com/aigc-apps
</details>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_baseline_00000008.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/b3228456-1ee9-4f4c-bfd2-03e13df980d7" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_hpsv2.1_00000008.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/4d85e13d-4741-437b-991a-195802cb9485" width="100%" controls autoplay loop></video>
</td>
<td>
<video src="https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset/v1.0/reward_lora/14B_mps_00000008.mp4" width="100%" controls autoplay loop></video>
<video src="https://github.com/user-attachments/assets/6dd31da7-4f1a-44b4-bdbb-14fe62e2796b" width="100%" controls autoplay loop></video>
</td>
</tr>
</table>
> [!NOTE]
> The above test prompts are from <a href="https://github.com/KaiyueSun98/T2V-CompBench">T2V-CompBench</a> and expanded into detailed prompts by Llama-3.3.
> Videos are generated with HPSv2.1 Reward LoRA weight 0.7 and MPS Reward LoRA weight 0.7.
@@ -205,4 +206,4 @@ Set `lora_path` and `lora_weight` in [examples/wan2.1_fun/predict_t2v.py](https
<ol>
<li id="ref1">Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.</li>
<li id="ref2">Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).</li>
</ol>
</ol>
+89 -86
View File
@@ -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))
+81 -116
View File
@@ -21,6 +21,7 @@ import logging
import math
import os
import pickle
import random
import shutil
import sys
@@ -47,8 +48,6 @@ from einops import rearrange
from omegaconf import OmegaConf
from packaging import version
from PIL import Image
from torch.distributed.fsdp.fully_sharded_data_parallel import (
FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig)
from torch.utils.data import RandomSampler
from torch.utils.tensorboard import SummaryWriter
from torchvision import transforms
@@ -72,10 +71,11 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset,
ImageVideoSampler,
get_random_mask)
from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
WanTransformer3DModel)
from videox_fun.pipeline import WanFunControlPipeline, WanFunInpaintPipeline
WanTransformer3DModel)
from videox_fun.pipeline import WanFunControlPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import get_video_to_video_latent, save_videos_grid
from videox_fun.utils.utils import (get_video_to_video_latent,
save_videos_grid)
if is_wandb_available():
import wandb
@@ -88,73 +88,18 @@ def filter_kwargs(cls, kwargs):
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
return filtered_kwargs
def get_random_downsample_ratio(sample_size, image_ratio=[],
all_choices=False, rng=None):
def _create_special_list(length):
if length == 1:
return [1.0]
if length >= 2:
first_element = 0.75
remaining_sum = 1.0 - first_element
other_elements_value = remaining_sum / (length - 1)
special_list = [first_element] + [other_elements_value] * (length - 1)
return special_list
if sample_size >= 1536:
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
elif sample_size >= 1024:
number_list = [1, 1.25, 1.5, 2] + image_ratio
elif sample_size >= 768:
number_list = [1, 1.25, 1.5] + image_ratio
elif sample_size >= 512:
number_list = [1] + image_ratio
else:
number_list = [1]
def linear_decay(initial_value, final_value, total_steps, current_step):
if current_step >= total_steps:
return final_value
current_step = max(0, current_step)
step_size = (final_value - initial_value) / total_steps
current_value = initial_value + step_size * current_step
return current_value
if all_choices:
return number_list
number_list_prob = np.array(_create_special_list(len(number_list)))
if rng is None:
return np.random.choice(number_list, p = number_list_prob)
else:
return rng.choice(number_list, p = number_list_prob)
def resize_mask(mask, latent, process_first_frame_only=True):
latent_size = latent.size()
batch_size, channels, num_frames, height, width = mask.shape
if process_first_frame_only:
target_size = list(latent_size[2:])
target_size[0] = 1
first_frame_resized = F.interpolate(
mask[:, :, 0:1, :, :],
size=target_size,
mode='trilinear',
align_corners=False
)
target_size = list(latent_size[2:])
target_size[0] = target_size[0] - 1
if target_size[0] != 0:
remaining_frames_resized = F.interpolate(
mask[:, :, 1:, :, :],
size=target_size,
mode='trilinear',
align_corners=False
)
resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2)
else:
resized_mask = first_frame_resized
else:
target_size = list(latent_size[2:])
resized_mask = F.interpolate(
mask,
size=target_size,
mode='trilinear',
align_corners=False
)
return resized_mask
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
return torch.clip(t.to(torch.int32), low, high - 1)
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.18.0.dev0")
@@ -222,19 +167,6 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
print(f"Eval error with info {e}")
return None
def linear_decay(initial_value, final_value, total_steps, current_step):
if current_step >= total_steps:
return final_value
current_step = max(0, current_step)
step_size = (final_value - initial_value) / total_steps
current_value = initial_value + step_size * current_step
return current_value
def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None):
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
return torch.clip(t.to(torch.int32), low, high - 1)
def parse_args():
parser = argparse.ArgumentParser(description="Simple example of a training script.")
parser.add_argument(
@@ -1040,6 +972,14 @@ def main():
image_sample_size=args.image_sample_size,
enable_bucket=args.enable_bucket, enable_inpaint=False,
)
def worker_init_fn(_seed):
_seed = _seed * 256
def _worker_init_fn(worker_id):
print(f"worker_init_fn with {_seed + worker_id}")
np.random.seed(_seed + worker_id)
random.seed(_seed + worker_id)
return _worker_init_fn
if args.enable_bucket:
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
@@ -1050,22 +990,54 @@ def main():
aspect_ratios=aspect_ratio_sample_size,
)
def get_length_to_frame_num(token_length):
if args.image_sample_size > args.video_sample_size:
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
if sample_sizes[-1] != args.image_sample_size:
sample_sizes.append(args.image_sample_size)
else:
sample_sizes = [args.image_sample_size]
length_to_frame_num = {
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
}
return length_to_frame_num
def collate_fn(examples):
def get_length_to_frame_num(token_length):
if args.image_sample_size > args.video_sample_size:
sample_sizes = list(range(args.video_sample_size, args.image_sample_size + 1, 128))
if sample_sizes[-1] != args.image_sample_size:
sample_sizes.append(args.image_sample_size)
else:
sample_sizes = [args.image_sample_size]
length_to_frame_num = {
sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes
}
return length_to_frame_num
def get_random_downsample_ratio(sample_size, image_ratio=[],
all_choices=False, rng=None):
def _create_special_list(length):
if length == 1:
return [1.0]
if length >= 2:
first_element = 0.90
remaining_sum = 1.0 - first_element
other_elements_value = remaining_sum / (length - 1)
special_list = [first_element] + [other_elements_value] * (length - 1)
return special_list
if sample_size >= 1536:
number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio
elif sample_size >= 1024:
number_list = [1, 1.25, 1.5, 2] + image_ratio
elif sample_size >= 768:
number_list = [1, 1.25, 1.5] + image_ratio
elif sample_size >= 512:
number_list = [1] + image_ratio
else:
number_list = [1]
if all_choices:
return number_list
number_list_prob = np.array(_create_special_list(len(number_list)))
if rng is None:
return np.random.choice(number_list, p = number_list_prob)
else:
return rng.choice(number_list, p = number_list_prob)
# Get token length
target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size
length_to_frame_num = get_length_to_frame_num(target_token_length)
@@ -1087,7 +1059,7 @@ def main():
data_type = examples[0]["data_type"]
f, h, w, c = np.shape(pixel_value)
if data_type == 'image':
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size], rng=rng)
random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size])
aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
@@ -1101,15 +1073,11 @@ def main():
choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25]
if len(choice_list) == 0:
choice_list = list(length_to_frame_num.keys())
if rng is None:
local_video_sample_size = np.random.choice(choice_list)
else:
local_video_sample_size = rng.choice(choice_list)
local_video_sample_size = np.random.choice(choice_list)
batch_video_length = length_to_frame_num[local_video_sample_size]
random_downsample_ratio = args.video_sample_size / local_video_sample_size
else:
random_downsample_ratio = get_random_downsample_ratio(
args.video_sample_size, rng=rng)
random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size)
batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval
else:
random_downsample_ratio = 1
@@ -1121,14 +1089,9 @@ def main():
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
else:
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
for example in examples:
@@ -1239,6 +1202,7 @@ def main():
collate_fn=collate_fn,
persistent_workers=True if args.dataloader_num_workers != 0 else False,
num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
)
else:
# DataLoaders creation:
@@ -1249,6 +1213,7 @@ def main():
batch_sampler=batch_sampler,
persistent_workers=True if args.dataloader_num_workers != 0 else False,
num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index)
)
# Scheduler and math around the number of training steps.
@@ -1400,7 +1365,7 @@ def main():
control_pixel_values = batch["control_pixel_values"].to(weight_dtype)
# Increase the batch size when the length of the latent sequence of the current sample is small
if args.training_with_video_token_length and not zero_stage == 3:
if args.training_with_video_token_length and zero_stage != 3:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1))
@@ -1422,7 +1387,7 @@ def main():
ref_pixel_values = batch["ref_pixel_values"].to(weight_dtype)
clip_pixel_values = batch["clip_pixel_values"]
# Increase the batch size when the length of the latent sequence of the current sample is small
if args.training_with_video_token_length and not zero_stage == 3:
if args.training_with_video_token_length and zero_stage != 3:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1))
File diff suppressed because it is too large Load Diff
+39
View File
@@ -0,0 +1,39 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-14B-Control"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora.py \
--config_path="config/wan2.1/wan_civitai.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--train_mode="control_ref" \
--control_ref_image="first_frame" \
--low_vram
+117 -109
View File
@@ -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))
+1 -1
View File
@@ -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):
+139 -27
View File
@@ -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):
-4
View File
@@ -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()
-4
View File
@@ -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()
+1 -5
View File
@@ -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:
+4 -1
View File
@@ -57,7 +57,8 @@ class Fun_Controller:
self, GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
config_path=None, ulysses_degree=1, ring_degree=1,
enable_teacache=None, teacache_threshold=None,
num_skip_start_steps=None, teacache_offload=None, weight_dtype=None,
num_skip_start_steps=None, teacache_offload=None,
enable_riflex=None, riflex_k=None, weight_dtype=None,
):
# config dirs
self.basedir = os.getcwd()
@@ -81,6 +82,8 @@ class Fun_Controller:
self.teacache_threshold = teacache_threshold
self.num_skip_start_steps = num_skip_start_steps
self.teacache_offload = teacache_offload
self.enable_riflex = enable_riflex
self.riflex_k = riflex_k
self.weight_dtype = weight_dtype
self.device = set_multi_gpus_devices(self.ulysses_degree, self.ring_degree)
Regular → Executable
+25 -83
View File
@@ -186,91 +186,31 @@ class Wan_Fun_Controller(Fun_Controller):
else: seed_textbox = np.random.randint(0, 1e10)
generator = torch.Generator(device=self.device).manual_seed(int(seed_textbox))
if self.enable_riflex:
self.pipeline.transformer.enable_riflex(k = self.riflex_k, L_test = length_slider if not is_image else 1)
try:
if self.model_type == "Inpaint":
if self.transformer.config.in_channels != self.vae.config.latent_channels:
if generation_method == "Long Video Generation":
if validation_video is not None:
raise gr.Error(f"Video to Video is not Support Long Video Generation now.")
init_frames = 0
last_frames = init_frames + partial_video_length
while init_frames < length_slider:
if last_frames >= length_slider:
_partial_video_length = length_slider - init_frames
_partial_video_length = int((_partial_video_length - 1) // self.vae.config.temporal_compression_ratio * self.vae.config.temporal_compression_ratio) + 1
if _partial_video_length <= 0:
break
else:
_partial_video_length = partial_video_length
if last_frames >= length_slider:
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=_partial_video_length, sample_size=(height_slider, width_slider))
else:
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, None, video_length=_partial_video_length, sample_size=(height_slider, width_slider))
with torch.no_grad():
sample = self.pipeline(
prompt_textbox,
negative_prompt = negative_prompt_textbox,
num_inference_steps = sample_step_slider,
guidance_scale = cfg_scale_slider,
width = width_slider,
height = height_slider,
num_frames = _partial_video_length,
generator = generator,
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image
).videos
if init_frames != 0:
mix_ratio = torch.from_numpy(
np.array([float(_index) / float(overlap_video_length) for _index in range(overlap_video_length)], np.float32)
).unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
new_sample[:, :, -overlap_video_length:] = new_sample[:, :, -overlap_video_length:] * (1 - mix_ratio) + \
sample[:, :, :overlap_video_length] * mix_ratio
new_sample = torch.cat([new_sample, sample[:, :, overlap_video_length:]], dim = 2)
sample = new_sample
else:
new_sample = sample
if last_frames >= length_slider:
break
start_image = [
Image.fromarray(
(sample[0, :, _index].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
) for _index in range(-overlap_video_length, 0)
]
init_frames = init_frames + _partial_video_length - overlap_video_length
last_frames = init_frames + _partial_video_length
if validation_video is not None:
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=16)
else:
if validation_video is not None:
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=16)
strength = denoise_strength
else:
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
strength = 1
input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider))
sample = self.pipeline(
prompt_textbox,
negative_prompt = negative_prompt_textbox,
num_inference_steps = sample_step_slider,
guidance_scale = cfg_scale_slider,
width = width_slider,
height = height_slider,
num_frames = length_slider if not is_image else 1,
generator = generator,
sample = self.pipeline(
prompt_textbox,
negative_prompt = negative_prompt_textbox,
num_inference_steps = sample_step_slider,
guidance_scale = cfg_scale_slider,
width = width_slider,
height = height_slider,
num_frames = length_slider if not is_image else 1,
generator = generator,
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image
).videos
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image
).videos
else:
sample = self.pipeline(
prompt_textbox,
@@ -335,12 +275,13 @@ class Wan_Fun_Controller(Fun_Controller):
Wan_Fun_Controller_Host = Wan_Fun_Controller
Wan_Fun_Controller_Client = Fun_Controller_Client
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype):
def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype):
controller = Wan_Fun_Controller(
GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint",
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
enable_teacache=enable_teacache, teacache_threshold=teacache_threshold,
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, weight_dtype=weight_dtype,
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload,
enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype,
)
with gr.Blocks(css=css) as demo:
@@ -460,12 +401,13 @@ def ui(GPU_memory_mode, scheduler_dict, config_path, ulysses_degree, ring_degree
)
return demo, controller
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, weight_dtype):
def ui_host(GPU_memory_mode, scheduler_dict, model_name, model_type, config_path, ulysses_degree, ring_degree, enable_teacache, teacache_threshold, num_skip_start_steps, teacache_offload, enable_riflex, riflex_k, weight_dtype):
controller = Wan_Fun_Controller_Host(
GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type,
config_path=config_path, ulysses_degree=ulysses_degree, ring_degree=ring_degree,
enable_teacache=enable_teacache, teacache_threshold=teacache_threshold,
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload, weight_dtype=weight_dtype,
num_skip_start_steps=num_skip_start_steps, teacache_offload=teacache_offload,
enable_riflex=enable_riflex, riflex_k=riflex_k, weight_dtype=weight_dtype,
)
with gr.Blocks(css=css) as demo:
Regular → Executable
+25 -83
View File
@@ -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:
+5 -6
View File
@@ -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