Fix: improve batch_video_length calculation logic to ensure generated mask behaves as expected (#333)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user