Update auto_tile_batch_size args

This commit is contained in:
bubbliiiing
2025-07-30 19:39:31 +08:00
parent e258d4158b
commit 924dd8528a
11 changed files with 52 additions and 19 deletions
+4 -1
View File
@@ -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:
+4 -1
View File
@@ -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))
+4 -1
View File
@@ -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:
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))
+5 -2
View File
@@ -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))