diff --git a/scripts/cogvideox_fun/train.py b/scripts/cogvideox_fun/train.py index 1443e5b..6220534 100755 --- a/scripts/cogvideox_fun/train.py +++ b/scripts/cogvideox_fun/train.py @@ -1099,6 +1099,14 @@ def main(): ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + for example in examples: if args.random_ratio_crop: # To 0~1 @@ -1137,10 +1145,6 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) - new_examples["text"].append(example["text"]) - - batch_video_length = int(min(batch_video_length, len(pixel_values))) # Magvae needs the number of frames to be 4n + 1. local_latent_length = (batch_video_length - 1) // sample_n_frames_bucket_interval + 1 @@ -1156,6 +1160,9 @@ def main(): if batch_video_length <= 0: batch_video_length = 1 + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) + new_examples["text"].append(example["text"]) + if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size()) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) + torch.ones_like(new_examples["pixel_values"][-1]) * -1 * mask @@ -1163,10 +1170,10 @@ def main(): new_examples["mask"].append(mask) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) # Encode prompts when enable_text_encoder_in_dataloader=True if args.enable_text_encoder_in_dataloader: diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index f33b6f6..2369638 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -1037,6 +1037,14 @@ def main(): ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + for example in examples: if args.random_ratio_crop: # To 0~1 @@ -1075,10 +1083,6 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) - new_examples["text"].append(example["text"]) - - batch_video_length = int(min(batch_video_length, len(pixel_values))) # Magvae needs the number of frames to be 4n + 1. local_latent_length = (batch_video_length - 1) // sample_n_frames_bucket_interval + 1 @@ -1093,6 +1097,9 @@ def main(): if batch_video_length <= 0: batch_video_length = 1 + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) + new_examples["text"].append(example["text"]) if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size()) @@ -1101,10 +1108,10 @@ def main(): new_examples["mask"].append(mask) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) # Encode prompts when enable_text_encoder_in_dataloader=True if args.enable_text_encoder_in_dataloader: diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 90be7ad..813eaa6 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -1253,6 +1253,17 @@ def main(): 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] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1303,17 +1314,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size(), image_start_only=True) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1327,10 +1331,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 1f9aa3e..1696e26 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -1178,6 +1178,17 @@ def main(): 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] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1228,17 +1239,9 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size(), image_start_only=True) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1252,10 +1255,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index b06678f..04ec5ae 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -1248,6 +1248,17 @@ def main(): 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] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1298,17 +1309,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size()) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1322,10 +1326,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index 1d57141..51e47cd 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -1176,6 +1176,17 @@ def main(): 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] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1226,17 +1237,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size()) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1250,10 +1254,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.2/train.py b/scripts/wan2.2/train.py index 11bd975..e2d3296 100644 --- a/scripts/wan2.2/train.py +++ b/scripts/wan2.2/train.py @@ -1250,6 +1250,17 @@ def main(): closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1300,17 +1311,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size(), image_start_only=True) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1324,10 +1328,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.2/train_lora.py b/scripts/wan2.2/train_lora.py index 2f8a3e1..83f19dd 100755 --- a/scripts/wan2.2/train_lora.py +++ b/scripts/wan2.2/train_lora.py @@ -1186,6 +1186,17 @@ def main(): closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1236,17 +1247,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size(), image_start_only=True) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1260,10 +1264,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.2_fun/train.py b/scripts/wan2.2_fun/train.py index 37c6a16..39a97d0 100644 --- a/scripts/wan2.2_fun/train.py +++ b/scripts/wan2.2_fun/train.py @@ -1256,6 +1256,17 @@ def main(): closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1306,17 +1317,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size()) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1330,10 +1334,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True diff --git a/scripts/wan2.2_fun/train_control.py b/scripts/wan2.2_fun/train_control.py index 3858f8f..a650d22 100644 --- a/scripts/wan2.2_fun/train_control.py +++ b/scripts/wan2.2_fun/train_control.py @@ -1264,6 +1264,17 @@ def main(): closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: # To 0~1 pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() @@ -1324,8 +1335,8 @@ def main(): transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC transforms.CenterCrop(closest_size), ]) - - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["control_pixel_values"].append(transform(control_pixel_values)) if args.train_mode == "control_camera_ref": @@ -1344,15 +1355,6 @@ def main(): new_examples["control_camera_values"].append(transform_no_normalize(local_control_camera_values)) new_examples["text"].append(example["text"]) - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = int( - min( - batch_video_length, - (len(pixel_values) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1, - ) - ) - if batch_video_length == 0: - batch_video_length = 1 if args.train_mode != "control": if args.control_ref_image == "first_frame": @@ -1387,17 +1389,17 @@ def main(): new_examples["mask"].append(mask) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) new_examples["control_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_pixel_values"]]) if args.train_mode != "control": - new_examples["ref_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["ref_pixel_values"]]) + new_examples["ref_pixel_values"] = torch.stack([example for example in new_examples["ref_pixel_values"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) new_examples["clip_idx"] = torch.tensor(new_examples["clip_idx"]) if args.train_mode == "control_camera_ref": new_examples["control_camera_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_camera_values"]]) if args.add_inpaint_info: - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) # Encode prompts when enable_text_encoder_in_dataloader=True if args.enable_text_encoder_in_dataloader: diff --git a/scripts/wan2.2_fun/train_control_lora.py b/scripts/wan2.2_fun/train_control_lora.py index 4311af2..d9c5f8c 100644 --- a/scripts/wan2.2_fun/train_control_lora.py +++ b/scripts/wan2.2_fun/train_control_lora.py @@ -1198,6 +1198,17 @@ def main(): closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: # To 0~1 pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() @@ -1258,8 +1269,8 @@ def main(): transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC transforms.CenterCrop(closest_size), ]) - - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["control_pixel_values"].append(transform(control_pixel_values)) if args.train_mode == "control_camera_ref": @@ -1278,15 +1289,6 @@ def main(): new_examples["control_camera_values"].append(transform_no_normalize(local_control_camera_values)) new_examples["text"].append(example["text"]) - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = int( - min( - batch_video_length, - (len(pixel_values) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1, - ) - ) - if batch_video_length == 0: - batch_video_length = 1 if args.train_mode != "control": if args.control_ref_image == "first_frame": @@ -1321,17 +1323,17 @@ def main(): new_examples["mask"].append(mask) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) new_examples["control_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_pixel_values"]]) if args.train_mode != "control": - new_examples["ref_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["ref_pixel_values"]]) + new_examples["ref_pixel_values"] = torch.stack([example for example in new_examples["ref_pixel_values"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) new_examples["clip_idx"] = torch.tensor(new_examples["clip_idx"]) if args.train_mode == "control_camera_ref": new_examples["control_camera_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_camera_values"]]) if args.add_inpaint_info: - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) # Encode prompts when enable_text_encoder_in_dataloader=True if args.enable_text_encoder_in_dataloader: diff --git a/scripts/wan2.2_fun/train_lora.py b/scripts/wan2.2_fun/train_lora.py index 9396c3b..e0bb0a7 100644 --- a/scripts/wan2.2_fun/train_lora.py +++ b/scripts/wan2.2_fun/train_lora.py @@ -1185,6 +1185,17 @@ def main(): closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size] + min_example_length = min( + [example["pixel_values"].shape[0] for example in examples] + ) + batch_video_length = int(min(batch_video_length, min_example_length)) + + # Magvae needs the number of frames to be 4n + 1. + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + + if batch_video_length <= 0: + batch_video_length = 1 + for example in examples: if args.fix_sample_size is not None: # To 0~1 @@ -1235,17 +1246,10 @@ def main(): transforms.CenterCrop(closest_size), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), ]) - new_examples["pixel_values"].append(transform(pixel_values)) + + new_examples["pixel_values"].append(transform(pixel_values)[:batch_video_length]) new_examples["text"].append(example["text"]) - batch_video_length = int(min(batch_video_length, len(pixel_values))) - - # Magvae needs the number of frames to be 4n + 1. - batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 - - if batch_video_length <= 0: - batch_video_length = 1 - if args.train_mode != "normal": mask = get_random_mask(new_examples["pixel_values"][-1].size()) mask_pixel_values = new_examples["pixel_values"][-1] * (1 - mask) @@ -1259,10 +1263,10 @@ def main(): new_examples["clip_pixel_values"].append(clip_pixel_values) # Limit the number of frames to the same - new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) if args.train_mode != "normal": - new_examples["mask_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["mask_pixel_values"]]) - new_examples["mask"] = torch.stack([example[:batch_video_length] for example in new_examples["mask"]]) + new_examples["mask_pixel_values"] = torch.stack([example for example in new_examples["mask_pixel_values"]]) + new_examples["mask"] = torch.stack([example for example in new_examples["mask"]]) new_examples["clip_pixel_values"] = torch.stack([example for example in new_examples["clip_pixel_values"]]) # Encode prompts when enable_text_encoder_in_dataloader=True