From 924dd8528a4493f5fc8b9a419a47f769f6cf05ec Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Wed, 30 Jul 2025 19:39:31 +0800 Subject: [PATCH] Update auto_tile_batch_size args --- scripts/cogvideox_fun/train.py | 5 ++++- scripts/cogvideox_fun/train_control.py | 5 ++++- scripts/cogvideox_fun/train_lora.py | 5 ++++- scripts/wan2.1/train.py | 7 +++++-- scripts/wan2.1/train_lora.py | 7 +++++-- scripts/wan2.1_fun/train.py | 7 +++++-- scripts/wan2.1_fun/train_control.py | 7 +++++-- scripts/wan2.1_fun/train_control_lora.py | 7 +++++-- scripts/wan2.1_fun/train_lora.py | 7 +++++-- scripts/wan2.2/train.py | 7 +++++-- scripts/wan2.2/train_lora.py | 7 +++++-- 11 files changed, 52 insertions(+), 19 deletions(-) diff --git a/scripts/cogvideox_fun/train.py b/scripts/cogvideox_fun/train.py index 2cc8d4a..1443e5b 100755 --- a/scripts/cogvideox_fun/train.py +++ b/scripts/cogvideox_fun/train.py @@ -570,6 +570,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1343,7 +1346,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.auto_tile_batch_size and args.training_with_video_token_length: 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: diff --git a/scripts/cogvideox_fun/train_control.py b/scripts/cogvideox_fun/train_control.py index 7e1cfeb..3eafdd5 100755 --- a/scripts/cogvideox_fun/train_control.py +++ b/scripts/cogvideox_fun/train_control.py @@ -520,6 +520,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1273,7 +1276,7 @@ def main(): pixel_values = batch["pixel_values"].to(weight_dtype) control_pixel_values = batch["control_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.auto_tile_batch_size and args.training_with_video_token_length: 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)) control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1)) diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index 0502ae4..f33b6f6 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -582,6 +582,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames." ) @@ -1338,7 +1341,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.auto_tile_batch_size and args.training_with_video_token_length: 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: diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index e148ae5..958e094 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -579,6 +579,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1516,7 +1519,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 zero_stage != 3: + if args.auto_tile_batch_size and 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: @@ -1537,7 +1540,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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index fff622b..c0a5377 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -586,6 +586,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames." ) @@ -1512,7 +1515,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 zero_stage != 3: + if args.auto_tile_batch_size and 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: @@ -1533,7 +1536,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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index bf19ea3..4f83aaf 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -543,6 +543,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1513,7 +1516,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 zero_stage != 3: + if args.auto_tile_batch_size and 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: @@ -1534,7 +1537,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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index d45fd27..0c2334d 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -459,6 +459,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1524,7 +1527,7 @@ def main(): control_camera_values = batch["control_camera_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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1)) @@ -1551,7 +1554,7 @@ def main(): clip_pixel_values = batch["clip_pixel_values"] clip_idx = batch["clip_idx"] # 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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1)) diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index 700dc6c..43dd088 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -477,6 +477,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1528,7 +1531,7 @@ def main(): control_camera_values = batch["control_camera_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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1)) @@ -1555,7 +1558,7 @@ def main(): clip_pixel_values = batch["clip_pixel_values"] clip_idx = batch["clip_idx"] # 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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) ref_pixel_values = torch.tile(ref_pixel_values, (4, 1, 1, 1, 1)) diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index e14150c..0136f08 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -555,6 +555,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1512,7 +1515,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 zero_stage != 3: + if args.auto_tile_batch_size and 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: @@ -1533,7 +1536,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 zero_stage != 3: + if args.auto_tile_batch_size and 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)) diff --git a/scripts/wan2.2/train.py b/scripts/wan2.2/train.py index 3a89d72..011b4d3 100644 --- a/scripts/wan2.2/train.py +++ b/scripts/wan2.2/train.py @@ -572,6 +572,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." ) @@ -1515,7 +1518,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 zero_stage != 3: + if args.auto_tile_batch_size and 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: @@ -1535,7 +1538,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 zero_stage != 3: + if args.auto_tile_batch_size and 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]: mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) mask = torch.tile(mask, (4, 1, 1, 1, 1)) diff --git a/scripts/wan2.2/train_lora.py b/scripts/wan2.2/train_lora.py index b4c1eec..7235d2f 100755 --- a/scripts/wan2.2/train_lora.py +++ b/scripts/wan2.2/train_lora.py @@ -585,6 +585,9 @@ def parse_args(): parser.add_argument( "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", ) + parser.add_argument( + "--auto_tile_batch_size", action="store_true", help="Whether to auto tile batch size.", + ) parser.add_argument( "--noise_share_in_frames", action="store_true", help="Whether enable noise share in frames." ) @@ -1522,7 +1525,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 zero_stage != 3: + if args.auto_tile_batch_size and 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: @@ -1542,7 +1545,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 zero_stage != 3: + if args.auto_tile_batch_size and 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]: mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) mask = torch.tile(mask, (4, 1, 1, 1, 1))