File size: 4,436 Bytes
59630ba | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 | from pathlib import Path
from omegaconf import DictConfig
import torch
from einops import rearrange
from lightning.pytorch.utilities.types import STEP_OUTPUT
from algorithms.common.base_pytorch_algo import BasePytorchAlgo
from utils.storage_utils import safe_torch_save
from utils.logging_utils import log_video
from ..common.distribution import DiagonalGaussianDistribution
from .trainer import ImageVAE
class ImageVAEPreprocessor(BasePytorchAlgo):
"""
An algorithm for preprocessing videos to latents using a pretrained ImageVAE model.
"""
def __init__(self, cfg: DictConfig):
self.pretrained_path = cfg.pretrained_path
self.pretrained_kwargs = cfg.pretrained_kwargs
self.use_fp16 = cfg.precision == "16-true"
self.max_encode_length = cfg.max_encode_length
self.max_decode_length = cfg.logging.max_video_length
self.log_every_n_batch = cfg.logging.every_n_batch
super().__init__(cfg)
def _build_model(self):
self.vae = ImageVAE.from_pretrained(
path=self.pretrained_path,
torch_dtype=torch.float16 if self.use_fp16 else torch.float32,
**self.pretrained_kwargs,
)
def training_step(self, batch, batch_idx) -> STEP_OUTPUT:
raise NotImplementedError(
"Training not implemented for VAEVideo. Only used for validation"
)
def test_step(self, batch, batch_idx) -> STEP_OUTPUT:
raise NotImplementedError(
"Testing not implemented for VAEVideo. Only used for validation"
)
def validation_step(self, batch, batch_idx, dataloader_idx=0) -> STEP_OUTPUT:
videos, latent_paths = batch
latent_paths = [Path(path) for path in latent_paths]
batch_size = videos.shape[0]
videos = self._rearrange_and_normalize(videos)
# Encode the video data into a latent space
# always convert to float16 (as they will be saved as float16 tensors)
latent_dist = self._encode_videos(videos)
latents = latent_dist.sample().to(torch.float16)
# just to see the progress in wandb
if batch_idx % 100 == 0:
self.log("dummy", 0.0)
# log gt vs reconstructed video to wandb
if batch_idx % self.log_every_n_batch == 0 and self.logger:
videos = videos.detach().cpu()[: self.max_decode_length]
reconstructed_videos = self.vae.decode(latents[: self.max_decode_length])
reconstructed_videos = reconstructed_videos.detach().cpu()
videos = self._rearrange_and_unnormalize(videos, batch_size)
reconstructed_videos = self._rearrange_and_unnormalize(
reconstructed_videos, batch_size
)
log_video(
reconstructed_videos,
videos,
step=self.global_step,
namespace="reconstruction_vis",
logger=self.logger.experiment,
captions=[
f"{p.parent.parent.name}/{p.parent.name}/{p.stem}"
for p in latent_paths
],
)
# save the latent to disk
latents_to_save = (
rearrange(
latents,
"(b f) c h w -> b f c h w",
b=batch_size,
)
.detach()
.cpu()
)
for latent, latent_path in zip(latents_to_save, latent_paths):
# should clone latent to avoid having large file size
safe_torch_save(latent.clone(), latent_path)
return None
def _encode_videos(self, video: torch.Tensor) -> DiagonalGaussianDistribution:
chunks = video.chunk(
(len(video) + self.max_encode_length - 1) // self.max_encode_length, dim=0
)
latent_dist_list = []
for chunk in chunks:
latent_dist_list.append(self.vae.encode(chunk))
return DiagonalGaussianDistribution.cat(latent_dist_list)
def _rearrange_and_normalize(self, videos: torch.Tensor) -> torch.Tensor:
videos = rearrange(videos, "b f c h w -> (b f) c h w")
videos = 2.0 * videos - 1.0
return videos
def _rearrange_and_unnormalize(
self, videos: torch.Tensor, batch_size: int
) -> torch.Tensor:
videos = 0.5 * videos + 0.5
videos = rearrange(videos, "(b f) c h w -> b f c h w", b=batch_size)
return videos
|