Vace training && Comfyui update && Fix flux readme && Update MV2V demo (#328)
This commit is contained in:
Binary file not shown.
Binary file not shown.
@@ -198,13 +198,18 @@ class LoadCogVideoXFunLora:
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def load_lora(self, cogvideoxfun_model, lora_name, strength_model, lora_cache):
|
||||
new_funmodels = dict(cogvideoxfun_model)
|
||||
|
||||
if lora_name is not None:
|
||||
cogvideoxfun_model['lora_cache'] = lora_cache
|
||||
cogvideoxfun_model['loras'] = cogvideoxfun_model.get("loras", []) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
cogvideoxfun_model['strength_model'] = cogvideoxfun_model.get("strength_model", []) + [strength_model]
|
||||
return (cogvideoxfun_model,)
|
||||
else:
|
||||
return (cogvideoxfun_model,)
|
||||
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
if lora_path is None:
|
||||
raise FileNotFoundError(f"LoRA 文件未找到: {lora_name}")
|
||||
|
||||
new_funmodels['lora_cache'] = lora_cache
|
||||
new_funmodels['loras'] = cogvideoxfun_model.get("loras", []) + [lora_path]
|
||||
new_funmodels['strength_model'] = cogvideoxfun_model.get("strength_model", []) + [strength_model]
|
||||
|
||||
return (new_funmodels,)
|
||||
|
||||
class CogVideoXFunT2VSampler:
|
||||
@classmethod
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
@@ -83,20 +84,28 @@ class FunCompile:
|
||||
def compile(self, cache_size_limit, funmodels):
|
||||
torch._dynamo.config.cache_size_limit = cache_size_limit
|
||||
if hasattr(funmodels["pipeline"].transformer, "blocks"):
|
||||
for i in range(len(funmodels["pipeline"].transformer.blocks)):
|
||||
funmodels["pipeline"].transformer.blocks[i] = torch.compile(funmodels["pipeline"].transformer.blocks[i])
|
||||
for i, block in enumerate(funmodels["pipeline"].transformer.blocks):
|
||||
if hasattr(block, "_orig_mod"):
|
||||
block = block._orig_mod
|
||||
funmodels["pipeline"].transformer.blocks[i] = torch.compile(block)
|
||||
|
||||
if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None:
|
||||
for i in range(len(funmodels["pipeline"].transformer_2.blocks)):
|
||||
funmodels["pipeline"].transformer_2.blocks[i] = torch.compile(funmodels["pipeline"].transformer_2.blocks[i])
|
||||
for i, block in enumerate(funmodels["pipeline"].transformer_2.blocks):
|
||||
if hasattr(block, "_orig_mod"):
|
||||
block = block._orig_mod
|
||||
funmodels["pipeline"].transformer_2.blocks[i] = torch.compile(block)
|
||||
|
||||
elif hasattr(funmodels["pipeline"].transformer, "transformer_blocks"):
|
||||
for i in range(len(funmodels["pipeline"].transformer.transformer_blocks)):
|
||||
funmodels["pipeline"].transformer.transformer_blocks[i] = torch.compile(funmodels["pipeline"].transformer.transformer_blocks[i])
|
||||
|
||||
for i, block in enumerate(funmodels["pipeline"].transformer.transformer_blocks):
|
||||
if hasattr(block, "_orig_mod"):
|
||||
block = block._orig_mod
|
||||
funmodels["pipeline"].transformer.transformer_blocks[i] = torch.compile(block)
|
||||
|
||||
if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None:
|
||||
for i in range(len(funmodels["pipeline"].transformer_2.transformer_blocks)):
|
||||
funmodels["pipeline"].transformer_2.transformer_blocks[i] = torch.compile(funmodels["pipeline"].transformer_2.transformer_blocks[i])
|
||||
for i, block in enumerate(funmodels["pipeline"].transformer_2.transformer_blocks):
|
||||
if hasattr(block, "_orig_mod"):
|
||||
block = block._orig_mod
|
||||
funmodels["pipeline"].transformer_2.transformer_blocks[i] = torch.compile(block)
|
||||
|
||||
else:
|
||||
funmodels["pipeline"].transformer.forward = torch.compile(funmodels["pipeline"].transformer.forward)
|
||||
@@ -106,6 +115,31 @@ class FunCompile:
|
||||
|
||||
print("Add Compile")
|
||||
return (funmodels,)
|
||||
|
||||
class FunAttention:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"attention_type": (
|
||||
["flash", "sage", "torch"],
|
||||
{"default": "flash"},
|
||||
),
|
||||
"funmodels": ("FunModels",)
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("FunModels",)
|
||||
RETURN_NAMES = ("funmodels",)
|
||||
FUNCTION = "funattention"
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def funattention(self, attention_type, funmodels):
|
||||
os.environ['VIDEOX_ATTENTION_TYPE'] = {
|
||||
"flash": "FLASH_ATTENTION",
|
||||
"sage": "SAGE_ATTENTION",
|
||||
"torch": "TORCH_SCALED_DOT"
|
||||
}[attention_type]
|
||||
return (funmodels,)
|
||||
|
||||
class LoadConfig:
|
||||
@classmethod
|
||||
@@ -376,6 +410,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FunTextBox": FunTextBox,
|
||||
"FunRiflex": FunRiflex,
|
||||
"FunCompile": FunCompile,
|
||||
"FunAttention": FunAttention,
|
||||
|
||||
"LoadCogVideoXFunModel": LoadCogVideoXFunModel,
|
||||
"LoadCogVideoXFunLora": LoadCogVideoXFunLora,
|
||||
@@ -436,6 +471,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FunTextBox": "FunTextBox",
|
||||
"FunRiflex": "FunRiflex",
|
||||
"FunCompile": "FunCompile",
|
||||
"FunAttention": "FunAttention",
|
||||
|
||||
"LoadWanClipEncoderModel": "Load Wan ClipEncoder Model",
|
||||
"LoadWanTextEncoderModel": "Load Wan TextEncoder Model",
|
||||
|
||||
+10
-6
@@ -612,13 +612,17 @@ class LoadWanLora:
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def load_lora(self, funmodels, lora_name, strength_model, lora_cache):
|
||||
new_funmodels = dict(funmodels)
|
||||
|
||||
if lora_name is not None:
|
||||
funmodels['lora_cache'] = lora_cache
|
||||
funmodels['loras'] = funmodels.get("loras", []) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
funmodels['strength_model'] = funmodels.get("strength_model", []) + [strength_model]
|
||||
return (funmodels,)
|
||||
else:
|
||||
return (funmodels,)
|
||||
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
|
||||
new_funmodels['lora_cache'] = lora_cache
|
||||
new_funmodels['loras'] = funmodels.get("loras", []) + [lora_path]
|
||||
new_funmodels['strength_model'] = funmodels.get("strength_model", []) + [strength_model]
|
||||
|
||||
return (new_funmodels,)
|
||||
|
||||
|
||||
class WanT2VSampler:
|
||||
@classmethod
|
||||
|
||||
@@ -238,13 +238,16 @@ class LoadWanFunLora:
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def load_lora(self, funmodels, lora_name, strength_model, lora_cache):
|
||||
new_funmodels = dict(funmodels)
|
||||
|
||||
if lora_name is not None:
|
||||
funmodels['lora_cache'] = lora_cache
|
||||
funmodels['loras'] = funmodels.get("loras", []) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
funmodels['strength_model'] = funmodels.get("strength_model", []) + [strength_model]
|
||||
return (funmodels,)
|
||||
else:
|
||||
return (funmodels,)
|
||||
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
|
||||
new_funmodels['lora_cache'] = lora_cache
|
||||
new_funmodels['loras'] = funmodels.get("loras", []) + [lora_path]
|
||||
new_funmodels['strength_model'] = funmodels.get("strength_model", []) + [strength_model]
|
||||
|
||||
return (new_funmodels,)
|
||||
|
||||
class WanFunT2VSampler:
|
||||
@classmethod
|
||||
|
||||
@@ -463,14 +463,16 @@ class LoadWan2_2Lora:
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def load_lora(self, funmodels, lora_name, lora_high_name, strength_model, lora_cache):
|
||||
new_funmodels = dict(funmodels)
|
||||
if lora_name is not None:
|
||||
funmodels['lora_cache'] = lora_cache
|
||||
funmodels['loras'] = funmodels.get("loras", []) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
funmodels['loras_high'] = funmodels.get("loras_high", []) + [folder_paths.get_full_path("loras", lora_high_name)]
|
||||
funmodels['strength_model'] = funmodels.get("strength_model", []) + [strength_model]
|
||||
return (funmodels,)
|
||||
else:
|
||||
return (funmodels,)
|
||||
loras = list(new_funmodels.get("loras", [])) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
loras_high = list(new_funmodels.get("loras_high", [])) + [folder_paths.get_full_path("loras", lora_high_name)]
|
||||
strength_models = list(new_funmodels.get("strength_model", [])) + [strength_model]
|
||||
new_funmodels['loras'] = loras
|
||||
new_funmodels['loras_high'] = loras_high
|
||||
new_funmodels['strength_model'] = strength_models
|
||||
new_funmodels['lora_cache'] = lora_cache
|
||||
return (new_funmodels,)
|
||||
|
||||
class Wan2_2T2VSampler:
|
||||
@classmethod
|
||||
|
||||
@@ -258,14 +258,16 @@ class LoadWan2_2FunLora:
|
||||
CATEGORY = "CogVideoXFUNWrapper"
|
||||
|
||||
def load_lora(self, funmodels, lora_name, lora_high_name, strength_model, lora_cache):
|
||||
new_funmodels = dict(funmodels)
|
||||
if lora_name is not None:
|
||||
funmodels['lora_cache'] = lora_cache
|
||||
funmodels['loras'] = funmodels.get("loras", []) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
funmodels['loras_high'] = funmodels.get("loras_high", []) + [folder_paths.get_full_path("loras", lora_high_name)]
|
||||
funmodels['strength_model'] = funmodels.get("strength_model", []) + [strength_model]
|
||||
return (funmodels,)
|
||||
else:
|
||||
return (funmodels,)
|
||||
loras = list(new_funmodels.get("loras", [])) + [folder_paths.get_full_path("loras", lora_name)]
|
||||
loras_high = list(new_funmodels.get("loras_high", [])) + [folder_paths.get_full_path("loras", lora_high_name)]
|
||||
strength_models = list(new_funmodels.get("strength_model", [])) + [strength_model]
|
||||
new_funmodels['loras'] = loras
|
||||
new_funmodels['loras_high'] = loras_high
|
||||
new_funmodels['strength_model'] = strength_models
|
||||
new_funmodels['lora_cache'] = lora_cache
|
||||
return (new_funmodels,)
|
||||
|
||||
class Wan2_2FunT2VSampler:
|
||||
@classmethod
|
||||
|
||||
@@ -202,7 +202,7 @@ else:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
if partial_video_length is not None:
|
||||
partial_video_length = int((partial_video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
@@ -292,7 +292,7 @@ else:
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -194,7 +194,7 @@ else:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -232,7 +232,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -201,7 +201,7 @@ else:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -227,7 +227,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -188,7 +188,7 @@ else:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -212,7 +212,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -184,7 +184,7 @@ else:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
sample = pipeline(
|
||||
@@ -198,7 +198,7 @@ with torch.no_grad():
|
||||
).images
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -244,7 +244,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -272,7 +272,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -176,7 +176,7 @@ else:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
sample = pipeline(
|
||||
@@ -190,7 +190,7 @@ with torch.no_grad():
|
||||
).images
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -245,7 +245,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -273,7 +273,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -232,7 +232,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -254,7 +254,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -246,7 +246,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -274,7 +274,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -253,7 +253,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -293,7 +293,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -256,7 +256,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -305,7 +305,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -256,7 +256,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -305,7 +305,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -256,7 +256,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -305,7 +305,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -247,7 +247,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -283,7 +283,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -247,7 +247,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -283,7 +283,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -247,7 +247,7 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
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
|
||||
@@ -283,7 +283,7 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -278,8 +278,8 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -308,8 +308,8 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -302,9 +302,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
with torch.no_grad():
|
||||
video_length = video_length // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio if video_length != 1 else 1
|
||||
@@ -339,8 +339,8 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -273,8 +273,8 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -298,8 +298,8 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -290,9 +290,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -325,9 +325,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -292,9 +292,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -324,9 +324,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -294,9 +294,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -326,9 +326,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -277,8 +277,8 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -307,8 +307,8 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -288,9 +288,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -320,9 +320,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -305,9 +305,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -351,9 +351,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -305,9 +305,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -351,9 +351,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -305,9 +305,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -351,9 +351,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -305,9 +305,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -351,9 +351,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -305,9 +305,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -351,9 +351,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -305,9 +305,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -351,9 +351,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -108,12 +108,15 @@ fps = 16
|
||||
|
||||
# Use torch.float16 if GPU does not support torch.bfloat16
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = None
|
||||
start_image = "asset/1.png"
|
||||
end_image = None
|
||||
subject_ref_images = None
|
||||
vace_context_scale = 1.00
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = None
|
||||
start_image = "asset/1.png"
|
||||
end_image = None
|
||||
# Use inpaint video instead of start image and end image.
|
||||
inpaint_video = None
|
||||
inpaint_video_mask = None
|
||||
subject_ref_images = None
|
||||
vace_context_scale = 1.00
|
||||
# Sometimes, when generating a video from a reference image, white borders appear.
|
||||
# Because the padding is mistakenly treated as part of the image.
|
||||
# If the aspect ratio of the reference image is close to the final output, you can omit the white padding.
|
||||
@@ -306,9 +309,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -323,7 +326,14 @@ with torch.no_grad():
|
||||
subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=padding_in_subject_ref_images) for _subject_ref_image in subject_ref_images]
|
||||
subject_ref_images = torch.cat(subject_ref_images, dim=2)
|
||||
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
if inpaint_video is not None:
|
||||
if inpaint_video_mask is None:
|
||||
raise ValueError("inpaint_video_mask is required when inpaint_video is provided")
|
||||
inpaint_video, _, _, _ = get_video_to_video_latent(inpaint_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask, _, _, _ = get_video_to_video_latent(inpaint_video_mask, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask = inpaint_video_mask[:, :1]
|
||||
else:
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
|
||||
@@ -347,9 +357,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -108,12 +108,15 @@ fps = 16
|
||||
|
||||
# Use torch.float16 if GPU does not support torch.bfloat16
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = None
|
||||
start_image = None
|
||||
end_image = None
|
||||
subject_ref_images = ["asset/8.png", "asset/ref_1.png"]
|
||||
vace_context_scale = 1.00
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = None
|
||||
start_image = None
|
||||
end_image = None
|
||||
# Use inpaint video instead of start image and end image.
|
||||
inpaint_video = None
|
||||
inpaint_video_mask = None
|
||||
subject_ref_images = ["asset/8.png", "asset/ref_1.png"]
|
||||
vace_context_scale = 1.00
|
||||
# Sometimes, when generating a video from a reference image, white borders appear.
|
||||
# Because the padding is mistakenly treated as part of the image.
|
||||
# If the aspect ratio of the reference image is close to the final output, you can omit the white padding.
|
||||
@@ -306,9 +309,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -323,7 +326,14 @@ with torch.no_grad():
|
||||
subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=padding_in_subject_ref_images) for _subject_ref_image in subject_ref_images]
|
||||
subject_ref_images = torch.cat(subject_ref_images, dim=2)
|
||||
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
if inpaint_video is not None:
|
||||
if inpaint_video_mask is None:
|
||||
raise ValueError("inpaint_video_mask is required when inpaint_video is provided")
|
||||
inpaint_video, _, _, _ = get_video_to_video_latent(inpaint_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask, _, _, _ = get_video_to_video_latent(inpaint_video_mask, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask = inpaint_video_mask[:, :1]
|
||||
else:
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
|
||||
@@ -347,9 +357,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -108,12 +108,15 @@ fps = 16
|
||||
|
||||
# Use torch.float16 if GPU does not support torch.bfloat16
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = "asset/pose.mp4"
|
||||
start_image = None
|
||||
end_image = None
|
||||
subject_ref_images = None
|
||||
vace_context_scale = 1.00
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = "asset/pose.mp4"
|
||||
start_image = None
|
||||
end_image = None
|
||||
# Use inpaint video instead of start image and end image.
|
||||
inpaint_video = None
|
||||
inpaint_video_mask = None
|
||||
subject_ref_images = None
|
||||
vace_context_scale = 1.00
|
||||
# Sometimes, when generating a video from a reference image, white borders appear.
|
||||
# Because the padding is mistakenly treated as part of the image.
|
||||
# If the aspect ratio of the reference image is close to the final output, you can omit the white padding.
|
||||
@@ -306,9 +309,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -323,7 +326,14 @@ with torch.no_grad():
|
||||
subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=padding_in_subject_ref_images) for _subject_ref_image in subject_ref_images]
|
||||
subject_ref_images = torch.cat(subject_ref_images, dim=2)
|
||||
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
if inpaint_video is not None:
|
||||
if inpaint_video_mask is None:
|
||||
raise ValueError("inpaint_video_mask is required when inpaint_video is provided")
|
||||
inpaint_video, _, _, _ = get_video_to_video_latent(inpaint_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask, _, _, _ = get_video_to_video_latent(inpaint_video_mask, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask = inpaint_video_mask[:, :1]
|
||||
else:
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
|
||||
@@ -347,9 +357,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -108,12 +108,15 @@ fps = 16
|
||||
|
||||
# Use torch.float16 if GPU does not support torch.bfloat16
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = "asset/pose.mp4"
|
||||
start_image = None
|
||||
end_image = None
|
||||
subject_ref_images = ["asset/8.png"]
|
||||
vace_context_scale = 1.00
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = "asset/pose.mp4"
|
||||
start_image = None
|
||||
end_image = None
|
||||
# Use inpaint video instead of start image and end image.
|
||||
inpaint_video = None
|
||||
inpaint_video_mask = None
|
||||
subject_ref_images = ["asset/8.png"]
|
||||
vace_context_scale = 1.00
|
||||
# Sometimes, when generating a video from a reference image, white borders appear.
|
||||
# Because the padding is mistakenly treated as part of the image.
|
||||
# If the aspect ratio of the reference image is close to the final output, you can omit the white padding.
|
||||
@@ -306,9 +309,9 @@ if cfg_skip_ratio is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
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
|
||||
@@ -323,7 +326,14 @@ with torch.no_grad():
|
||||
subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=padding_in_subject_ref_images) for _subject_ref_image in subject_ref_images]
|
||||
subject_ref_images = torch.cat(subject_ref_images, dim=2)
|
||||
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
if inpaint_video is not None:
|
||||
if inpaint_video_mask is None:
|
||||
raise ValueError("inpaint_video_mask is required when inpaint_video is provided")
|
||||
inpaint_video, _, _, _ = get_video_to_video_latent(inpaint_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask, _, _, _ = get_video_to_video_latent(inpaint_video_mask, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask = inpaint_video_mask[:, :1]
|
||||
else:
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
|
||||
@@ -347,9 +357,9 @@ with torch.no_grad():
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
|
||||
@@ -0,0 +1,387 @@
|
||||
import os
|
||||
import sys
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from omegaconf import OmegaConf
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
|
||||
for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel,
|
||||
WanT5EncoderModel, VaceWanTransformer3DModel)
|
||||
from videox_fun.data.dataset_image_video import process_pose_file
|
||||
from videox_fun.models.cache_utils import get_teacache_coefficients
|
||||
from videox_fun.pipeline import Wan2_2VaceFunPipeline, WanPipeline
|
||||
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
|
||||
convert_weight_dtype_wrapper,
|
||||
replace_parameters_by_name)
|
||||
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
|
||||
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent,
|
||||
get_video_to_video_latent,
|
||||
save_videos_grid)
|
||||
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
|
||||
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
|
||||
# GPU memory mode, which can be chosen in [model_full_load, model_cpu_offload_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
|
||||
# model_full_load means that the entire model will be moved to the GPU.
|
||||
#
|
||||
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
|
||||
# and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
#
|
||||
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
|
||||
#
|
||||
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
|
||||
# and the transformer model has been quantized to float8, which can save more GPU memory.
|
||||
#
|
||||
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
|
||||
# resulting in slower speeds but saving a large amount of GPU memory.
|
||||
GPU_memory_mode = "sequential_cpu_offload"
|
||||
# Multi GPUs config
|
||||
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
|
||||
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
|
||||
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
# Use FSDP to save more GPU memory in multi gpus.
|
||||
fsdp_dit = False
|
||||
fsdp_text_encoder = True
|
||||
# Compile will give a speedup in fixed resolution and need a little GPU memory.
|
||||
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
|
||||
compile_dit = False
|
||||
|
||||
# Support TeaCache.
|
||||
enable_teacache = True
|
||||
# Recommended to be set between 0.05 and 0.30. 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.
|
||||
# # --------------------------------------------------------------------------------------------------- #
|
||||
# | Model Name | threshold | Model Name | threshold |
|
||||
# | Wan2.1-VACE-1.3B | 0.05~0.10 | Wan2.1-VACE-14B | 0.10~0.15 |
|
||||
# # --------------------------------------------------------------------------------------------------- #
|
||||
teacache_threshold = 0.10
|
||||
# The number of steps to skip TeaCache at the beginning of the inference process, which can
|
||||
# reduce the impact of TeaCache on generated video quality.
|
||||
num_skip_start_steps = 5
|
||||
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
|
||||
teacache_offload = False
|
||||
|
||||
# Skip some cfg steps in inference for acceleration
|
||||
# Recommended to be set between 0.00 and 0.25
|
||||
cfg_skip_ratio = 0
|
||||
|
||||
# Riflex config
|
||||
enable_riflex = False
|
||||
# Index of intrinsic frequency
|
||||
riflex_k = 6
|
||||
|
||||
# Config and model path
|
||||
config_path = "config/wan2.2/wan_civitai_t2v.yaml"
|
||||
# model path
|
||||
model_name = "models/Diffusion_Transformer/Wan2.2-VACE-Fun-A14B"
|
||||
|
||||
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
|
||||
sampler_name = "Flow"
|
||||
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
|
||||
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
|
||||
shift = 12.0
|
||||
|
||||
# Load pretrained model if need
|
||||
# The transformer_path is used for low noise model, the transformer_high_path is used for high noise model.
|
||||
transformer_path = None
|
||||
transformer_high_path = None
|
||||
vae_path = None
|
||||
# Load lora model if need
|
||||
# The lora_path is used for low noise model, the lora_high_path is used for high noise model.
|
||||
lora_path = None
|
||||
lora_high_path = None
|
||||
|
||||
# Other params
|
||||
sample_size = [480, 832]
|
||||
video_length = 81
|
||||
fps = 16
|
||||
|
||||
# Use torch.float16 if GPU does not support torch.bfloat16
|
||||
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
|
||||
weight_dtype = torch.bfloat16
|
||||
control_video = None
|
||||
start_image = None
|
||||
end_image = None
|
||||
# Use inpaint video instead of start image and end image.
|
||||
inpaint_video = "asset/inpaint_video.mp4"
|
||||
inpaint_video_mask = "asset/inpaint_video_mask.mp4"
|
||||
subject_ref_images = None
|
||||
vace_context_scale = 1.00
|
||||
# Sometimes, when generating a video from a reference image, white borders appear.
|
||||
# Because the padding is mistakenly treated as part of the image.
|
||||
# If the aspect ratio of the reference image is close to the final output, you can omit the white padding.
|
||||
padding_in_subject_ref_images = True
|
||||
|
||||
# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性
|
||||
# 在neg prompt中添加"安静,固定"等词语可以增加动态性。
|
||||
prompt = "一只棕色的兔子舔了一下它的舌头,坐在舒适房间里的浅色沙发上。在兔子的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。"
|
||||
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
# Using longer neg prompt such as "Blurring, mutation, deformation, distortion, dark and solid, comics, text subtitles, line art." can increase stability
|
||||
# Adding words such as "quiet, solid" to the neg prompt can increase dynamism.
|
||||
# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical."
|
||||
# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code."
|
||||
guidance_scale = 5.0
|
||||
seed = 43
|
||||
num_inference_steps = 50
|
||||
# The lora_weight is used for low noise model, the lora_high_weight is used for high noise model.
|
||||
lora_weight = 0.55
|
||||
lora_high_weight = 0.55
|
||||
save_path = "samples/vace-videos-fun"
|
||||
|
||||
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
|
||||
config = OmegaConf.load(config_path)
|
||||
boundary = config['transformer_additional_kwargs'].get('boundary', 0.875)
|
||||
|
||||
transformer = VaceWanTransformer3DModel.from_pretrained(
|
||||
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')),
|
||||
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=weight_dtype,
|
||||
)
|
||||
if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe":
|
||||
transformer_2 = VaceWanTransformer3DModel.from_pretrained(
|
||||
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
|
||||
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=weight_dtype,
|
||||
)
|
||||
else:
|
||||
transformer_2 = None
|
||||
|
||||
if transformer_path is not None:
|
||||
print(f"From checkpoint: {transformer_path}")
|
||||
if transformer_path.endswith("safetensors"):
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(transformer_path)
|
||||
else:
|
||||
state_dict = torch.load(transformer_path, map_location="cpu")
|
||||
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
|
||||
|
||||
m, u = transformer.load_state_dict(state_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
|
||||
if transformer_2 is not None:
|
||||
if transformer_high_path is not None:
|
||||
print(f"From checkpoint: {transformer_high_path}")
|
||||
if transformer_high_path.endswith("safetensors"):
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(transformer_high_path)
|
||||
else:
|
||||
state_dict = torch.load(transformer_high_path, map_location="cpu")
|
||||
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
|
||||
|
||||
m, u = transformer_2.load_state_dict(state_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
|
||||
# Get Vae
|
||||
Chosen_AutoencoderKL = {
|
||||
"AutoencoderKLWan": AutoencoderKLWan,
|
||||
"AutoencoderKLWan3_8": AutoencoderKLWan3_8
|
||||
}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')]
|
||||
vae = Chosen_AutoencoderKL.from_pretrained(
|
||||
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
|
||||
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
|
||||
).to(weight_dtype)
|
||||
|
||||
if vae_path is not None:
|
||||
print(f"From checkpoint: {vae_path}")
|
||||
if vae_path.endswith("safetensors"):
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(vae_path)
|
||||
else:
|
||||
state_dict = torch.load(vae_path, map_location="cpu")
|
||||
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
|
||||
|
||||
m, u = vae.load_state_dict(state_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
|
||||
# Get Tokenizer
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
|
||||
)
|
||||
|
||||
# Get Text encoder
|
||||
text_encoder = WanT5EncoderModel.from_pretrained(
|
||||
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
|
||||
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
|
||||
low_cpu_mem_usage=True,
|
||||
torch_dtype=weight_dtype,
|
||||
)
|
||||
text_encoder = text_encoder.eval()
|
||||
|
||||
# Get Scheduler
|
||||
Chosen_Scheduler = scheduler_dict = {
|
||||
"Flow": FlowMatchEulerDiscreteScheduler,
|
||||
"Flow_Unipc": FlowUniPCMultistepScheduler,
|
||||
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
|
||||
}[sampler_name]
|
||||
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
|
||||
config['scheduler_kwargs']['shift'] = 1
|
||||
scheduler = Chosen_Scheduler(
|
||||
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
|
||||
)
|
||||
|
||||
# Get Pipeline
|
||||
pipeline = Wan2_2VaceFunPipeline(
|
||||
transformer=transformer,
|
||||
transformer_2=transformer_2,
|
||||
vae=vae,
|
||||
tokenizer=tokenizer,
|
||||
text_encoder=text_encoder,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
if ulysses_degree > 1 or ring_degree > 1:
|
||||
from functools import partial
|
||||
transformer.enable_multi_gpus_inference()
|
||||
if transformer_2 is not None:
|
||||
transformer_2.enable_multi_gpus_inference()
|
||||
if fsdp_dit:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.transformer = shard_fn(pipeline.transformer)
|
||||
if transformer_2 is not None:
|
||||
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
|
||||
print("Add FSDP DIT")
|
||||
if fsdp_text_encoder:
|
||||
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
|
||||
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
|
||||
print("Add FSDP TEXT ENCODER")
|
||||
|
||||
if compile_dit:
|
||||
for i in range(len(pipeline.transformer.blocks)):
|
||||
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
|
||||
if transformer_2 is not None:
|
||||
for i in range(len(pipeline.transformer_2.blocks)):
|
||||
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
|
||||
print("Add Compile")
|
||||
|
||||
if GPU_memory_mode == "sequential_cpu_offload":
|
||||
replace_parameters_by_name(transformer, ["modulation",], device=device)
|
||||
transformer.freqs = transformer.freqs.to(device=device)
|
||||
if transformer_2 is not None:
|
||||
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
|
||||
transformer_2.freqs = transformer_2.freqs.to(device=device)
|
||||
pipeline.enable_sequential_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload":
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
if transformer_2 is not None:
|
||||
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
|
||||
pipeline.to(device=device)
|
||||
else:
|
||||
pipeline.to(device=device)
|
||||
|
||||
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
|
||||
if coefficients is not None:
|
||||
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
|
||||
pipeline.transformer.enable_teacache(
|
||||
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
|
||||
)
|
||||
if transformer_2 is not None:
|
||||
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
|
||||
|
||||
if cfg_skip_ratio is not None:
|
||||
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
|
||||
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
|
||||
if transformer_2 is not None:
|
||||
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
|
||||
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
if transformer_2 is not None:
|
||||
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
|
||||
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_2 is not None:
|
||||
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
|
||||
|
||||
if subject_ref_images is not None:
|
||||
subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=padding_in_subject_ref_images) for _subject_ref_image in subject_ref_images]
|
||||
subject_ref_images = torch.cat(subject_ref_images, dim=2)
|
||||
|
||||
if inpaint_video is not None:
|
||||
if inpaint_video_mask is None:
|
||||
raise ValueError("inpaint_video_mask is required when inpaint_video is provided")
|
||||
inpaint_video, _, _, _ = get_video_to_video_latent(inpaint_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask, _, _, _ = get_video_to_video_latent(inpaint_video_mask, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
inpaint_video_mask = inpaint_video_mask[:, :1]
|
||||
else:
|
||||
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size)
|
||||
|
||||
control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
|
||||
|
||||
sample = pipeline(
|
||||
prompt,
|
||||
num_frames = video_length,
|
||||
negative_prompt = negative_prompt,
|
||||
height = sample_size[0],
|
||||
width = sample_size[1],
|
||||
generator = generator,
|
||||
guidance_scale = guidance_scale,
|
||||
num_inference_steps = num_inference_steps,
|
||||
|
||||
video = inpaint_video,
|
||||
mask_video = inpaint_video_mask,
|
||||
control_video = control_video,
|
||||
subject_ref_images = subject_ref_images,
|
||||
boundary = boundary,
|
||||
shift = shift,
|
||||
vace_context_scale = vace_context_scale
|
||||
).videos
|
||||
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
|
||||
if transformer_2 is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
|
||||
|
||||
def save_results():
|
||||
if not os.path.exists(save_path):
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
|
||||
index = len([path for path in os.listdir(save_path)]) + 1
|
||||
prefix = str(index).zfill(8)
|
||||
if video_length == 1:
|
||||
video_path = os.path.join(save_path, prefix + ".png")
|
||||
|
||||
image = sample[0, :, 0]
|
||||
image = image.transpose(0, 1).transpose(1, 2)
|
||||
image = (image * 255).numpy().astype(np.uint8)
|
||||
image = Image.fromarray(image)
|
||||
image.save(video_path)
|
||||
else:
|
||||
video_path = os.path.join(save_path, prefix + ".mp4")
|
||||
save_videos_grid(sample, video_path, fps=fps)
|
||||
|
||||
if ulysses_degree * ring_degree > 1:
|
||||
import torch.distributed as dist
|
||||
if dist.get_rank() == 0:
|
||||
save_results()
|
||||
else:
|
||||
save_results()
|
||||
Executable
+165
@@ -0,0 +1,165 @@
|
||||
## Training Code
|
||||
|
||||
We can choose whether to use deepspeed or fsdp in flux, which can save a lot of video memory.
|
||||
|
||||
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 at the center, but instead, it trains the entire images after grouping them into buckets based on resolution.
|
||||
- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum.
|
||||
- For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`
|
||||
- `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.
|
||||
|
||||
Without deepspeed:
|
||||
|
||||
Training flux without DeepSpeed may result in insufficient GPU memory.
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/flux/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
With deepspeed zero-2:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# 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/flux/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
Deepspeed zero-3:
|
||||
|
||||
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/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/flux/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
With FSDP:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap FluxSingleTransformerBlock,FluxTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/flux/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
Executable
+153
@@ -0,0 +1,153 @@
|
||||
## Lora Training Code
|
||||
|
||||
We can choose whether to use deepspeed or fsdp in flux, which can save a lot of video memory.
|
||||
|
||||
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 at the center, but instead, it trains the entire images after grouping them into buckets based on resolution.
|
||||
- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum.
|
||||
- For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`
|
||||
- `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.
|
||||
|
||||
Without deepspeed:
|
||||
|
||||
Training flux without DeepSpeed may result in insufficient GPU memory.
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling
|
||||
```
|
||||
|
||||
With deepspeed zero-2:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# 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/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling
|
||||
```
|
||||
|
||||
Deepspeed zero-3:
|
||||
|
||||
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/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling
|
||||
```
|
||||
|
||||
With FSDP:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap FluxSingleTransformerBlock,FluxTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--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 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling
|
||||
```
|
||||
@@ -1576,7 +1576,6 @@ def main():
|
||||
|
||||
# Predict the noise residual
|
||||
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device):
|
||||
print(noisy_latents.size(), prompt_embeds.size(), pooled_prompt_embeds.size(), text_ids.size(), latent_image_ids.size())
|
||||
noise_pred = transformer3d(
|
||||
hidden_states=noisy_latents,
|
||||
timestep=timesteps / 1000,
|
||||
|
||||
@@ -1571,7 +1571,6 @@ def main():
|
||||
|
||||
# Predict the noise residual
|
||||
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device):
|
||||
print(noisy_latents.size(), prompt_embeds.size(), pooled_prompt_embeds.size(), text_ids.size(), latent_image_ids.size())
|
||||
noise_pred = transformer3d(
|
||||
hidden_states=noisy_latents,
|
||||
timestep=timesteps / 1000,
|
||||
|
||||
Executable
+257
@@ -0,0 +1,257 @@
|
||||
## Training Code
|
||||
|
||||
We can choose whether to use deepspeed or fsdp 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",
|
||||
"object_file_path": ["object/1.jpg", "object/2.jpg"],
|
||||
"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",
|
||||
"object_file_path": ["object/3.jpg", "object/4.jpg"],
|
||||
"text": "Ba Da Ba Ba Ba Ba.",
|
||||
"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=256`, and `image_sample_size=1024`, 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=256`, and `image_sample_size=1024`, 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=256`, and `image_sample_size=1024`, 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.
|
||||
- `train_mode` is used to set the training mode.
|
||||
- The models named `Wan2.1-Fun-*-Control` are trained in the `control_ref` mode.
|
||||
- The models named `Wan2.1-Fun-*-Control-Camera` are trained in the `control_ref_camera` mode.
|
||||
- `control_ref_image` is used to specify the type of control image. The available options are `first_frame` and `random`.
|
||||
- `first_frame` is used in V1.0 because V1.0 supports using a specified start frame as the control image. The Control-Camera models use the first frame as the control image.
|
||||
- `random` is used in V1.1 because V1.1 supports both using a specified start frame and a reference image as the control image.
|
||||
- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode).
|
||||
|
||||
When train model with multi machines, please set the params as follows:
|
||||
```sh
|
||||
export MASTER_ADDR="your master address"
|
||||
export MASTER_PORT=10086
|
||||
export WORLD_SIZE=1 # The number of machines
|
||||
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
|
||||
export RANK=0 # The rank of this machine
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
|
||||
```
|
||||
|
||||
Wan-Fun-Control without deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-VACE-Fun-A14B"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.2_vace_fun/train.py \
|
||||
--config_path="config/wan2.2/wan_civitai_t2v.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=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--low_vram \
|
||||
--control_ref_image="random" \
|
||||
--boundary_type="low" \
|
||||
--trainable_modules "vace"
|
||||
```
|
||||
|
||||
Wan-Fun-Control with deepspeed:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-VACE-Fun-A14B"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# 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.2_vace_fun/train.py \
|
||||
--config_path="config/wan2.2/wan_civitai_t2v.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=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--low_vram \
|
||||
--control_ref_image="random" \
|
||||
--boundary_type="low" \
|
||||
--trainable_modules "vace"
|
||||
```
|
||||
|
||||
Wan-Fun-Control with deepspeed zero-3:
|
||||
|
||||
Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. 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.2-VACE-Fun-A14B"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# 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_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.2_vace_fun/train.py \
|
||||
--config_path="config/wan2.2/wan_civitai_t2v.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=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--low_vram \
|
||||
--control_ref_image="random" \
|
||||
--boundary_type="low" \
|
||||
--trainable_modules "vace"
|
||||
```
|
||||
|
||||
Wan-Fun-Control with FSDP:
|
||||
|
||||
Wan with FSDP is suitable for 14B Wan at high resolutions. Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-VACE-Fun-A14B"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=VaceWanAttentionBlock,BaseWanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.2_vace_fun/train.py \
|
||||
--config_path="config/wan2.2/wan_civitai_t2v.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=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--low_vram \
|
||||
--control_ref_image="random" \
|
||||
--boundary_type="low" \
|
||||
--trainable_modules "vace"
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,43 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-VACE-Fun-A14B"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/wan2.2_vace_fun/train.py \
|
||||
--config_path="config/wan2.2/wan_civitai_t2v.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=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--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 \
|
||||
--low_vram \
|
||||
--control_ref_image="random" \
|
||||
--boundary_type="low" \
|
||||
--trainable_modules "vace"
|
||||
@@ -3,6 +3,7 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import Any, Dict
|
||||
|
||||
import os
|
||||
import math
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
@@ -14,6 +15,8 @@ from .wan_transformer3d import (WanAttentionBlock, WanTransformer3DModel,
|
||||
sinusoidal_embedding_1d)
|
||||
|
||||
|
||||
VIDEOX_OFFLOAD_VACE_LATENTS = os.environ.get("VIDEOX_OFFLOAD_VACE_LATENTS", False)
|
||||
|
||||
class VaceWanAttentionBlock(WanAttentionBlock):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -45,8 +48,16 @@ class VaceWanAttentionBlock(WanAttentionBlock):
|
||||
all_c = list(torch.unbind(c))
|
||||
c = all_c.pop(-1)
|
||||
|
||||
if VIDEOX_OFFLOAD_VACE_LATENTS:
|
||||
c = c.to(x.device)
|
||||
|
||||
c = super().forward(c, **kwargs)
|
||||
c_skip = self.after_proj(c)
|
||||
|
||||
if VIDEOX_OFFLOAD_VACE_LATENTS:
|
||||
c_skip = c_skip.to("cpu")
|
||||
c = c.to("cpu")
|
||||
|
||||
all_c += [c_skip, c]
|
||||
c = torch.stack(all_c)
|
||||
return c
|
||||
@@ -71,7 +82,10 @@ class BaseWanAttentionBlock(WanAttentionBlock):
|
||||
def forward(self, x, hints, context_scale=1.0, **kwargs):
|
||||
x = super().forward(x, **kwargs)
|
||||
if self.block_id is not None:
|
||||
x = x + hints[self.block_id] * context_scale
|
||||
if VIDEOX_OFFLOAD_VACE_LATENTS:
|
||||
x = x + hints[self.block_id].to(x.device) * context_scale
|
||||
else:
|
||||
x = x + hints[self.block_id] * context_scale
|
||||
return x
|
||||
|
||||
|
||||
|
||||
@@ -383,6 +383,14 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
|
||||
key = key.replace(".self_attn.", "_self_attn_")
|
||||
key = key.replace(".cross_attn.", "_cross_attn_")
|
||||
key = key.replace(".ffn.", "_ffn_")
|
||||
if "lora_A" in key or "lora_B" in key:
|
||||
key = "lora_unet__" + key
|
||||
key = key.replace("blocks.", "blocks_")
|
||||
key = key.replace(".self_attn.", "_self_attn_")
|
||||
key = key.replace(".cross_attn.", "_cross_attn_")
|
||||
key = key.replace(".ffn.", "_ffn_")
|
||||
key = key.replace(".lora_A.default.", ".lora_down.")
|
||||
key = key.replace(".lora_B.default.", ".lora_up.")
|
||||
layer, elem = key.split('.', 1)
|
||||
updates[layer][elem] = value
|
||||
|
||||
@@ -496,6 +504,14 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
|
||||
key = key.replace(".self_attn.", "_self_attn_")
|
||||
key = key.replace(".cross_attn.", "_cross_attn_")
|
||||
key = key.replace(".ffn.", "_ffn_")
|
||||
if "lora_A" in key or "lora_B" in key:
|
||||
key = "lora_unet__" + key
|
||||
key = key.replace("blocks.", "blocks_")
|
||||
key = key.replace(".self_attn.", "_self_attn_")
|
||||
key = key.replace(".cross_attn.", "_cross_attn_")
|
||||
key = key.replace(".ffn.", "_ffn_")
|
||||
key = key.replace(".lora_A.default.", ".lora_down.")
|
||||
key = key.replace(".lora_B.default.", ".lora_up.")
|
||||
layer, elem = key.split('.', 1)
|
||||
updates[layer][elem] = value
|
||||
|
||||
|
||||
Reference in New Issue
Block a user