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