more dataset fixes from stashed changes (#49)

This commit is contained in:
Aryan
2024-10-19 06:30:26 +05:30
committed by GitHub
parent f2a1626de2
commit 8da9638896
2 changed files with 35 additions and 38 deletions
+20 -14
View File
@@ -229,32 +229,38 @@ class VideoDataset(Dataset):
pt_filename = f"{filename_without_ext}.pt"
# The current path is something like: /a/b/c/d/videos/00001.mp4
# We need to reach: /a/b/c/d/latents/00001.pt
images_path = path.parent.parent.joinpath("image_latents")
latents_path = path.parent.parent.joinpath("latents")
embeds_path = path.parent.parent.joinpath("embeddings")
# We need to reach: /a/b/c/d/video_latents/00001.pt
image_latents_path = path.parent.parent.joinpath("image_latents")
video_latents_path = path.parent.parent.joinpath("video_latents")
embeds_path = path.parent.parent.joinpath("prompt_embeds")
if not latents_path.exists() or not embeds_path.exists() or (self.image_to_video and not images_path.exists()):
if (
not video_latents_path.exists()
or not embeds_path.exists()
or (self.image_to_video and not image_latents_path.exists())
):
raise ValueError(
f"When setting the load_tensors parameter to `True`, it is expected that the `{self.data_root=}` contains two folders named `latents` and `embeddings`. However, these folders were not found. Please make sure to have prepared your data correctly using `prepare_data.py`. Additionally, if you're training image-to-video, it is expected that an `image_latents` folder is also present."
f"When setting the load_tensors parameter to `True`, it is expected that the `{self.data_root=}` contains two folders named `video_latents` and `prompt_embeds`. However, these folders were not found. Please make sure to have prepared your data correctly using `prepare_data.py`. Additionally, if you're training image-to-video, it is expected that an `image_latents` folder is also present."
)
if self.image_to_video:
image_filepath = images_path.joinpath(pt_filename)
latent_filepath = latents_path.joinpath(pt_filename)
image_latent_filepath = image_latents_path.joinpath(pt_filename)
video_latent_filepath = video_latents_path.joinpath(pt_filename)
embeds_filepath = embeds_path.joinpath(pt_filename)
if not latent_filepath.is_file() or not embeds_filepath.is_file():
if not video_latent_filepath.is_file() or not embeds_filepath.is_file():
if self.image_to_video:
image_filepath = image_filepath.as_posix()
latent_filepath = latent_filepath.as_posix()
image_latent_filepath = image_latent_filepath.as_posix()
video_latent_filepath = video_latent_filepath.as_posix()
embeds_filepath = embeds_filepath.as_posix()
raise ValueError(
f"The file {latent_filepath=} or {embeds_filepath=} could not be found. Please ensure that you've correctly executed `prepare_dataset.py`."
f"The file {video_latent_filepath=} or {embeds_filepath=} could not be found. Please ensure that you've correctly executed `prepare_dataset.py`."
)
images = torch.load(image_filepath, map_location="cpu", weights_only=True) if self.image_to_video else None
latents = torch.load(latent_filepath, map_location="cpu", weights_only=True)
images = (
torch.load(image_latent_filepath, map_location="cpu", weights_only=True) if self.image_to_video else None
)
latents = torch.load(video_latent_filepath, map_location="cpu", weights_only=True)
embeds = torch.load(embeds_filepath, map_location="cpu", weights_only=True)
return images, latents, embeds
+15 -24
View File
@@ -477,8 +477,6 @@ def main():
# 3. Prepare models
device = f"cuda:{rank}"
generator = torch.Generator(device).manual_seed(args.seed)
if args.save_latents_and_embeddings:
tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer")
text_encoder = T5EncoderModel.from_pretrained(
@@ -520,29 +518,22 @@ def main():
# Encode videos & images
if args.save_latents_and_embeddings:
if args.save_image_latents:
image_noise_sigma = torch.normal(
mean=-3.0,
std=0.5,
size=(images.size(0),),
generator=generator,
device=device,
dtype=weight_dtype,
)
image_noise_sigma = torch.exp(image_noise_sigma)
noisy_images = (
images
+ torch.empty_like(images).normal_(generator=generator)
* image_noise_sigma[:, None, None, None, None]
)
image_latent_dist = vae.encode(noisy_images).latent_dist
image_latents = image_latent_dist.sample() * vae.config.scaling_factor
image_latents = image_latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
if args.use_slicing:
if args.save_image_latents:
encoded_slices = [vae._encode(image_slice) for image_slice in images.split(1)]
image_latents = torch.cat(encoded_slices)
image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
encoded_slices = [vae._encode(video_slice) for video_slice in videos.split(1)]
video_latents = torch.cat(encoded_slices)
else:
if args.save_image_latents:
image_latents = vae._encode(images)
image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
video_latents = vae._encode(videos)
latent_dist = vae.encode(videos).latent_dist
video_latents = latent_dist.sample(generator=generator) * vae.config.scaling_factor
video_latents = video_latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
video_latents = video_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
# Encode prompts