Update fix sample size && Fix vram memory bug in wan2.1 lora training (#227)

This commit is contained in:
Bubbliiiing
2025-06-24 10:09:23 +08:00
committed by GitHub
parent ba1da31fb5
commit 86028a9ce6
14 changed files with 464 additions and 182 deletions
+4 -4
View File
@@ -1738,10 +1738,10 @@ def main():
if args.use_ema:
ema_transformer3d.copy_to(transformer3d.parameters())
if args.use_deepspeed or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
if args.use_deepspeed or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
accelerator.end_training()
+5 -5
View File
@@ -944,7 +944,7 @@ def main():
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
video_repeat=args.video_repeat,
image_sample_size=args.image_sample_size,
enable_bucket=args.enable_bucket, enable_inpaint=False,
enable_bucket=args.enable_bucket,
)
if args.enable_bucket:
@@ -1629,10 +1629,10 @@ def main():
if args.use_ema:
ema_transformer3d.copy_to(transformer3d.parameters())
if args.use_deepspeed or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
if args.use_deepspeed or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
accelerator.end_training()
+6 -5
View File
@@ -1697,11 +1697,12 @@ def main():
# Create the pipeline using the trained modules and save it.
accelerator.wait_for_everyone()
if accelerator.is_main_process:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
if args.save_state:
if args.use_deepspeed or accelerator.is_main_process:
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
logger.info(f"Saved state to {accelerator_save_path}")
+58 -18
View File
@@ -606,7 +606,12 @@ def parse_args():
"--image_sample_size",
type=int,
default=512,
help="Sample size of the video.",
help="Sample size of the image.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn."
)
parser.add_argument(
"--video_sample_stride",
@@ -768,6 +773,9 @@ def main():
print(f"Using DeepSpeed Zero stage: {zero_stage}")
args.use_deepspeed = True
if zero_stage == 3:
print(f"Auto set save_state to True because zero_stage == 3")
args.save_state = True
elif fsdp_plugin is not None:
from torch.distributed.fsdp import ShardingStrategy
zero_stage = 0
@@ -782,6 +790,9 @@ def main():
print(f"Using FSDP stage: {fsdp_stage}")
args.use_fsdp = True
if fsdp_stage == 3:
print(f"Auto set save_state to True because fsdp_stage == 3")
args.save_state = True
else:
zero_stage = 0
fsdp_stage = 0
@@ -875,6 +886,13 @@ def main():
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Clip Image Encoder
if args.train_mode != "normal":
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Get Transformer
transformer3d = WanTransformer3DModel.from_pretrained(
@@ -882,13 +900,6 @@ def main():
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
if args.train_mode != "normal":
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Freeze vae and text_encoder and set transformer3d to trainable
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
@@ -1120,6 +1131,13 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
# Get the dataset
train_dataset = ImageVideoDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
@@ -1204,9 +1222,9 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
@@ -1216,9 +1234,24 @@ def main():
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
for example in examples:
if args.random_ratio_crop:
if args.fix_sample_size is not None:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
# Get adapt hw for resize
fix_sample_size = list(map(lambda x: int(x), fix_sample_size))
transform = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
elif args.random_ratio_crop:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
@@ -1341,6 +1374,12 @@ def main():
transformer3d, optimizer, train_dataloader, lr_scheduler
)
if fsdp_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
text_encoder = shard_fn(text_encoder)
if args.use_ema:
ema_transformer3d.to(accelerator.device)
@@ -1365,6 +1404,7 @@ def main():
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1465,7 +1505,7 @@ def main():
pixel_values = batch["pixel_values"].to(weight_dtype)
# Increase the batch size when the length of the latent sequence of the current sample is small
if args.training_with_video_token_length and not zero_stage == 3:
if args.training_with_video_token_length and zero_stage != 3:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
if args.enable_text_encoder_in_dataloader:
@@ -1486,7 +1526,7 @@ def main():
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
mask = batch["mask"].to(weight_dtype)
# Increase the batch size when the length of the latent sequence of the current sample is small
if args.training_with_video_token_length and not zero_stage == 3:
if args.training_with_video_token_length and zero_stage != 3:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
@@ -1856,10 +1896,10 @@ def main():
if args.use_ema:
ema_transformer3d.copy_to(transformer3d.parameters())
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
accelerator.end_training()
+84 -30
View File
@@ -625,7 +625,12 @@ def parse_args():
"--image_sample_size",
type=int,
default=512,
help="Sample size of the video.",
help="Sample size of the image.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn."
)
parser.add_argument(
"--video_sample_stride",
@@ -653,12 +658,6 @@ def parse_args():
"The config of the model in training."
),
)
parser.add_argument(
"--image_repeat_in_forward",
type=int,
default=0,
help="Num of repeat image in forward.",
)
parser.add_argument(
"--transformer_path",
type=str,
@@ -676,7 +675,7 @@ def parse_args():
parser.add_argument(
'--tokenizer_max_length',
type=int,
default=226,
default=512,
help='Max length of tokenizer'
)
parser.add_argument(
@@ -694,7 +693,7 @@ def parse_args():
default="normal",
help=(
'The format of training data. Support `"normal"`'
' (default), `"inpaint"`.'
' (default), `"i2v"`.'
),
)
parser.add_argument(
@@ -716,6 +715,12 @@ def parse_args():
default=1.29,
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
parser.add_argument(
"--lora_skip_name",
type=str,
default=None,
help=("The module is not trained in loras. "),
)
args = parser.parse_args()
env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
@@ -767,6 +772,9 @@ def main():
print(f"Using DeepSpeed Zero stage: {zero_stage}")
args.use_deepspeed = True
if zero_stage == 3:
print(f"Auto set save_state to True because zero_stage == 3")
args.save_state = True
elif fsdp_plugin is not None:
from torch.distributed.fsdp import ShardingStrategy
zero_stage = 0
@@ -781,6 +789,9 @@ def main():
print(f"Using FSDP stage: {fsdp_stage}")
args.use_fsdp = True
if fsdp_stage == 3:
print(f"Auto set save_state to True because fsdp_stage == 3")
args.save_state = True
else:
zero_stage = 0
fsdp_stage = 0
@@ -868,11 +879,19 @@ def main():
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Clip Image Encoder
if args.train_mode != "normal":
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Get Transformer
transformer3d = WanTransformer3DModel.from_pretrained(
@@ -880,13 +899,6 @@ def main():
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
if args.train_mode != "normal":
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Freeze vae and text_encoder and set transformer3d to trainable
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
@@ -902,7 +914,7 @@ def main():
text_encoder,
transformer3d,
neuron_dropout=None,
add_lora_in_attn_temporal=True,
skip_name=args.lora_skip_name,
)
network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True)
@@ -1055,12 +1067,20 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
# Get the dataset
train_dataset = ImageVideoDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
video_repeat=args.video_repeat,
image_sample_size=args.image_sample_size,
enable_bucket=args.enable_bucket, enable_inpaint=True if args.train_mode != "normal" else False,
enable_bucket=args.enable_bucket,
enable_inpaint=True if args.train_mode != "normal" else False,
)
if args.enable_bucket:
@@ -1086,6 +1106,7 @@ def main():
}
return length_to_frame_num
def collate_fn(examples):
# Get token length
target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size
@@ -1138,9 +1159,9 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
@@ -1150,9 +1171,24 @@ def main():
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
for example in examples:
if args.random_ratio_crop:
if args.fix_sample_size is not None:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
# Get adapt hw for resize
fix_sample_size = list(map(lambda x: int(x), fix_sample_size))
transform = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
elif args.random_ratio_crop:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
@@ -1282,6 +1318,18 @@ def main():
network, optimizer, train_dataloader, lr_scheduler
)
if zero_stage == 3:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
transformer3d = shard_fn(transformer3d)
if fsdp_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
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)
transformer3d.to(accelerator.device, dtype=weight_dtype)
@@ -1302,6 +1350,7 @@ def main():
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1354,7 +1403,7 @@ def main():
first_epoch = global_step // num_update_steps_per_epoch
print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
if zero_stage != 3:
if zero_stage != 3 and not args.use_fsdp:
from safetensors.torch import load_file
state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device))
m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False)
@@ -1463,7 +1512,7 @@ def main():
pixel_values = batch["pixel_values"].to(weight_dtype)
# Increase the batch size when the length of the latent sequence of the current sample is small
if args.training_with_video_token_length:
if args.training_with_video_token_length and zero_stage != 3:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1))
if args.enable_text_encoder_in_dataloader:
@@ -1484,7 +1533,7 @@ def main():
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
mask = batch["mask"].to(weight_dtype)
# Increase the batch size when the length of the latent sequence of the current sample is small
if args.training_with_video_token_length:
if args.training_with_video_token_length and zero_stage != 3:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
@@ -1563,6 +1612,8 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
if args.train_mode != "normal":
clip_image_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1619,6 +1670,8 @@ def main():
if args.low_vram:
vae.to('cpu')
if args.train_mode != "normal":
clip_image_encoder.to('cpu')
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1813,11 +1866,12 @@ def main():
# Create the pipeline using the trained modules and save it.
accelerator.wait_for_everyone()
if accelerator.is_main_process:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
if args.save_state:
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
logger.info(f"Saved state to {accelerator_save_path}")
+64 -19
View File
@@ -576,7 +576,12 @@ def parse_args():
"--image_sample_size",
type=int,
default=512,
help="Sample size of the video.",
help="Sample size of the image.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn."
)
parser.add_argument(
"--video_sample_stride",
@@ -738,6 +743,9 @@ def main():
print(f"Using DeepSpeed Zero stage: {zero_stage}")
args.use_deepspeed = True
if zero_stage == 3:
print(f"Auto set save_state to True because zero_stage == 3")
args.save_state = True
elif fsdp_plugin is not None:
from torch.distributed.fsdp import ShardingStrategy
zero_stage = 0
@@ -752,6 +760,9 @@ def main():
print(f"Using FSDP stage: {fsdp_stage}")
args.use_fsdp = True
if fsdp_stage == 3:
print(f"Auto set save_state to True because fsdp_stage == 3")
args.save_state = True
else:
zero_stage = 0
fsdp_stage = 0
@@ -845,6 +856,13 @@ def main():
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Clip Image Encoder
if args.train_mode != "normal":
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Get Transformer
transformer3d = WanTransformer3DModel.from_pretrained(
@@ -852,13 +870,6 @@ def main():
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
if args.train_mode != "normal":
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Freeze vae and text_encoder and set transformer3d to trainable
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
@@ -1090,6 +1101,13 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
# Get the dataset
train_dataset = ImageVideoDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
@@ -1210,16 +1228,36 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
else:
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
for example in examples:
if args.random_ratio_crop:
if args.fix_sample_size is not None:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
# Get adapt hw for resize
fix_sample_size = list(map(lambda x: int(x), fix_sample_size))
transform = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
elif args.random_ratio_crop:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
@@ -1344,6 +1382,12 @@ def main():
transformer3d, optimizer, train_dataloader, lr_scheduler
)
if fsdp_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
text_encoder = shard_fn(text_encoder)
if args.use_ema:
ema_transformer3d.to(accelerator.device)
@@ -1368,6 +1412,7 @@ def main():
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1865,10 +1910,10 @@ def main():
if args.use_ema:
ema_transformer3d.copy_to(transformer3d.parameters())
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
accelerator.end_training()
+65 -20
View File
@@ -492,7 +492,12 @@ def parse_args():
"--image_sample_size",
type=int,
default=512,
help="Sample size of the video.",
help="Sample size of the image.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn."
)
parser.add_argument(
"--video_sample_stride",
@@ -670,6 +675,9 @@ def main():
print(f"Using DeepSpeed Zero stage: {zero_stage}")
args.use_deepspeed = True
if zero_stage == 3:
print(f"Auto set save_state to True because zero_stage == 3")
args.save_state = True
elif fsdp_plugin is not None:
from torch.distributed.fsdp import ShardingStrategy
zero_stage = 0
@@ -684,6 +692,9 @@ def main():
print(f"Using FSDP stage: {fsdp_stage}")
args.use_fsdp = True
if fsdp_stage == 3:
print(f"Auto set save_state to True because fsdp_stage == 3")
args.save_state = True
else:
zero_stage = 0
fsdp_stage = 0
@@ -777,6 +788,13 @@ def main():
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Clip Image Encoder
if args.train_mode != "normal":
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Get Transformer
transformer3d = WanTransformer3DModel.from_pretrained(
@@ -784,13 +802,6 @@ def main():
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
if args.train_mode != "normal":
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Freeze vae and text_encoder and set transformer3d to trainable
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
@@ -1020,13 +1031,19 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
# Get the dataset
train_dataset = ImageVideoControlDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
video_repeat=args.video_repeat,
image_sample_size=args.image_sample_size,
enable_bucket=args.enable_bucket, enable_inpaint=False,
enable_bucket=args.enable_bucket,
enable_camera_info=args.train_mode == "control_camera_ref"
)
@@ -1170,13 +1187,21 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
else:
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
for example in examples:
# To 0~1
@@ -1186,7 +1211,20 @@ def main():
control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous()
control_pixel_values = control_pixel_values / 255.
if args.random_ratio_crop:
if args.fix_sample_size is not None:
# Get adapt hw for resize
fix_sample_size = list(map(lambda x: int(x), fix_sample_size))
transform = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
transform_no_normalize = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
])
elif args.random_ratio_crop:
# Get adapt hw for resize
b, c, h, w = pixel_values.size()
th, tw = random_sample_size
@@ -1347,6 +1385,12 @@ def main():
transformer3d, optimizer, train_dataloader, lr_scheduler
)
if fsdp_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
text_encoder = shard_fn(text_encoder)
if args.use_ema:
ema_transformer3d.to(accelerator.device)
@@ -1371,6 +1415,7 @@ def main():
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1902,10 +1947,10 @@ def main():
if args.use_ema:
ema_transformer3d.copy_to(transformer3d.parameters())
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
accelerator.end_training()
+68 -22
View File
@@ -510,7 +510,12 @@ def parse_args():
"--image_sample_size",
type=int,
default=512,
help="Sample size of the video.",
help="Sample size of the image.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn."
)
parser.add_argument(
"--video_sample_stride",
@@ -668,6 +673,9 @@ def main():
print(f"Using DeepSpeed Zero stage: {zero_stage}")
args.use_deepspeed = True
if zero_stage == 3:
print(f"Auto set save_state to True because zero_stage == 3")
args.save_state = True
elif fsdp_plugin is not None:
from torch.distributed.fsdp import ShardingStrategy
zero_stage = 0
@@ -682,6 +690,9 @@ def main():
print(f"Using FSDP stage: {fsdp_stage}")
args.use_fsdp = True
if fsdp_stage == 3:
print(f"Auto set save_state to True because fsdp_stage == 3")
args.save_state = True
else:
zero_stage = 0
fsdp_stage = 0
@@ -775,6 +786,13 @@ def main():
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Clip Image Encoder
if args.train_mode != "normal":
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Get Transformer
transformer3d = WanTransformer3DModel.from_pretrained(
@@ -782,13 +800,6 @@ def main():
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
if args.train_mode != "normal":
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Freeze vae and text_encoder and set transformer3d to trainable
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
@@ -955,13 +966,19 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
# Get the dataset
train_dataset = ImageVideoControlDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
video_repeat=args.video_repeat,
image_sample_size=args.image_sample_size,
enable_bucket=args.enable_bucket, enable_inpaint=False,
enable_bucket=args.enable_bucket,
enable_camera_info=args.train_mode == "control_camera_ref"
)
@@ -1105,13 +1122,21 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
else:
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
for example in examples:
# To 0~1
@@ -1121,7 +1146,20 @@ def main():
control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous()
control_pixel_values = control_pixel_values / 255.
if args.random_ratio_crop:
if args.fix_sample_size is not None:
# Get adapt hw for resize
fix_sample_size = list(map(lambda x: int(x), fix_sample_size))
transform = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
transform_no_normalize = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
])
elif args.random_ratio_crop:
# Get adapt hw for resize
b, c, h, w = pixel_values.size()
th, tw = random_sample_size
@@ -1295,6 +1333,12 @@ def main():
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
transformer3d = shard_fn(transformer3d)
if fsdp_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
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)
transformer3d.to(accelerator.device, dtype=weight_dtype)
@@ -1315,6 +1359,7 @@ def main():
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1367,7 +1412,7 @@ def main():
first_epoch = global_step // num_update_steps_per_epoch
print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
if zero_stage != 3:
if zero_stage != 3 and not args.use_fsdp:
from safetensors.torch import load_file
state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device))
m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False)
@@ -1870,11 +1915,12 @@ def main():
# Create the pipeline using the trained modules and save it.
accelerator.wait_for_everyone()
if accelerator.is_main_process:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
if args.save_state:
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
logger.info(f"Saved state to {accelerator_save_path}")
+68 -22
View File
@@ -588,7 +588,12 @@ def parse_args():
"--image_sample_size",
type=int,
default=512,
help="Sample size of the video.",
help="Sample size of the image.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn."
)
parser.add_argument(
"--video_sample_stride",
@@ -730,6 +735,9 @@ def main():
print(f"Using DeepSpeed Zero stage: {zero_stage}")
args.use_deepspeed = True
if zero_stage == 3:
print(f"Auto set save_state to True because zero_stage == 3")
args.save_state = True
elif fsdp_plugin is not None:
from torch.distributed.fsdp import ShardingStrategy
zero_stage = 0
@@ -744,6 +752,9 @@ def main():
print(f"Using FSDP stage: {fsdp_stage}")
args.use_fsdp = True
if fsdp_stage == 3:
print(f"Auto set save_state to True because fsdp_stage == 3")
args.save_state = True
else:
zero_stage = 0
fsdp_stage = 0
@@ -837,20 +848,20 @@ def main():
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Clip Image Encoder
if args.train_mode != "normal":
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Get Transformer
transformer3d = WanTransformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
if args.train_mode != "normal":
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
)
clip_image_encoder = clip_image_encoder.eval()
# Freeze vae and text_encoder and set transformer3d to trainable
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
@@ -1019,6 +1030,13 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
# Get the dataset
train_dataset = ImageVideoDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames,
@@ -1139,16 +1157,36 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
if args.random_ratio_crop:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
else:
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
for example in examples:
if args.random_ratio_crop:
if args.fix_sample_size is not None:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
# Get adapt hw for resize
fix_sample_size = list(map(lambda x: int(x), fix_sample_size))
transform = transforms.Compose([
transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC
transforms.CenterCrop(fix_sample_size),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
elif args.random_ratio_crop:
# To 0~1
pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
@@ -1286,6 +1324,12 @@ def main():
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
transformer3d = shard_fn(transformer3d)
if fsdp_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
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)
transformer3d.to(accelerator.device, dtype=weight_dtype)
@@ -1306,6 +1350,7 @@ def main():
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1358,7 +1403,7 @@ def main():
first_epoch = global_step // num_update_steps_per_epoch
print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
if zero_stage != 3:
if zero_stage != 3 and not args.use_fsdp:
from safetensors.torch import load_file
state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device))
m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False)
@@ -1827,11 +1872,12 @@ def main():
# Create the pipeline using the trained modules and save it.
accelerator.wait_for_everyone()
if accelerator.is_main_process:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
if args.save_state:
if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process:
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
logger.info(f"Saved state to {accelerator_save_path}")