Fix bug in wan-animate and fantasytalking multi gpus inference && Update Training Codes && Fix bug in s2v lora merging && Update qwen image quick loading (#368)

This commit is contained in:
Bubbliiiing
2025-11-11 14:16:41 +08:00
committed by GitHub
parent df77df019e
commit 7fd9594919
41 changed files with 4293 additions and 122 deletions
+2 -1
View File
@@ -345,7 +345,8 @@ with torch.no_grad():
if lora_path is not None:
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")
if transformer_2 is not None:
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):
+2 -2
View File
@@ -1229,9 +1229,9 @@ def main():
ema_transformer3d.to(accelerator.device)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+2 -2
View File
@@ -1164,9 +1164,9 @@ def main():
ema_transformer3d.to(accelerator.device)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+2 -2
View File
@@ -1164,10 +1164,10 @@ def main():
)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1357,7 +1357,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
audio_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=torch.float32)
+2 -2
View File
@@ -1348,8 +1348,8 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder_2.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
text_encoder_2.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+3 -3
View File
@@ -1280,11 +1280,11 @@ def main():
# text_encoder_2 = shard_fn(text_encoder_2)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder_2.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
text_encoder_2.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1215,7 +1215,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1260,7 +1260,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+2 -2
View File
@@ -1209,10 +1209,10 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+2 -2
View File
@@ -1157,10 +1157,10 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1405,7 +1405,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if args.train_mode != "normal":
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
+3 -3
View File
@@ -1337,12 +1337,12 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if args.train_mode != "normal":
clip_image_encoder.to(accelerator.device, dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1402,7 +1402,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if args.train_mode != "normal":
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
+1 -1
View File
@@ -1408,7 +1408,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
+3 -3
View File
@@ -1350,11 +1350,11 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
clip_image_encoder.to(accelerator.device, dtype=weight_dtype)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+3 -3
View File
@@ -1338,12 +1338,12 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if args.train_mode != "normal":
clip_image_encoder.to(accelerator.device, dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1054,7 +1054,7 @@ def main():
vae.to(accelerator.device, dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
text_encoder.to(accelerator.device)
clip_image_encoder.to(accelerator.device, dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(prompt_list) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1404,7 +1404,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1439,7 +1439,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -33
View File
@@ -679,38 +679,6 @@ def parse_args():
'The initial gradient is relative to the multiple of the max_grad_norm. '
),
)
parser.add_argument(
"--train_mode",
type=str,
default="control",
help=(
'The format of training data. Support `"control"`'
' (default), `"control_ref"`, `"control_camera_ref"`.'
),
)
parser.add_argument(
"--control_ref_image",
type=str,
default="first_frame",
help=(
'The format of training data. Support `"first_frame"`'
' (default), `"random"`.'
),
)
parser.add_argument(
"--add_full_ref_image_in_self_attention",
action="store_true",
help=(
'Whether enable add full ref image in self attention.'
),
)
parser.add_argument(
"--add_inpaint_info",
action="store_true",
help=(
'Whether enable add inpaint info in self attention.'
),
)
parser.add_argument(
"--weighting_scheme",
type=str,
@@ -1464,7 +1432,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
clip_image_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
+2 -3
View File
@@ -11,9 +11,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate.py \
--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_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
File diff suppressed because it is too large Load Diff
+38
View File
@@ -0,0 +1,38 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Animate-14B/"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.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/train_animate_lora.py \
--config_path="config/wan2.2/wan_civitai_animate.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--boundary_type="full" \
--low_vram
+2 -2
View File
@@ -1383,10 +1383,10 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1460,7 +1460,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
audio_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=torch.float32)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
File diff suppressed because it is too large Load Diff
+39
View File
@@ -0,0 +1,39 @@
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-S2V-14B"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.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/train_s2v_lora.py \
--config_path="config/wan2.2/wan_civitai_s2v.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--boundary_type="full" \
--control_ref_image="random" \
--low_vram
+1 -1
View File
@@ -1456,7 +1456,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1513,7 +1513,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+2 -2
View File
@@ -1458,10 +1458,10 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+2 -2
View File
@@ -1384,10 +1384,10 @@ def main():
text_encoder = shard_fn(text_encoder)
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device, dtype=weight_dtype)
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1461,7 +1461,7 @@ def main():
# Move text_encode and vae to gpu and cast to weight_dtype
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu")
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
+1 -1
View File
@@ -1,5 +1,5 @@
from .dataset_image import CC15M, ImageEditDataset
from .dataset_image_video import (ImageVideoControlDataset, ImageVideoDataset,
from .dataset_image_video import (ImageVideoControlDataset, ImageVideoDataset, TextDataset,
ImageVideoSampler)
from .dataset_video import VideoDataset, VideoSpeechDataset, VideoAnimateDataset, WebVid10M
from .utils import (VIDEO_READER_TIMEOUT, Camera, VideoReader_contextmanager,
+35 -2
View File
@@ -530,7 +530,7 @@ class ImageVideoControlDataset(Dataset):
shuffle(subject_id)
subject_images = []
for i in range(min(len(subject_id), 4)):
subject_image = Image.open(subject_id[i])
subject_image = Image.open(subject_id[i]).convert('RGB')
width, height = subject_image.size
total_pixels = width * height
@@ -621,4 +621,37 @@ class ImageVideoSafetensorsDataset(Dataset):
else:
path = os.path.join(self.data_root, self.dataset[idx]["file_path"])
state_dict = load_file(path)
return state_dict
return state_dict
class TextDataset(Dataset):
def __init__(self, ann_path, text_drop_ratio=0.0):
print(f"loading annotations from {ann_path} ...")
with open(ann_path, 'r') as f:
self.dataset = json.load(f)
self.length = len(self.dataset)
print(f"data scale: {self.length}")
self.text_drop_ratio = text_drop_ratio
def __len__(self):
return self.length
def __getitem__(self, idx):
while True:
try:
item = self.dataset[idx]
text = item['text']
# Randomly drop text (for classifier-free guidance)
if random.random() < self.text_drop_ratio:
text = ''
sample = {
"text": text,
"idx": idx
}
return sample
except Exception as e:
print(f"Error at index {idx}: {e}, retrying with random index...")
idx = np.random.randint(0, self.length - 1)
@@ -4,6 +4,7 @@ import math
import os
from typing import Any, Dict
import numpy as np
import torch
import torch.cuda.amp as amp
import torch.nn as nn
@@ -45,6 +46,10 @@ class AudioCrossAttentionProcessor(nn.Module):
nn.init.zeros_(self.k_proj.weight)
nn.init.zeros_(self.v_proj.weight)
self.sp_world_size = 1
self.sp_world_rank = 0
self.all_gather = None
def __call__(
self,
attn: nn.Module,
@@ -80,7 +85,14 @@ class AudioCrossAttentionProcessor(nn.Module):
img_x = img_x.flatten(2)
if len(audio_proj.shape) == 4:
q = sequence_parallel_all_gather(q, dim=1)
if self.sp_world_size > 1:
q = self.all_gather(q, dim=1)
length = int(np.floor(q.size()[1] / latents_num_frames) * latents_num_frames)
origin_length = q.size()[1]
if origin_length > length:
q_pad = q[:, length:]
q = q[:, :length]
audio_q = q.view(b * latents_num_frames, -1, n, d) # [b, 21, l1, n, d]
ip_key = self.k_proj(audio_proj).view(b * latents_num_frames, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b * latents_num_frames, -1, n, d)
@@ -88,8 +100,11 @@ class AudioCrossAttentionProcessor(nn.Module):
audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_type="NORMAL"
)
audio_x = audio_x.view(b, q.size(1), n, d)
if self.sp_world_size > 1:
if origin_length > length:
audio_x = torch.cat([audio_x, q_pad], dim=1)
audio_x = torch.chunk(audio_x, self.sp_world_size, dim=1)[self.sp_world_rank]
audio_x = audio_x.flatten(2)
audio_x = sequence_parallel_chunk(audio_x, dim=1)
elif len(audio_proj.shape) == 3:
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
@@ -361,6 +376,14 @@ class FantasyTalkingTransformer3DModel(WanTransformer3DModel):
k_lens_list, dtype=torch.long
)
def enable_multi_gpus_inference(self,):
super().enable_multi_gpus_inference()
for name, module in self.named_modules():
if module.__class__.__name__ == 'AudioCrossAttentionProcessor':
module.sp_world_size = self.sp_world_size
module.sp_world_rank = self.sp_world_rank
module.all_gather = self.all_gather
@cfg_skip()
def forward(
self,
+55 -13
View File
@@ -994,30 +994,72 @@ class QwenImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fro
for key in _state_dict:
state_dict[key] = _state_dict[key]
filtered_state_dict = {}
for key in state_dict:
if key in model.state_dict() and model.state_dict()[key].size() == state_dict[key].size():
filtered_state_dict[key] = state_dict[key]
else:
print(f"Skipping key '{key}' due to size mismatch or absence in model.")
model_keys = set(model.state_dict().keys())
loaded_keys = set(filtered_state_dict.keys())
missing_keys = model_keys - loaded_keys
def initialize_missing_parameters(missing_keys, model_state_dict, torch_dtype=None):
initialized_dict = {}
with torch.no_grad():
for key in missing_keys:
param_shape = model_state_dict[key].shape
param_dtype = torch_dtype if torch_dtype is not None else model_state_dict[key].dtype
if 'weight' in key:
if any(norm_type in key for norm_type in ['norm', 'ln_', 'layer_norm', 'group_norm', 'batch_norm']):
initialized_dict[key] = torch.ones(param_shape, dtype=param_dtype)
elif 'embedding' in key or 'embed' in key:
initialized_dict[key] = torch.randn(param_shape, dtype=param_dtype) * 0.02
elif 'head' in key or 'output' in key or 'proj_out' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
elif len(param_shape) >= 2:
initialized_dict[key] = torch.empty(param_shape, dtype=param_dtype)
nn.init.xavier_uniform_(initialized_dict[key])
else:
initialized_dict[key] = torch.randn(param_shape, dtype=param_dtype) * 0.02
elif 'bias' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
elif 'running_mean' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
elif 'running_var' in key:
initialized_dict[key] = torch.ones(param_shape, dtype=param_dtype)
elif 'num_batches_tracked' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=torch.long)
else:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
return initialized_dict
if missing_keys:
print(f"Missing keys will be initialized: {sorted(missing_keys)}")
initialized_params = initialize_missing_parameters(
missing_keys,
model.state_dict(),
torch_dtype
)
filtered_state_dict.update(initialized_params)
if diffusers_version >= "0.33.0":
# Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit:
# https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785.
load_model_dict_into_meta(
model,
state_dict,
filtered_state_dict,
dtype=torch_dtype,
model_name_or_path=pretrained_model_path,
)
else:
model._convert_deprecated_attention_blocks(state_dict)
# move the params from meta device to cpu
missing_keys = set(model.state_dict().keys()) - set(state_dict.keys())
if len(missing_keys) > 0:
raise ValueError(
f"Cannot load {cls} from {pretrained_model_path} because the following keys are"
f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass"
" `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize"
" those weights or else make sure your checkpoint file is correct."
)
model._convert_deprecated_attention_blocks(filtered_state_dict)
unexpected_keys = load_model_dict_into_meta(
model,
state_dict,
filtered_state_dict,
device=param_device,
dtype=torch_dtype,
model_name_or_path=pretrained_model_path,
+19 -5
View File
@@ -4,6 +4,7 @@ from typing import Optional, Tuple
import torch
import torch.nn.functional as F
import numpy as np
from einops import rearrange
from torch import nn
@@ -337,8 +338,11 @@ class FaceBlock(nn.Module):
motion_vec: torch.Tensor,
motion_mask: Optional[torch.Tensor] = None,
use_context_parallel=False,
all_gather=None,
sp_world_size=1,
sp_world_rank=0,
) -> torch.Tensor:
dtype = x.dtype
B, T, N, C = motion_vec.shape
T_comp = T
@@ -358,10 +362,17 @@ class FaceBlock(nn.Module):
k = rearrange(k, "B L N H D -> (B L) N H D")
v = rearrange(v, "B L N H D -> (B L) N H D")
# if use_context_parallel:
# q = gather_forward(q, dim=1)
if use_context_parallel:
q = all_gather(q, dim=1)
length = int(np.floor(q.size()[1] / T_comp) * T_comp)
origin_length = q.size()[1]
if origin_length > length:
q_pad = q[:, length:]
q = q[:, :length]
q = rearrange(q, "B (L S) H D -> (B L) S H D", L=T_comp)
q, k, v = q.to(dtype), k.to(dtype), v.to(dtype)
# Compute attention.
attn = attention(
q,
@@ -372,8 +383,11 @@ class FaceBlock(nn.Module):
)
attn = rearrange(attn, "(B L) S C -> B (L S) C", L=T_comp)
# if use_context_parallel:
# attn = torch.chunk(attn, get_world_size(), dim=1)[get_rank()]
if use_context_parallel:
q_pad = rearrange(q_pad, "B L H D -> B L (H D)")
if origin_length > length:
attn = torch.cat([attn, q_pad], dim=1)
attn = torch.chunk(attn, sp_world_size, dim=1)[sp_world_rank]
output = self.linear2(attn)
+61 -13
View File
@@ -669,6 +669,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.current_steps = 0
self.num_inference_steps = None
self.gradient_checkpointing = False
self.all_gather = None
self.sp_world_size = 1
self.sp_world_rank = 0
self.init_weights()
@@ -1159,30 +1160,77 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
for key in _state_dict:
state_dict[key] = _state_dict[key]
if model.state_dict()['patch_embedding.weight'].size() != state_dict['patch_embedding.weight'].size():
model.state_dict()['patch_embedding.weight'][:, :state_dict['patch_embedding.weight'].size()[1], :, :] = state_dict['patch_embedding.weight'][:, :model.state_dict()['patch_embedding.weight'].size()[1], :, :]
model.state_dict()['patch_embedding.weight'][:, state_dict['patch_embedding.weight'].size()[1]:, :, :] = 0
state_dict['patch_embedding.weight'] = model.state_dict()['patch_embedding.weight']
filtered_state_dict = {}
for key in state_dict:
if key in model.state_dict() and model.state_dict()[key].size() == state_dict[key].size():
filtered_state_dict[key] = state_dict[key]
else:
print(f"Skipping key '{key}' due to size mismatch or absence in model.")
model_keys = set(model.state_dict().keys())
loaded_keys = set(filtered_state_dict.keys())
missing_keys = model_keys - loaded_keys
def initialize_missing_parameters(missing_keys, model_state_dict, torch_dtype=None):
initialized_dict = {}
with torch.no_grad():
for key in missing_keys:
param_shape = model_state_dict[key].shape
param_dtype = torch_dtype if torch_dtype is not None else model_state_dict[key].dtype
if 'weight' in key:
if any(norm_type in key for norm_type in ['norm', 'ln_', 'layer_norm', 'group_norm', 'batch_norm']):
initialized_dict[key] = torch.ones(param_shape, dtype=param_dtype)
elif 'embedding' in key or 'embed' in key:
initialized_dict[key] = torch.randn(param_shape, dtype=param_dtype) * 0.02
elif 'head' in key or 'output' in key or 'proj_out' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
elif len(param_shape) >= 2:
initialized_dict[key] = torch.empty(param_shape, dtype=param_dtype)
nn.init.xavier_uniform_(initialized_dict[key])
else:
initialized_dict[key] = torch.randn(param_shape, dtype=param_dtype) * 0.02
elif 'bias' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
elif 'running_mean' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
elif 'running_var' in key:
initialized_dict[key] = torch.ones(param_shape, dtype=param_dtype)
elif 'num_batches_tracked' in key:
initialized_dict[key] = torch.zeros(param_shape, dtype=torch.long)
else:
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
return initialized_dict
if missing_keys:
print(f"Missing keys will be initialized: {sorted(missing_keys)}")
initialized_params = initialize_missing_parameters(
missing_keys,
model.state_dict(),
torch_dtype
)
filtered_state_dict.update(initialized_params)
if diffusers_version >= "0.33.0":
# Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit:
# https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785.
load_model_dict_into_meta(
model,
state_dict,
filtered_state_dict,
dtype=torch_dtype,
model_name_or_path=pretrained_model_path,
)
else:
model._convert_deprecated_attention_blocks(state_dict)
# move the params from meta device to cpu
missing_keys = set(model.state_dict().keys()) - set(state_dict.keys())
if len(missing_keys) > 0:
raise ValueError(
f"Cannot load {cls} from {pretrained_model_path} because the following keys are"
f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass"
" `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize"
" those weights or else make sure your checkpoint file is correct."
)
model._convert_deprecated_attention_blocks(filtered_state_dict)
unexpected_keys = load_model_dict_into_meta(
model,
state_dict,
filtered_state_dict,
device=param_device,
dtype=torch_dtype,
model_name_or_path=pretrained_model_path,
@@ -46,7 +46,6 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
cross_attn_norm=True,
eps=1e-6,
motion_encoder_dim=512,
use_context_parallel=False,
use_img_emb=True
):
model_type = "i2v" # TODO: Hard code for both preview and official versions.
@@ -54,7 +53,6 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
num_heads, num_layers, window_size, qk_norm, cross_attn_norm, eps)
self.motion_encoder_dim = motion_encoder_dim
self.use_context_parallel = use_context_parallel
self.use_img_emb = use_img_emb
self.pose_patch_embedding = nn.Conv3d(
@@ -100,15 +98,14 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
motion_vec = torch.cat([pad_face, motion_vec], dim=1)
return x, motion_vec
def after_transformer_block(self, block_idx, x, motion_vec, motion_masks=None):
if block_idx % 5 == 0:
adapter_args = [x, motion_vec, motion_masks, self.use_context_parallel]
use_context_parallel = self.sp_world_size > 1
adapter_args = [x, motion_vec, motion_masks, use_context_parallel, self.all_gather, self.sp_world_size, self.sp_world_rank]
residual_out = self.face_adapter.fuser_blocks[block_idx // 5](*adapter_args)
x = residual_out + x
return x
@cfg_skip()
def forward(
self,
@@ -139,6 +136,8 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
if self.sp_world_size > 1:
seq_len = int(math.ceil(seq_len / self.sp_world_size)) * self.sp_world_size
assert seq_lens.max() <= seq_len
x = torch.cat([
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
@@ -227,6 +226,7 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
t,
**ckpt_kwargs,
)
x, motion_vec = x.to(dtype), motion_vec.to(dtype)
x = self.after_transformer_block(idx, x, motion_vec)
else:
# arguments
@@ -241,6 +241,7 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
t=t
)
x = block(x, **kwargs)
x, motion_vec = x.to(dtype), motion_vec.to(dtype)
x = self.after_transformer_block(idx, x, motion_vec)
if cond_flag:
@@ -270,6 +271,7 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
t,
**ckpt_kwargs,
)
x, motion_vec = x.to(dtype), motion_vec.to(dtype)
x = self.after_transformer_block(idx, x, motion_vec)
else:
# arguments
@@ -284,6 +286,7 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
t=t
)
x = block(x, **kwargs)
x, motion_vec = x.to(dtype), motion_vec.to(dtype)
x = self.after_transformer_block(idx, x, motion_vec)
# head
+14 -1
View File
@@ -158,7 +158,8 @@ def precalculate_safetensors_hashes(tensors, metadata):
class LoRANetwork(torch.nn.Module):
TRANSFORMER_TARGET_REPLACE_MODULE = [
"CogVideoXTransformer3DModel", "WanTransformer3DModel", \
"Wan2_2Transformer3DModel", "FluxTransformer2DModel", "QwenImageTransformer2DModel"
"Wan2_2Transformer3DModel", "FluxTransformer2DModel", "QwenImageTransformer2DModel", \
"Wan2_2Transformer3DModel_Animate", "Wan2_2Transformer3DModel_S2V", "FantasyTalkingTransformer3DModel",
]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF", "BertEncoder", "T5SelfAttention", "T5CrossAttention"]
LORA_PREFIX_TRANSFORMER = "lora_unet"
@@ -173,6 +174,7 @@ class LoRANetwork(torch.nn.Module):
dropout: Optional[float] = None,
module_class: Type[object] = LoRAModule,
skip_name: str = None,
target_name: str = None,
varbose: Optional[bool] = False,
) -> None:
super().__init__()
@@ -207,6 +209,15 @@ class LoRANetwork(torch.nn.Module):
if skip_name is not None and skip_name in child_name:
continue
if target_name is not None:
target_name_in = False
if isinstance(target_name, str):
target_name_in = target_name in child_name
elif isinstance(target_name, list):
target_name_in = any([_target_name in child_name for _target_name in target_name])
if not target_name_in:
continue
if is_linear or is_conv2d:
lora_name = prefix + "." + name + "." + child_name
@@ -349,6 +360,7 @@ def create_network(
transformer,
neuron_dropout: Optional[float] = None,
skip_name: str = None,
target_name = None,
**kwargs,
):
if network_dim is None:
@@ -364,6 +376,7 @@ def create_network(
alpha=network_alpha,
dropout=neuron_dropout,
skip_name=skip_name,
target_name=target_name,
varbose=True,
)
return network