Download GeometryForcing/algorithms/vae/video_vae/model.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 19.7 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/algorithms/vae/video_vae/model.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/GeometryForcing/algorithms/vae/video_vae/model.py
-
curl -L -o model.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/algorithms/vae/video_vae/model.py
19.7 kB
| """ | |
| Adapted from Open-Sora-Plan-v1.2.0 | |
| https://github.com/PKU-YuanGroup/Open-Sora-Plan/blob/main/opensora/models/causalvideovae/model/causal_vae/modeling_causalvae.py | |
| """ | |
| from typing import Tuple, Optional, Literal | |
| from functools import partial | |
| import torch | |
| import torch.nn as nn | |
| from einops import rearrange, repeat | |
| from utils.ckpt_utils import ( | |
| is_wandb_run_path, | |
| is_hf_path, | |
| wandb_to_local_path, | |
| download_pretrained as hf_to_local_path, | |
| ) | |
| from ..common.modules import Normalize, nonlinearity | |
| from ..common.modules.utils import ( | |
| resolve_str_to_module as _resolve_str_to_module, | |
| Module, | |
| ) | |
| from ..common.distribution import DiagonalGaussianDistribution | |
| from ..common.base_vae import VAE | |
| class Encoder(nn.Module): | |
| def __init__( | |
| self, | |
| z_channels: int, | |
| hidden_size: int, | |
| hidden_size_mult: Tuple[int] = (1, 2, 4, 4), | |
| attn_resolutions: Tuple[int] = (16,), | |
| conv_in: Module = "Conv2d", | |
| conv_out: Module = "PaddedConv3D", | |
| attention: Module = "AttnBlock", | |
| resnet_blocks: Tuple[Module] = ( | |
| "ResnetBlock2D", | |
| "ResnetBlock2D", | |
| "ResnetBlock2D", | |
| "ResnetBlock3D", | |
| ), | |
| spatial_downsample: Tuple[Module] = ( | |
| "Downsample", | |
| "Downsample", | |
| "Downsample", | |
| "", | |
| ), | |
| temporal_downsample: Tuple[Module] = ("", "", "TimeDownsampleRes2x", ""), | |
| mid_resnet: Module = "ResnetBlock3D", | |
| dropout: float = 0.0, | |
| resolution: int = 256, | |
| num_res_blocks: int = 2, | |
| double_z: bool = True, | |
| is_causal: bool = True, | |
| ) -> None: | |
| super().__init__() | |
| assert len(resnet_blocks) == len(hidden_size_mult), print( | |
| hidden_size_mult, resnet_blocks | |
| ) | |
| resolve_str_to_module = partial(_resolve_str_to_module, is_causal=is_causal) | |
| # ---- Config ---- | |
| self.num_resolutions = len(hidden_size_mult) | |
| self.resolution = resolution | |
| self.num_res_blocks = num_res_blocks | |
| # ---- In ---- | |
| self.conv_in = resolve_str_to_module(conv_in)( | |
| 3, hidden_size, kernel_size=3, stride=1, padding=1 | |
| ) | |
| # ---- Downsample ---- | |
| curr_res = resolution | |
| in_ch_mult = (1,) + tuple(hidden_size_mult) | |
| self.in_ch_mult = in_ch_mult | |
| self.down = nn.ModuleList() | |
| for i_level in range(self.num_resolutions): | |
| block = nn.ModuleList() | |
| attn = nn.ModuleList() | |
| block_in = hidden_size * in_ch_mult[i_level] | |
| block_out = hidden_size * hidden_size_mult[i_level] | |
| for i_block in range(self.num_res_blocks): | |
| block.append( | |
| resolve_str_to_module(resnet_blocks[i_level])( | |
| in_channels=block_in, | |
| out_channels=block_out, | |
| dropout=dropout, | |
| ) | |
| ) | |
| block_in = block_out | |
| if curr_res in attn_resolutions: | |
| attn.append(resolve_str_to_module(attention)(block_in)) | |
| down = nn.Module() | |
| down.block = block | |
| down.attn = attn | |
| if spatial_downsample[i_level]: | |
| down.downsample = resolve_str_to_module(spatial_downsample[i_level])( | |
| block_in, block_in | |
| ) | |
| curr_res = curr_res // 2 | |
| if temporal_downsample[i_level]: | |
| down.time_downsample = resolve_str_to_module( | |
| temporal_downsample[i_level] | |
| )(block_in, block_in) | |
| self.down.append(down) | |
| # ---- Mid ---- | |
| self.mid = nn.Module() | |
| self.mid.block_1 = resolve_str_to_module(mid_resnet)( | |
| in_channels=block_in, | |
| out_channels=block_in, | |
| dropout=dropout, | |
| ) | |
| self.mid.attn_1 = resolve_str_to_module(attention)(block_in) | |
| self.mid.block_2 = resolve_str_to_module(mid_resnet)( | |
| in_channels=block_in, | |
| out_channels=block_in, | |
| dropout=dropout, | |
| ) | |
| # ---- Out ---- | |
| self.norm_out = Normalize(block_in) | |
| self.conv_out = resolve_str_to_module(conv_out)( | |
| block_in, | |
| 2 * z_channels if double_z else z_channels, | |
| kernel_size=3, | |
| stride=1, | |
| padding=1, | |
| ) | |
| def forward(self, x): | |
| hs = [self.conv_in(x)] | |
| for i_level in range(self.num_resolutions): | |
| for i_block in range(self.num_res_blocks): | |
| h = self.down[i_level].block[i_block](hs[-1]) | |
| if len(self.down[i_level].attn) > 0: | |
| h = self.down[i_level].attn[i_block](h) | |
| hs.append(h) | |
| if hasattr(self.down[i_level], "downsample"): | |
| hs.append(self.down[i_level].downsample(hs[-1])) | |
| if hasattr(self.down[i_level], "time_downsample"): | |
| hs_down = self.down[i_level].time_downsample(hs[-1]) | |
| hs.append(hs_down) | |
| h = self.mid.block_1(h) | |
| h = self.mid.attn_1(h) | |
| h = self.mid.block_2(h) | |
| h = self.norm_out(h) | |
| h = nonlinearity(h) | |
| h = self.conv_out(h) | |
| return h | |
| class Decoder(nn.Module): | |
| def __init__( | |
| self, | |
| z_channels: int, | |
| hidden_size: int, | |
| hidden_size_mult: Tuple[int] = (1, 2, 4, 4), | |
| attn_resolutions: Tuple[int] = (16,), | |
| conv_in: Module = "Conv2d", | |
| conv_out: Module = "PaddedConv3D", | |
| attention: Module = "AttnBlock", | |
| resnet_blocks: Tuple[Module] = ( | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| ), | |
| spatial_upsample: Tuple[Module] = ( | |
| "", | |
| "SpatialUpsample2x", | |
| "SpatialUpsample2x", | |
| "SpatialUpsample2x", | |
| ), | |
| temporal_upsample: Tuple[Module] = ("", "", "", "TimeUpsampleRes2x"), | |
| mid_resnet: Module = "ResnetBlock3D", | |
| dropout: float = 0.0, | |
| resolution: int = 256, | |
| num_res_blocks: int = 2, | |
| is_causal: bool = True, | |
| ): | |
| super().__init__() | |
| resolve_str_to_module = partial(_resolve_str_to_module, is_causal=is_causal) | |
| # ---- Config ---- | |
| self.num_resolutions = len(hidden_size_mult) | |
| self.resolution = resolution | |
| self.num_res_blocks = num_res_blocks | |
| # ---- In ---- | |
| block_in = hidden_size * hidden_size_mult[self.num_resolutions - 1] | |
| curr_res = resolution // 2 ** (self.num_resolutions - 1) | |
| self.conv_in = resolve_str_to_module(conv_in)( | |
| z_channels, block_in, kernel_size=3, padding=1 | |
| ) | |
| # ---- Mid ---- | |
| self.mid = nn.Module() | |
| self.mid.block_1 = resolve_str_to_module(mid_resnet)( | |
| in_channels=block_in, | |
| out_channels=block_in, | |
| dropout=dropout, | |
| ) | |
| self.mid.attn_1 = resolve_str_to_module(attention)(block_in) | |
| self.mid.block_2 = resolve_str_to_module(mid_resnet)( | |
| in_channels=block_in, | |
| out_channels=block_in, | |
| dropout=dropout, | |
| ) | |
| # ---- Upsample ---- | |
| self.up = nn.ModuleList() | |
| for i_level in reversed(range(self.num_resolutions)): | |
| block = nn.ModuleList() | |
| attn = nn.ModuleList() | |
| block_out = hidden_size * hidden_size_mult[i_level] | |
| for i_block in range(self.num_res_blocks + 1): | |
| block.append( | |
| resolve_str_to_module(resnet_blocks[i_level])( | |
| in_channels=block_in, | |
| out_channels=block_out, | |
| dropout=dropout, | |
| ) | |
| ) | |
| block_in = block_out | |
| if curr_res in attn_resolutions: | |
| attn.append(resolve_str_to_module(attention)(block_in)) | |
| up = nn.Module() | |
| up.block = block | |
| up.attn = attn | |
| if spatial_upsample[i_level]: | |
| upsample_kwargs = {} | |
| if spatial_upsample[i_level] == "Spatial2xTime2x3DUpsample": | |
| upsample_kwargs["is_first"] = ( | |
| i_level == len(range(self.num_resolutions)) - 1 | |
| ) | |
| up.upsample = resolve_str_to_module(spatial_upsample[i_level])( | |
| block_in, block_in, **upsample_kwargs | |
| ) | |
| curr_res = curr_res * 2 | |
| if temporal_upsample[i_level]: | |
| up.time_upsample = resolve_str_to_module(temporal_upsample[i_level])( | |
| block_in, block_in | |
| ) | |
| self.up.insert(0, up) | |
| # ---- Out ---- | |
| self.norm_out = Normalize(block_in) | |
| self.conv_out = resolve_str_to_module(conv_out)( | |
| block_in, 3, kernel_size=3, padding=1 | |
| ) | |
| def forward(self, z): | |
| h = self.conv_in(z) | |
| h = self.mid.block_1(h) | |
| h = self.mid.attn_1(h) | |
| h = self.mid.block_2(h) | |
| for i_level in reversed(range(self.num_resolutions)): | |
| for i_block in range(self.num_res_blocks + 1): | |
| h = self.up[i_level].block[i_block](h) | |
| if len(self.up[i_level].attn) > 0: | |
| h = self.up[i_level].attn[i_block](h) | |
| if hasattr(self.up[i_level], "upsample"): | |
| h = self.up[i_level].upsample(h) | |
| if hasattr(self.up[i_level], "time_upsample"): | |
| h = self.up[i_level].time_upsample(h) | |
| h = self.norm_out(h) | |
| h = nonlinearity(h) | |
| h = self.conv_out(h) | |
| return h | |
| class VideoVAE(VAE): | |
| """ | |
| Generic Video VAE model with two options: causal and non-causal | |
| (Temporal compression factor of `f_t = 2 ** num_temporal_downsample`). | |
| - Causal: Latents are only dependent on past frames, requires the input to be of length `f_t * k + 1`, which will be encoded to `k + 1` latents. Maps the first frame to the first latent. Causal model allows input to be f_t * l + 1, where 0 <= l <= k. | |
| - Noncausal: Latents are dependent on both past and future frames, requires the input to be of length `f_t * k`, which will be encoded to `k` latents. Maps the first frame to the first latent. Noncausal model allows input to be exactly f_t * k. | |
| TODO: Implement encoding and decoding with tiling for very long videos. | |
| """ | |
| def __init__( | |
| self, | |
| hidden_size: int = 128, | |
| z_channels: int = 4, | |
| hidden_size_mult: Tuple[int] = (1, 2, 4, 4), | |
| attn_resolutions: Tuple[int] = (), | |
| dropout: float = 0.0, | |
| resolution: int = 256, | |
| temporal_length: int = 17, | |
| double_z: bool = True, | |
| embed_dim: int = 4, | |
| num_res_blocks: int = 2, | |
| q_conv: Module = "PaddedConv3D", | |
| encoder_conv_in: Module = "Conv2d", | |
| encoder_conv_out: Module = "PaddedConv3D", | |
| encoder_attention: Module = "AttnBlock3D", | |
| encoder_resnet_blocks: Tuple[Module] = ( | |
| "ResnetBlock2D", | |
| "ResnetBlock2D", | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| ), | |
| encoder_spatial_downsample: Tuple[Module] = ( | |
| "Downsample", | |
| "Spatial2xTime2x3DDownsample", | |
| "Spatial2xTime2x3DDownsample", | |
| "", | |
| ), | |
| encoder_temporal_downsample: Tuple[Module] = ( | |
| "", | |
| "", | |
| "", | |
| "", | |
| ), | |
| encoder_mid_resnet: Module = "ResnetBlock3D", | |
| decoder_conv_in: Module = "PaddedConv3D", | |
| decoder_conv_out: Module = "PaddedConv3D", | |
| decoder_attention: Module = "AttnBlock3D", | |
| decoder_resnet_blocks: Tuple[Module] = ( | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| "ResnetBlock3D", | |
| ), | |
| decoder_spatial_upsample: Tuple[Module] = ( | |
| "", | |
| "SpatialUpsample2x", | |
| "Spatial2xTime2x3DUpsample", | |
| "Spatial2xTime2x3DUpsample", | |
| ), | |
| decoder_temporal_upsample: Tuple[Module] = ( | |
| "", | |
| "", | |
| "", | |
| "", | |
| ), | |
| decoder_mid_resnet: Module = "ResnetBlock3D", | |
| use_quant_layer: bool = True, | |
| is_causal: bool = True, | |
| first_padding_mode: Literal["zero", "same"] = "same", | |
| ) -> None: | |
| super().__init__() | |
| self.is_causal = is_causal | |
| self.temporal_pixel_length = temporal_length | |
| self.temporal_downsampling_factor = 2 ** ( | |
| len([d for d in encoder_spatial_downsample if "Time" in d]) | |
| + len([d for d in encoder_temporal_downsample if d != ""]) | |
| ) | |
| if is_causal: | |
| assert ( | |
| self.temporal_pixel_length % self.temporal_downsampling_factor == 1 | |
| ), f"For causal model, temporal length must be {self.temporal_downsampling_factor} * k + 1" | |
| else: | |
| assert ( | |
| self.temporal_pixel_length % self.temporal_downsampling_factor == 0 | |
| ), f"For non-causal model, temporal length must be {self.temporal_downsampling_factor} * k" | |
| self.temporal_latent_length = ( | |
| self.temporal_pixel_length // self.temporal_downsampling_factor | |
| + (1 if is_causal else 0) | |
| ) | |
| self.use_quant_layer = use_quant_layer | |
| self.first_padding_mode = first_padding_mode | |
| self.encoder = Encoder( | |
| z_channels=z_channels, | |
| hidden_size=hidden_size, | |
| hidden_size_mult=hidden_size_mult, | |
| attn_resolutions=attn_resolutions, | |
| conv_in=encoder_conv_in, | |
| conv_out=encoder_conv_out, | |
| attention=encoder_attention, | |
| resnet_blocks=encoder_resnet_blocks, | |
| spatial_downsample=encoder_spatial_downsample, | |
| temporal_downsample=encoder_temporal_downsample, | |
| mid_resnet=encoder_mid_resnet, | |
| dropout=dropout, | |
| resolution=resolution, | |
| num_res_blocks=num_res_blocks, | |
| double_z=double_z, | |
| is_causal=is_causal, | |
| ) | |
| self.decoder = Decoder( | |
| z_channels=z_channels, | |
| hidden_size=hidden_size, | |
| hidden_size_mult=hidden_size_mult, | |
| attn_resolutions=attn_resolutions, | |
| conv_in=decoder_conv_in, | |
| conv_out=decoder_conv_out, | |
| attention=decoder_attention, | |
| resnet_blocks=decoder_resnet_blocks, | |
| spatial_upsample=decoder_spatial_upsample, | |
| temporal_upsample=decoder_temporal_upsample, | |
| mid_resnet=decoder_mid_resnet, | |
| dropout=dropout, | |
| resolution=resolution, | |
| num_res_blocks=num_res_blocks, | |
| is_causal=is_causal, | |
| ) | |
| if self.use_quant_layer: | |
| quant_conv_cls = _resolve_str_to_module(q_conv, is_causal) | |
| self.quant_conv = quant_conv_cls(2 * z_channels, 2 * embed_dim, 1) | |
| self.post_quant_conv = quant_conv_cls(embed_dim, z_channels, 1) | |
| def _encode(self, x: torch.Tensor) -> torch.Tensor: | |
| h = self.encoder(x) | |
| if self.use_quant_layer: | |
| h = self.quant_conv(h) | |
| return h | |
| def encode(self, x: torch.Tensor) -> DiagonalGaussianDistribution: | |
| """ | |
| Encode a batch of videos to a batch of DiagonalGaussianDistributions. | |
| """ | |
| batch_size, _, temporal_length = x.shape[:3] | |
| if self.is_causal: | |
| assert ( | |
| temporal_length <= self.temporal_pixel_length | |
| and temporal_length % self.temporal_downsampling_factor == 1 | |
| ), f"Temporal length must be {self.temporal_downsampling_factor} * k + 1 where 0 <= k <= {self.temporal_latent_length - 1}, got {temporal_length}" | |
| else: | |
| if temporal_length % self.temporal_pixel_length != 0: | |
| pad = repeat( | |
| ( | |
| x[:, :, :1] | |
| if self.first_padding_mode == "same" | |
| else torch.zeros_like(x[:, :, :1]) | |
| ), | |
| "b c 1 h w -> b c t h w", | |
| t=self.temporal_pixel_length | |
| - temporal_length % self.temporal_pixel_length, | |
| ) | |
| x = torch.cat([pad, x], dim=2) | |
| x = rearrange( | |
| x, "b c (m t) h w -> (b m) c t h w", t=self.temporal_pixel_length | |
| ) | |
| h = self._encode(x) | |
| if h.shape[0] != batch_size: | |
| h = rearrange( | |
| h, | |
| "(b m) c t h w -> b c (m t) h w", | |
| b=batch_size, | |
| ) | |
| return DiagonalGaussianDistribution(h) | |
| def _decode(self, z: torch.Tensor) -> torch.Tensor: | |
| if self.use_quant_layer: | |
| z = self.post_quant_conv(z) | |
| dec = self.decoder(z) | |
| return dec | |
| def decode( | |
| self, z: torch.Tensor, desired_length: Optional[int] = None | |
| ) -> torch.Tensor: | |
| batch_size, _, temporal_latent_length = z.shape[:3] | |
| if not self.is_causal: | |
| assert ( | |
| temporal_latent_length % self.temporal_latent_length == 0 | |
| ), f"Temporal latent length must be a multiple of {self.temporal_latent_length}, got {temporal_latent_length}" | |
| z = rearrange( | |
| z, | |
| "b c (m t) h w -> (b m) c t h w", | |
| t=self.temporal_latent_length, | |
| ) | |
| dec = self._decode(z) | |
| if dec.shape[0] != batch_size: | |
| dec = rearrange( | |
| dec, | |
| "(b m) c t h w -> b c (m t) h w", | |
| b=batch_size, | |
| ) | |
| if desired_length is not None: | |
| dec = dec[:, :, -desired_length:] | |
| assert ( | |
| dec.shape[2] == desired_length | |
| ), f"Desired length {desired_length} does not match decoded length {dec.shape[2]}" | |
| return dec | |
| def forward( | |
| self, sample: torch.Tensor, sample_posterior: bool = True | |
| ) -> Tuple[torch.Tensor, DiagonalGaussianDistribution]: | |
| posterior = self.encode(sample) | |
| if sample_posterior: | |
| z = posterior.sample() | |
| else: | |
| z = posterior.mode() | |
| dec = self.decode(z, desired_length=sample.shape[2]) | |
| return dec, posterior | |
| def get_last_layer(self): | |
| if hasattr(self.decoder.conv_out, "conv"): | |
| return self.decoder.conv_out.conv.weight | |
| else: | |
| return self.decoder.conv_out.weight | |
| def from_pretrained(cls, path: str, **kwargs) -> "VideoVAE": | |
| if is_wandb_run_path(path): | |
| path = wandb_to_local_path(path) | |
| elif is_hf_path(path): | |
| path = hf_to_local_path(path) | |
| checkpoint = torch.load(path, map_location="cpu") | |
| model_cfg = checkpoint["model_cfg"] | |
| for key, value in model_cfg.items(): | |
| if isinstance(value, list): | |
| model_cfg[key] = tuple(value) | |
| model = cls(**model_cfg) | |
| if ( | |
| len(checkpoint["optimizer_states"]) > 0 | |
| and "ema" in checkpoint["optimizer_states"][0] | |
| ): | |
| state_dict = dict( | |
| zip( | |
| [name for name, _ in model.named_parameters()], | |
| checkpoint["optimizer_states"][0]["ema"], | |
| ) | |
| ) | |
| else: | |
| state_dict = { | |
| key.replace("vae.", ""): value | |
| for key, value in checkpoint["state_dict"].items() | |
| if key.startswith("vae.") | |
| } | |
| model.load_state_dict(state_dict) | |
| return model | |