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:
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user