Fix: improve batch_video_length calculation logic to ensure generated mask behaves as expected (#333)

This commit is contained in:
王泽鹏
2025-09-23 15:19:12 +08:00
committed by GitHub
parent 3e8d6867d3
commit add6b679a4
12 changed files with 189 additions and 140 deletions
+14 -7
View File
@@ -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:
+14 -7
View File
@@ -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:
+16 -12
View File
@@ -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
+15 -12
View File
@@ -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
+16 -12
View File
@@ -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
+16 -12
View File
@@ -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
+16 -12
View File
@@ -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
+16 -12
View File
@@ -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
+16 -12
View File
@@ -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
+17 -15
View File
@@ -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:
+17 -15
View File
@@ -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:
+16 -12
View File
@@ -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