diff --git a/configs_3b/main.yaml b/configs_3b/main.yaml index b5a6993..d2ffda4 100644 --- a/configs_3b/main.yaml +++ b/configs_3b/main.yaml @@ -6,9 +6,9 @@ dit: model: __object__: path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit_v2.nadit" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit_v2.nadit" - - "models.dit_v2.nadit" + - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit_v2.nadit" + - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit_v2.nadit" + - "src.models.dit_v2.nadit" name: "NaDiT" args: "as_params" vid_in_channels: 33 @@ -49,9 +49,9 @@ vae: model: __object__: path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae" - - "models.video_vae_v3.modules.attn_video_vae" + - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" + - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" + - "src.models.video_vae_v3.modules.attn_video_vae" name: "VideoAutoencoderKLWrapper" args: "as_params" freeze_encoder: False diff --git a/configs_7b/main.yaml b/configs_7b/main.yaml index d0cb203..3f80813 100644 --- a/configs_7b/main.yaml +++ b/configs_7b/main.yaml @@ -6,9 +6,9 @@ dit: model: __object__: path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit.nadit" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit.nadit" - - "models.dit.nadit" + - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit.nadit" + - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit.nadit" + - "src.models.dit.nadit" name: "NaDiT" args: "as_params" vid_in_channels: 33 @@ -46,9 +46,9 @@ vae: model: __object__: path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae" - - "models.video_vae_v3.modules.attn_video_vae" + - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" + - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" + - "src.models.video_vae_v3.modules.attn_video_vae" name: "VideoAutoencoderKLWrapper" args: "as_params" freeze_encoder: False diff --git a/common/__init__.py b/src/common/__init__.py similarity index 100% rename from common/__init__.py rename to src/common/__init__.py diff --git a/common/cache.py b/src/common/cache.py similarity index 100% rename from common/cache.py rename to src/common/cache.py diff --git a/common/config.py b/src/common/config.py similarity index 100% rename from common/config.py rename to src/common/config.py diff --git a/common/decorators.py b/src/common/decorators.py similarity index 100% rename from common/decorators.py rename to src/common/decorators.py diff --git a/common/diffusion/__init__.py b/src/common/diffusion/__init__.py similarity index 100% rename from common/diffusion/__init__.py rename to src/common/diffusion/__init__.py diff --git a/common/diffusion/config.py b/src/common/diffusion/config.py similarity index 100% rename from common/diffusion/config.py rename to src/common/diffusion/config.py diff --git a/common/diffusion/samplers/base.py b/src/common/diffusion/samplers/base.py similarity index 100% rename from common/diffusion/samplers/base.py rename to src/common/diffusion/samplers/base.py diff --git a/common/diffusion/samplers/euler.py b/src/common/diffusion/samplers/euler.py similarity index 100% rename from common/diffusion/samplers/euler.py rename to src/common/diffusion/samplers/euler.py diff --git a/common/diffusion/schedules/base.py b/src/common/diffusion/schedules/base.py similarity index 100% rename from common/diffusion/schedules/base.py rename to src/common/diffusion/schedules/base.py diff --git a/common/diffusion/schedules/lerp.py b/src/common/diffusion/schedules/lerp.py similarity index 100% rename from common/diffusion/schedules/lerp.py rename to src/common/diffusion/schedules/lerp.py diff --git a/common/diffusion/timesteps/base.py b/src/common/diffusion/timesteps/base.py similarity index 100% rename from common/diffusion/timesteps/base.py rename to src/common/diffusion/timesteps/base.py diff --git a/common/diffusion/timesteps/sampling/trailing.py b/src/common/diffusion/timesteps/sampling/trailing.py similarity index 100% rename from common/diffusion/timesteps/sampling/trailing.py rename to src/common/diffusion/timesteps/sampling/trailing.py diff --git a/common/diffusion/types.py b/src/common/diffusion/types.py similarity index 100% rename from common/diffusion/types.py rename to src/common/diffusion/types.py diff --git a/common/diffusion/utils.py b/src/common/diffusion/utils.py similarity index 100% rename from common/diffusion/utils.py rename to src/common/diffusion/utils.py diff --git a/common/distributed/__init__.py b/src/common/distributed/__init__.py similarity index 100% rename from common/distributed/__init__.py rename to src/common/distributed/__init__.py diff --git a/common/distributed/advanced.py b/src/common/distributed/advanced.py similarity index 100% rename from common/distributed/advanced.py rename to src/common/distributed/advanced.py diff --git a/common/distributed/basic.py b/src/common/distributed/basic.py similarity index 100% rename from common/distributed/basic.py rename to src/common/distributed/basic.py diff --git a/common/distributed/meta_init_utils.py b/src/common/distributed/meta_init_utils.py similarity index 100% rename from common/distributed/meta_init_utils.py rename to src/common/distributed/meta_init_utils.py diff --git a/common/distributed/ops.py b/src/common/distributed/ops.py similarity index 100% rename from common/distributed/ops.py rename to src/common/distributed/ops.py diff --git a/common/half_precision_fixes.py b/src/common/half_precision_fixes.py similarity index 100% rename from common/half_precision_fixes.py rename to src/common/half_precision_fixes.py diff --git a/common/logger.py b/src/common/logger.py similarity index 100% rename from common/logger.py rename to src/common/logger.py diff --git a/common/partition.py b/src/common/partition.py similarity index 100% rename from common/partition.py rename to src/common/partition.py diff --git a/common/seed.py b/src/common/seed.py similarity index 100% rename from common/seed.py rename to src/common/seed.py diff --git a/src/core/generation.py b/src/core/generation.py index 90e3514..bf5f1e8 100644 --- a/src/core/generation.py +++ b/src/core/generation.py @@ -28,15 +28,15 @@ from src.optimization.performance import ( optimized_video_rearrange, optimized_single_video_rearrange, optimized_sample_to_image_format, temporal_latent_blending ) -from common.seed import set_seed +from src.common.seed import set_seed import comfy.model_management # Get script directory for embeddings script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) # Import transforms and color fix -from data.image.transforms.divisible_crop import DivisibleCrop -from data.image.transforms.na_resize import NaResize +from src.data.image.transforms.divisible_crop import DivisibleCrop +from src.data.image.transforms.na_resize import NaResize from src.utils.color_fix import wavelet_reconstruction diff --git a/src/core/infer.py b/src/core/infer.py index 39e5897..ffb255b 100644 --- a/src/core/infer.py +++ b/src/core/infer.py @@ -20,27 +20,27 @@ from einops import rearrange from omegaconf import DictConfig, ListConfig from torch import Tensor from src.optimization.memory_manager import clear_vram_cache -from models.video_vae_v3.modules.types import MemoryState +from src.models.video_vae_v3.modules.types import MemoryState -from common.config import create_object -from common.decorators import log_on_entry, log_runtime -from common.diffusion import ( +from src.common.config import create_object +from src.common.decorators import log_on_entry, log_runtime +from src.common.diffusion import ( classifier_free_guidance_dispatcher, create_sampler_from_config, create_sampling_timesteps_from_config, create_schedule_from_config, ) -from common.distributed import ( +from src.common.distributed import ( get_device, get_global_rank, ) -from common.distributed.meta_init_utils import ( +from src.common.distributed.meta_init_utils import ( meta_non_persistent_buffer_init_fn, ) # from common.fs import download -from models.dit_v2 import na +from src.models.dit_v2 import na def optimized_channels_to_last(tensor): diff --git a/src/core/model_manager.py b/src/core/model_manager.py index 76a0acb..452673f 100644 --- a/src/core/model_manager.py +++ b/src/core/model_manager.py @@ -31,7 +31,7 @@ except ImportError: from src.optimization.memory_manager import get_basic_vram_info, clear_vram_cache from src.optimization.compatibility import FP8CompatibleDiT from src.optimization.memory_manager import preinitialize_rope_cache -from common.config import load_config, create_object +from src.common.config import load_config, create_object from src.core.infer import VideoDiffusionInfer # NOUVEAU: Import des opérations ComfyUI pour FP8 @@ -77,7 +77,7 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False): # No need for dynamic path resolution here anymore! # Load and configure VAE with additional parameters - vae_config_path = os.path.join(script_directory, 'models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml') + vae_config_path = os.path.join(script_directory, 'src/models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml') t = time.time() vae_config = OmegaConf.load(vae_config_path) if debug: diff --git a/data/image/transforms/area_resize.py b/src/data/image/transforms/area_resize.py similarity index 100% rename from data/image/transforms/area_resize.py rename to src/data/image/transforms/area_resize.py diff --git a/data/image/transforms/divisible_crop.py b/src/data/image/transforms/divisible_crop.py similarity index 100% rename from data/image/transforms/divisible_crop.py rename to src/data/image/transforms/divisible_crop.py diff --git a/data/image/transforms/na_resize.py b/src/data/image/transforms/na_resize.py similarity index 100% rename from data/image/transforms/na_resize.py rename to src/data/image/transforms/na_resize.py diff --git a/data/image/transforms/side_resize.py b/src/data/image/transforms/side_resize.py similarity index 100% rename from data/image/transforms/side_resize.py rename to src/data/image/transforms/side_resize.py diff --git a/models/dit/attention.py b/src/models/dit/attention.py similarity index 100% rename from models/dit/attention.py rename to src/models/dit/attention.py diff --git a/models/dit/blocks/__init__.py b/src/models/dit/blocks/__init__.py similarity index 100% rename from models/dit/blocks/__init__.py rename to src/models/dit/blocks/__init__.py diff --git a/models/dit/blocks/mmdit_window_block.py b/src/models/dit/blocks/mmdit_window_block.py similarity index 100% rename from models/dit/blocks/mmdit_window_block.py rename to src/models/dit/blocks/mmdit_window_block.py diff --git a/models/dit/embedding.py b/src/models/dit/embedding.py similarity index 100% rename from models/dit/embedding.py rename to src/models/dit/embedding.py diff --git a/models/dit/mlp.py b/src/models/dit/mlp.py similarity index 100% rename from models/dit/mlp.py rename to src/models/dit/mlp.py diff --git a/models/dit/mm.py b/src/models/dit/mm.py similarity index 100% rename from models/dit/mm.py rename to src/models/dit/mm.py diff --git a/models/dit/modulation.py b/src/models/dit/modulation.py similarity index 100% rename from models/dit/modulation.py rename to src/models/dit/modulation.py diff --git a/models/dit/na.py b/src/models/dit/na.py similarity index 100% rename from models/dit/na.py rename to src/models/dit/na.py diff --git a/models/dit/nablocks/__init__.py b/src/models/dit/nablocks/__init__.py similarity index 100% rename from models/dit/nablocks/__init__.py rename to src/models/dit/nablocks/__init__.py diff --git a/models/dit/nablocks/mmsr_block.py b/src/models/dit/nablocks/mmsr_block.py similarity index 100% rename from models/dit/nablocks/mmsr_block.py rename to src/models/dit/nablocks/mmsr_block.py diff --git a/models/dit/nadit.py b/src/models/dit/nadit.py similarity index 100% rename from models/dit/nadit.py rename to src/models/dit/nadit.py diff --git a/models/dit/normalization.py b/src/models/dit/normalization.py similarity index 100% rename from models/dit/normalization.py rename to src/models/dit/normalization.py diff --git a/models/dit/patch.py b/src/models/dit/patch.py similarity index 100% rename from models/dit/patch.py rename to src/models/dit/patch.py diff --git a/models/dit/rope.py b/src/models/dit/rope.py similarity index 100% rename from models/dit/rope.py rename to src/models/dit/rope.py diff --git a/models/dit/window.py b/src/models/dit/window.py similarity index 100% rename from models/dit/window.py rename to src/models/dit/window.py diff --git a/models/dit_v2/attention.py b/src/models/dit_v2/attention.py similarity index 100% rename from models/dit_v2/attention.py rename to src/models/dit_v2/attention.py diff --git a/models/dit_v2/embedding.py b/src/models/dit_v2/embedding.py similarity index 100% rename from models/dit_v2/embedding.py rename to src/models/dit_v2/embedding.py diff --git a/models/dit_v2/mlp.py b/src/models/dit_v2/mlp.py similarity index 100% rename from models/dit_v2/mlp.py rename to src/models/dit_v2/mlp.py diff --git a/models/dit_v2/mm.py b/src/models/dit_v2/mm.py similarity index 100% rename from models/dit_v2/mm.py rename to src/models/dit_v2/mm.py diff --git a/models/dit_v2/modulation.py b/src/models/dit_v2/modulation.py similarity index 100% rename from models/dit_v2/modulation.py rename to src/models/dit_v2/modulation.py diff --git a/models/dit_v2/na.py b/src/models/dit_v2/na.py similarity index 100% rename from models/dit_v2/na.py rename to src/models/dit_v2/na.py diff --git a/models/dit_v2/nablocks/__init__.py b/src/models/dit_v2/nablocks/__init__.py similarity index 100% rename from models/dit_v2/nablocks/__init__.py rename to src/models/dit_v2/nablocks/__init__.py diff --git a/models/dit_v2/nablocks/attention/__init__.py b/src/models/dit_v2/nablocks/attention/__init__.py similarity index 100% rename from models/dit_v2/nablocks/attention/__init__.py rename to src/models/dit_v2/nablocks/attention/__init__.py diff --git a/models/dit_v2/nablocks/attention/mmattn.py b/src/models/dit_v2/nablocks/attention/mmattn.py similarity index 100% rename from models/dit_v2/nablocks/attention/mmattn.py rename to src/models/dit_v2/nablocks/attention/mmattn.py diff --git a/models/dit_v2/nablocks/mmsr_block.py b/src/models/dit_v2/nablocks/mmsr_block.py similarity index 100% rename from models/dit_v2/nablocks/mmsr_block.py rename to src/models/dit_v2/nablocks/mmsr_block.py diff --git a/models/dit_v2/nadit.py b/src/models/dit_v2/nadit.py similarity index 100% rename from models/dit_v2/nadit.py rename to src/models/dit_v2/nadit.py diff --git a/models/dit_v2/normalization.py b/src/models/dit_v2/normalization.py similarity index 100% rename from models/dit_v2/normalization.py rename to src/models/dit_v2/normalization.py diff --git a/models/dit_v2/patch/__init__.py b/src/models/dit_v2/patch/__init__.py similarity index 100% rename from models/dit_v2/patch/__init__.py rename to src/models/dit_v2/patch/__init__.py diff --git a/models/dit_v2/patch/patch_v1.py b/src/models/dit_v2/patch/patch_v1.py similarity index 100% rename from models/dit_v2/patch/patch_v1.py rename to src/models/dit_v2/patch/patch_v1.py diff --git a/models/dit_v2/rope.py b/src/models/dit_v2/rope.py similarity index 99% rename from models/dit_v2/rope.py rename to src/models/dit_v2/rope.py index 3525c88..dde5b47 100644 --- a/models/dit_v2/rope.py +++ b/src/models/dit_v2/rope.py @@ -19,7 +19,7 @@ from einops import rearrange from rotary_embedding_torch import RotaryEmbedding, apply_rotary_emb from torch import nn -from common.cache import Cache +from src.common.cache import Cache class RotaryEmbeddingBase(nn.Module): diff --git a/models/dit_v2/window.py b/src/models/dit_v2/window.py similarity index 100% rename from models/dit_v2/window.py rename to src/models/dit_v2/window.py diff --git a/models/video_vae_v3/modules/attn_video_vae.py b/src/models/video_vae_v3/modules/attn_video_vae.py similarity index 100% rename from models/video_vae_v3/modules/attn_video_vae.py rename to src/models/video_vae_v3/modules/attn_video_vae.py diff --git a/models/video_vae_v3/modules/causal_inflation_lib.py b/src/models/video_vae_v3/modules/causal_inflation_lib.py similarity index 100% rename from models/video_vae_v3/modules/causal_inflation_lib.py rename to src/models/video_vae_v3/modules/causal_inflation_lib.py diff --git a/models/video_vae_v3/modules/context_parallel_lib.py b/src/models/video_vae_v3/modules/context_parallel_lib.py similarity index 100% rename from models/video_vae_v3/modules/context_parallel_lib.py rename to src/models/video_vae_v3/modules/context_parallel_lib.py diff --git a/models/video_vae_v3/modules/global_config.py b/src/models/video_vae_v3/modules/global_config.py similarity index 100% rename from models/video_vae_v3/modules/global_config.py rename to src/models/video_vae_v3/modules/global_config.py diff --git a/models/video_vae_v3/modules/inflated_layers.py b/src/models/video_vae_v3/modules/inflated_layers.py similarity index 100% rename from models/video_vae_v3/modules/inflated_layers.py rename to src/models/video_vae_v3/modules/inflated_layers.py diff --git a/models/video_vae_v3/modules/inflated_lib.py b/src/models/video_vae_v3/modules/inflated_lib.py similarity index 100% rename from models/video_vae_v3/modules/inflated_lib.py rename to src/models/video_vae_v3/modules/inflated_lib.py diff --git a/models/video_vae_v3/modules/types.py b/src/models/video_vae_v3/modules/types.py similarity index 100% rename from models/video_vae_v3/modules/types.py rename to src/models/video_vae_v3/modules/types.py diff --git a/models/video_vae_v3/modules/video_vae.py b/src/models/video_vae_v3/modules/video_vae.py similarity index 100% rename from models/video_vae_v3/modules/video_vae.py rename to src/models/video_vae_v3/modules/video_vae.py diff --git a/models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml b/src/models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml similarity index 100% rename from models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml rename to src/models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml diff --git a/src/models/video_vae_v3_mine_bad/modules/attn_video_vae.py b/src/models/video_vae_v3_mine_bad/modules/attn_video_vae.py new file mode 100644 index 0000000..9584c83 --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/attn_video_vae.py @@ -0,0 +1,1361 @@ +# Copyright (c) 2023 HuggingFace Team +# Copyright (c) 2025 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache License, Version 2.0 (the "License") +# +# This file has been modified by ByteDance Ltd. and/or its affiliates. on 1st June 2025 +# +# Original file was released under Apache License, Version 2.0 (the "License"), with the full license text +# available at http://www.apache.org/licenses/LICENSE-2.0. +# +# This modified file is released under the same license. + + +from contextlib import nullcontext +from typing import Literal, Optional, Tuple, Union +import diffusers +import torch +import torch.nn as nn +import torch.nn.functional as F +from diffusers.models.attention_processor import Attention, SpatialNorm +from diffusers.models.autoencoders.vae import DecoderOutput, DiagonalGaussianDistribution +from diffusers.models.downsampling import Downsample2D +from diffusers.models.lora import LoRACompatibleConv +from diffusers.models.modeling_outputs import AutoencoderKLOutput +from diffusers.models.resnet import ResnetBlock2D +from diffusers.models.unets.unet_2d_blocks import DownEncoderBlock2D, UpDecoderBlock2D +from diffusers.models.upsampling import Upsample2D +from diffusers.utils import is_torch_version +from diffusers.utils.accelerate_utils import apply_forward_hook +from einops import rearrange +from ....common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation + +from ....common.distributed.advanced import get_sequence_parallel_world_size +from ....common.logger import get_logger +from .causal_inflation_lib import ( + InflatedCausalConv3d, + causal_norm_wrapper, + init_causal_conv3d, + remove_head, +) +from .context_parallel_lib import ( + causal_conv_gather_outputs, + causal_conv_slice_inputs, +) +from .global_config import set_norm_limit +from .types import ( + CausalAutoencoderOutput, + CausalDecoderOutput, + CausalEncoderOutput, + MemoryState, + _inflation_mode_t, + _memory_device_t, + _receptive_field_t, +) + +logger = get_logger(__name__) # pylint: disable=invalid-name + + +class Upsample3D(Upsample2D): + """A 3D upsampling layer with an optional convolution.""" + + def __init__( + self, + *args, + inflation_mode: _inflation_mode_t = "tail", + temporal_up: bool = False, + spatial_up: bool = True, + slicing: bool = False, + **kwargs, + ): + super().__init__(*args, **kwargs) + conv = self.conv if self.name == "conv" else self.Conv2d_0 + + assert type(conv) is not nn.ConvTranspose2d + # Note: lora_layer is not passed into constructor in the original implementation. + # So we make a simplification. + conv = init_causal_conv3d( + self.channels, + self.out_channels, + 3, + padding=1, + inflation_mode=inflation_mode, + ) + + self.temporal_up = temporal_up + self.spatial_up = spatial_up + self.temporal_ratio = 2 if temporal_up else 1 + self.spatial_ratio = 2 if spatial_up else 1 + self.slicing = slicing + + assert not self.interpolate + # [Override] MAGViT v2 implementation + if not self.interpolate: + upscale_ratio = (self.spatial_ratio**2) * self.temporal_ratio + self.upscale_conv = nn.Conv3d( + self.channels, self.channels * upscale_ratio, kernel_size=1, padding=0 + ) + identity = ( + torch.eye(self.channels) + .repeat(upscale_ratio, 1) + .reshape_as(self.upscale_conv.weight) + ) + self.upscale_conv.weight.data.copy_(identity) + nn.init.zeros_(self.upscale_conv.bias) + + if self.name == "conv": + self.conv = conv + else: + self.Conv2d_0 = conv + + def forward( + self, + hidden_states: torch.FloatTensor, + output_size: Optional[int] = None, + memory_state: MemoryState = MemoryState.DISABLED, + preserve_vram: bool = False, + **kwargs, + ) -> torch.FloatTensor: + assert hidden_states.shape[1] == self.channels + + if hasattr(self, "norm") and self.norm is not None: + # [Overridden] change to causal norm. + hidden_states = causal_norm_wrapper(self.norm, hidden_states) + + if self.use_conv_transpose: + return self.conv(hidden_states) + + if self.slicing: + split_size = hidden_states.size(2) // 2 + hidden_states = list( + hidden_states.split([split_size, hidden_states.size(2) - split_size], dim=2) + ) + else: + hidden_states = [hidden_states] + # ADD BY NUMZ + if preserve_vram: + torch.cuda.empty_cache() + for i in range(len(hidden_states)): + hidden_states[i] = self.upscale_conv(hidden_states[i]) + hidden_states[i] = rearrange( + hidden_states[i], + "b (x y z c) f h w -> b c (f z) (h x) (w y)", + x=self.spatial_ratio, + y=self.spatial_ratio, + z=self.temporal_ratio, + ) + + # [Overridden] For causal temporal conv + if self.temporal_up and memory_state != MemoryState.ACTIVE: + hidden_states[0] = remove_head(hidden_states[0]) + + if not self.slicing: + hidden_states = hidden_states[0] + # ADD BY NUMZ + if preserve_vram: + torch.cuda.empty_cache() + if self.use_conv: + if self.name == "conv": + hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram) + else: + hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state) + + if not self.slicing: + return hidden_states + else: + return torch.cat(hidden_states, dim=2) + + +class Downsample3D(Downsample2D): + """A 3D downsampling layer with an optional convolution.""" + + def __init__( + self, + *args, + inflation_mode: _inflation_mode_t = "tail", + spatial_down: bool = False, + temporal_down: bool = False, + **kwargs, + ): + super().__init__(*args, **kwargs) + conv = self.conv + self.temporal_down = temporal_down + self.spatial_down = spatial_down + + self.temporal_ratio = 2 if temporal_down else 1 + self.spatial_ratio = 2 if spatial_down else 1 + + self.temporal_kernel = 3 if temporal_down else 1 + self.spatial_kernel = 3 if spatial_down else 1 + + if type(conv) in [nn.Conv2d, LoRACompatibleConv]: + # Note: lora_layer is not passed into constructor in the original implementation. + # So we make a simplification. + conv = init_causal_conv3d( + self.channels, + self.out_channels, + kernel_size=(self.temporal_kernel, self.spatial_kernel, self.spatial_kernel), + stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), + padding=( + 1 if self.temporal_down else 0, + self.padding if self.spatial_down else 0, + self.padding if self.spatial_down else 0, + ), + inflation_mode=inflation_mode, + ) + elif type(conv) is nn.AvgPool2d: + assert self.channels == self.out_channels + conv = nn.AvgPool3d( + kernel_size=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), + stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), + ) + else: + raise NotImplementedError + + if self.name == "conv": + self.Conv2d_0 = conv + self.conv = conv + else: + self.conv = conv + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state: MemoryState = MemoryState.DISABLED, + preserve_vram: bool = False, + **kwargs, + ) -> torch.FloatTensor: + + assert hidden_states.shape[1] == self.channels + + if hasattr(self, "norm") and self.norm is not None: + # [Overridden] change to causal norm. + hidden_states = causal_norm_wrapper(self.norm, hidden_states) + + if self.use_conv and self.padding == 0 and self.spatial_down: + pad = (0, 1, 0, 1) + hidden_states = safe_pad_operation(hidden_states, pad, mode="constant", value=0) + + assert hidden_states.shape[1] == self.channels + + hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram) + + return hidden_states + + +class ResnetBlock3D(ResnetBlock2D): + def __init__( + self, + *args, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + slicing: bool = False, + **kwargs, + ): + super().__init__(*args, **kwargs) + self.conv1 = init_causal_conv3d( + self.in_channels, + self.out_channels, + kernel_size=(1, 3, 3) if time_receptive_field == "half" else (3, 3, 3), + stride=1, + padding=(0, 1, 1) if time_receptive_field == "half" else (1, 1, 1), + inflation_mode=inflation_mode, + ) + + self.conv2 = init_causal_conv3d( + self.out_channels, + self.conv2.out_channels, + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + if self.up: + assert type(self.upsample) is Upsample2D + self.upsample = Upsample3D( + self.in_channels, + use_conv=False, + inflation_mode=inflation_mode, + slicing=slicing, + ) + elif self.down: + assert type(self.downsample) is Downsample2D + self.downsample = Downsample3D( + self.in_channels, + use_conv=False, + padding=1, + name="op", + inflation_mode=inflation_mode, + ) + + if self.use_in_shortcut: + self.conv_shortcut = init_causal_conv3d( + self.in_channels, + self.conv_shortcut.out_channels, + kernel_size=1, + stride=1, + padding=0, + bias=(self.conv_shortcut.bias is not None), + inflation_mode=inflation_mode, + ) + + def forward( + self, input_tensor, temb, memory_state: MemoryState = MemoryState.DISABLED, preserve_vram: bool = False, **kwargs + ): + hidden_states = input_tensor + + hidden_states = causal_norm_wrapper(self.norm1, hidden_states, preserve_vram=preserve_vram) + + hidden_states = self.nonlinearity(hidden_states) + + if self.upsample is not None: + # upsample_nearest_nhwc fails with large batch sizes. + # see https://github.com/huggingface/diffusers/issues/984 + if hidden_states.shape[0] >= 64: + input_tensor = input_tensor.contiguous() + hidden_states = hidden_states.contiguous() + input_tensor = self.upsample(input_tensor, memory_state=memory_state) + hidden_states = self.upsample(hidden_states, memory_state=memory_state) + elif self.downsample is not None: + input_tensor = self.downsample(input_tensor, memory_state=memory_state) + hidden_states = self.downsample(hidden_states, memory_state=memory_state) + + hidden_states = self.conv1(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram) + + if self.time_emb_proj is not None: + if not self.skip_time_act: + temb = self.nonlinearity(temb) + temb = self.time_emb_proj(temb)[:, :, None, None] + + if temb is not None and self.time_embedding_norm == "default": + hidden_states = hidden_states + temb + + hidden_states = causal_norm_wrapper(self.norm2, hidden_states) + + if temb is not None and self.time_embedding_norm == "scale_shift": + scale, shift = torch.chunk(temb, 2, dim=1) + hidden_states = hidden_states * (1 + scale) + shift + + hidden_states = self.nonlinearity(hidden_states) + + hidden_states = self.dropout(hidden_states) + hidden_states = self.conv2(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram) + + if self.conv_shortcut is not None: + input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state, preserve_vram=preserve_vram) + + output_tensor = (input_tensor + hidden_states) / self.output_scale_factor + + return output_tensor + + +class DownEncoderBlock3D(DownEncoderBlock2D): + def __init__( + self, + in_channels: int, + out_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + output_scale_factor: float = 1.0, + add_downsample: bool = True, + downsample_padding: int = 1, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_down: bool = True, + spatial_down: bool = True, + ): + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + dropout=dropout, + num_layers=num_layers, + resnet_eps=resnet_eps, + resnet_time_scale_shift=resnet_time_scale_shift, + resnet_act_fn=resnet_act_fn, + resnet_groups=resnet_groups, + resnet_pre_norm=resnet_pre_norm, + output_scale_factor=output_scale_factor, + add_downsample=add_downsample, + downsample_padding=downsample_padding, + ) + resnets = [] + temporal_modules = [] + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + # [Override] Replace module. + ResnetBlock3D( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=None, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + temporal_modules.append(nn.Identity()) + + self.resnets = nn.ModuleList(resnets) + self.temporal_modules = nn.ModuleList(temporal_modules) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + # [Override] Replace module. + Downsample3D( + out_channels, + use_conv=True, + out_channels=out_channels, + padding=downsample_padding, + name="op", + temporal_down=temporal_down, + spatial_down=spatial_down, + inflation_mode=inflation_mode, + ) + ] + ) + else: + self.downsamplers = None + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state: MemoryState = MemoryState.DISABLED, + preserve_vram: bool = False, + **kwargs, + ) -> torch.FloatTensor: + for resnet, temporal in zip(self.resnets, self.temporal_modules): + hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, preserve_vram=preserve_vram) + hidden_states = temporal(hidden_states) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram) + + return hidden_states + + +class UpDecoderBlock3D(UpDecoderBlock2D): + def __init__( + self, + in_channels: int, + out_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", # default, spatial + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + output_scale_factor: float = 1.0, + add_upsample: bool = True, + temb_channels: Optional[int] = None, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up: bool = True, + spatial_up: bool = True, + slicing: bool = False, + ): + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + dropout=dropout, + num_layers=num_layers, + resnet_eps=resnet_eps, + resnet_time_scale_shift=resnet_time_scale_shift, + resnet_act_fn=resnet_act_fn, + resnet_groups=resnet_groups, + resnet_pre_norm=resnet_pre_norm, + output_scale_factor=output_scale_factor, + add_upsample=add_upsample, + temb_channels=temb_channels, + ) + resnets = [] + temporal_modules = [] + + for i in range(num_layers): + input_channels = in_channels if i == 0 else out_channels + + resnets.append( + # [Override] Replace module. + ResnetBlock3D( + in_channels=input_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + slicing=slicing, + ) + ) + + temporal_modules.append(nn.Identity()) + + self.resnets = nn.ModuleList(resnets) + self.temporal_modules = nn.ModuleList(temporal_modules) + + if add_upsample: + # [Override] Replace module & use learnable upsample + self.upsamplers = nn.ModuleList( + [ + Upsample3D( + out_channels, + use_conv=True, + out_channels=out_channels, + temporal_up=temporal_up, + spatial_up=spatial_up, + interpolate=False, + inflation_mode=inflation_mode, + slicing=slicing, + ) + ] + ) + else: + self.upsamplers = None + + def forward( + self, + hidden_states: torch.FloatTensor, + temb: Optional[torch.FloatTensor] = None, + memory_state: MemoryState = MemoryState.DISABLED, + preserve_vram: bool = False, + ) -> torch.FloatTensor: + for resnet, temporal in zip(self.resnets, self.temporal_modules): + hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, preserve_vram=preserve_vram) + hidden_states = temporal(hidden_states) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram) + + return hidden_states + + +class UNetMidBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", # default, spatial + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + add_attention: bool = True, + attention_head_dim: int = 1, + output_scale_factor: float = 1.0, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) + self.add_attention = add_attention + + # there is always at least one resnet + resnets = [ + # [Override] Replace module. + ResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ] + attentions = [] + + if attention_head_dim is None: + logger.warn( + f"It is not recommend to pass `attention_head_dim=None`. " + f"Defaulting `attention_head_dim` to `in_channels`: {in_channels}." + ) + attention_head_dim = in_channels + + for _ in range(num_layers): + if self.add_attention: + attentions.append( + Attention( + in_channels, + heads=in_channels // attention_head_dim, + dim_head=attention_head_dim, + rescale_output_factor=output_scale_factor, + eps=resnet_eps, + norm_num_groups=( + resnet_groups if resnet_time_scale_shift == "default" else None + ), + spatial_norm_dim=( + temb_channels if resnet_time_scale_shift == "spatial" else None + ), + residual_connection=True, + bias=True, + upcast_softmax=True, + _from_deprecated_attn_block=True, + ) + ) + else: + attentions.append(None) + + resnets.append( + ResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + def forward(self, hidden_states, temb=None, memory_state: MemoryState = MemoryState.DISABLED): + video_length, frame_height, frame_width = hidden_states.size()[-3:] + hidden_states = self.resnets[0](hidden_states, temb, memory_state=memory_state) + for attn, resnet in zip(self.attentions, self.resnets[1:]): + if attn is not None: + hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w") + hidden_states = attn(hidden_states, temb=temb) + hidden_states = rearrange( + hidden_states, "(b f) c h w -> b c f h w", f=video_length + ) + hidden_states = resnet(hidden_states, temb, memory_state=memory_state) + + return hidden_states + + +class Encoder3D(nn.Module): + r""" + [Override] override most logics to support extra condition input and causal conv + + The `Encoder` layer of a variational autoencoder that encodes + its input into a latent representation. + + Args: + in_channels (`int`, *optional*, defaults to 3): + The number of input channels. + out_channels (`int`, *optional*, defaults to 3): + The number of output channels. + down_block_types (`Tuple[str, ...]`, *optional*, defaults to `("DownEncoderBlock2D",)`): + The types of down blocks to use. + See `~diffusers.models.unet_2d_blocks.get_down_block` + for available options. + block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`): + The number of output channels for each block. + layers_per_block (`int`, *optional*, defaults to 2): + The number of layers per block. + norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups for normalization. + act_fn (`str`, *optional*, defaults to `"silu"`): + The activation function to use. + See `~diffusers.models.activations.get_activation` for available options. + double_z (`bool`, *optional*, defaults to `True`): + Whether to double the number of output channels for the last block. + """ + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + down_block_types: Tuple[str, ...] = ("DownEncoderBlock3D",), + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + norm_num_groups: int = 32, + act_fn: str = "silu", + double_z: bool = True, + mid_block_add_attention=True, + # [Override] add extra_cond_dim, temporal down num + temporal_down_num: int = 2, + extra_cond_dim: int = None, + gradient_checkpoint: bool = False, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_down_num = temporal_down_num + + self.conv_in = init_causal_conv3d( + in_channels, + block_out_channels[0], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.mid_block = None + self.down_blocks = nn.ModuleList([]) + self.extra_cond_dim = extra_cond_dim + + self.conv_extra_cond = nn.ModuleList([]) + + # down + output_channel = block_out_channels[0] + for i, down_block_type in enumerate(down_block_types): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + # [Override] to support temporal down block design + is_temporal_down_block = i >= len(block_out_channels) - self.temporal_down_num - 1 + # Note: take the last ones + + assert down_block_type == "DownEncoderBlock3D" + + down_block = DownEncoderBlock3D( + num_layers=self.layers_per_block, + in_channels=input_channel, + out_channels=output_channel, + add_downsample=not is_final_block, + resnet_eps=1e-6, + downsample_padding=0, + # Note: Don't know why set it as 0 + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + temporal_down=is_temporal_down_block, + spatial_down=True, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.down_blocks.append(down_block) + + def zero_module(module): + # Zero out the parameters of a module and return it. + for p in module.parameters(): + p.detach().zero_() + return module + + self.conv_extra_cond.append( + zero_module( + nn.Conv3d(extra_cond_dim, output_channel, kernel_size=1, stride=1, padding=0) + ) + if self.extra_cond_dim is not None and self.extra_cond_dim > 0 + else None + ) + + # mid + self.mid_block = UNetMidBlock3D( + in_channels=block_out_channels[-1], + resnet_eps=1e-6, + resnet_act_fn=act_fn, + output_scale_factor=1, + resnet_time_scale_shift="default", + attention_head_dim=block_out_channels[-1], + resnet_groups=norm_num_groups, + temb_channels=None, + add_attention=mid_block_add_attention, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + # out + self.conv_norm_out = nn.GroupNorm( + num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6 + ) + self.conv_act = nn.SiLU() + + conv_out_channels = 2 * out_channels if double_z else out_channels + self.conv_out = init_causal_conv3d( + block_out_channels[-1], conv_out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + self.gradient_checkpointing = gradient_checkpoint + + def forward( + self, + sample: torch.FloatTensor, + extra_cond=None, + memory_state: MemoryState = MemoryState.DISABLED, + preserve_vram: bool = False + ) -> torch.FloatTensor: + r"""The forward method of the `Encoder` class.""" + sample = self.conv_in(sample, memory_state=memory_state, preserve_vram=preserve_vram) + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + + # down + # [Override] add extra block and extra cond + for down_block, extra_block in zip(self.down_blocks, self.conv_extra_cond): + sample = torch.utils.checkpoint.checkpoint( + create_custom_forward(down_block), sample, memory_state, use_reentrant=False + ) + if extra_block is not None: + sample = sample + safe_interpolate_operation(extra_block(extra_cond), size=sample.shape[2:]) + + # middle + sample = self.mid_block(sample, memory_state=memory_state) + + # sample = torch.utils.checkpoint.checkpoint( + # create_custom_forward(self.mid_block), sample, use_reentrant=False + # ) + + else: + # down + # [Override] add extra block and extra cond + for down_block, extra_block in zip(self.down_blocks, self.conv_extra_cond): + sample = down_block(sample, memory_state=memory_state, preserve_vram=preserve_vram) + if extra_block is not None: + sample = sample + safe_interpolate_operation(extra_block(extra_cond), size=sample.shape[2:]) + + # middle + sample = self.mid_block(sample, memory_state=memory_state) + + # post-process + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state) + + return sample + + +class Decoder3D(nn.Module): + r""" + The `Decoder` layer of a variational autoencoder that + decodes its latent representation into an output sample. + + Args: + in_channels (`int`, *optional*, defaults to 3): + The number of input channels. + out_channels (`int`, *optional*, defaults to 3): + The number of output channels. + up_block_types (`Tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`): + The types of up blocks to use. + See `~diffusers.models.unet_2d_blocks.get_up_block` for available options. + block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`): + The number of output channels for each block. + layers_per_block (`int`, *optional*, defaults to 2): + The number of layers per block. + norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups for normalization. + act_fn (`str`, *optional*, defaults to `"silu"`): + The activation function to use. + See `~diffusers.models.activations.get_activation` for available options. + norm_type (`str`, *optional*, defaults to `"group"`): + The normalization type to use. Can be either `"group"` or `"spatial"`. + """ + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + up_block_types: Tuple[str, ...] = ("UpDecoderBlock3D",), + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + norm_num_groups: int = 32, + act_fn: str = "silu", + norm_type: str = "group", # group, spatial + mid_block_add_attention=True, + # [Override] add temporal up block + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up_num: int = 2, + slicing_up_num: int = 0, + gradient_checkpoint: bool = False, + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_up_num = temporal_up_num + + self.conv_in = init_causal_conv3d( + in_channels, + block_out_channels[-1], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.mid_block = None + self.up_blocks = nn.ModuleList([]) + + temb_channels = in_channels if norm_type == "spatial" else None + + # mid + self.mid_block = UNetMidBlock3D( + in_channels=block_out_channels[-1], + resnet_eps=1e-6, + resnet_act_fn=act_fn, + output_scale_factor=1, + resnet_time_scale_shift="default" if norm_type == "group" else norm_type, + attention_head_dim=block_out_channels[-1], + resnet_groups=norm_num_groups, + temb_channels=temb_channels, + add_attention=mid_block_add_attention, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + # up + reversed_block_out_channels = list(reversed(block_out_channels)) + output_channel = reversed_block_out_channels[0] + #print(f"slicing_up_num: {slicing_up_num}") + for i, up_block_type in enumerate(up_block_types): + prev_output_channel = output_channel + output_channel = reversed_block_out_channels[i] + + is_final_block = i == len(block_out_channels) - 1 + is_temporal_up_block = i < self.temporal_up_num + is_slicing_up_block = i >= len(block_out_channels) - slicing_up_num + # Note: Keep symmetric + + assert up_block_type == "UpDecoderBlock3D" + up_block = UpDecoderBlock3D( + num_layers=self.layers_per_block + 1, + in_channels=prev_output_channel, + out_channels=output_channel, + add_upsample=not is_final_block, + resnet_eps=1e-6, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + resnet_time_scale_shift=norm_type, + temb_channels=temb_channels, + temporal_up=is_temporal_up_block, + slicing=is_slicing_up_block, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.up_blocks.append(up_block) + prev_output_channel = output_channel + + # out + if norm_type == "spatial": + self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels) + else: + self.conv_norm_out = nn.GroupNorm( + num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6 + ) + self.conv_act = nn.SiLU() + self.conv_out = init_causal_conv3d( + block_out_channels[0], out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + self.gradient_checkpointing = gradient_checkpoint + + # Note: Just copy from Decoder. + def forward( + self, + sample: torch.FloatTensor, + latent_embeds: Optional[torch.FloatTensor] = None, + memory_state: MemoryState = MemoryState.DISABLED, + preserve_vram: bool = False + ) -> torch.FloatTensor: + r"""The forward method of the `Decoder` class.""" + + sample = self.conv_in(sample, memory_state=memory_state) + + #upscale_dtype = next(iter(self.up_blocks.parameters())).dtype + # ADD BY NUMZ + upscale_dtype = sample.dtype + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + + if is_torch_version(">=", "1.11.0"): + sample = self.mid_block(sample, latent_embeds, memory_state=memory_state) + sample = sample.to(upscale_dtype) + + # up + for up_block in self.up_blocks: + sample = torch.utils.checkpoint.checkpoint( + create_custom_forward(up_block), + sample, + latent_embeds, + memory_state, + use_reentrant=False, + ) + else: + # middle + sample = self.mid_block(sample, latent_embeds, memory_state=memory_state) + sample = sample.to(upscale_dtype) + + # up + for up_block in self.up_blocks: + sample = torch.utils.checkpoint.checkpoint( + create_custom_forward(up_block), sample, latent_embeds, memory_state + ) + else: + # middle + sample = self.mid_block(sample, latent_embeds, memory_state=memory_state) + sample = sample.to(upscale_dtype) + + # up + for up_block in self.up_blocks: + sample = up_block(sample, latent_embeds, memory_state=memory_state, preserve_vram=preserve_vram) + + # post-process + sample = causal_norm_wrapper(self.conv_norm_out, sample, preserve_vram=preserve_vram) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state) + + return sample + + +class AutoencoderKL(diffusers.AutoencoderKL): + """ + We simply inherit the model code from diffusers + """ + + def __init__(self, attention: bool = True, *args, **kwargs): + super().__init__(*args, **kwargs) + + # A hacky way to remove attention. + if not attention: + self.encoder.mid_block.attentions = torch.nn.ModuleList([None]) + self.decoder.mid_block.attentions = torch.nn.ModuleList([None]) + + def load_state_dict(self, state_dict, strict=True): + # Newer version of diffusers changed the model keys, + # causing incompatibility with old checkpoints. + # They provided a method for conversion. We call conversion before loading state_dict. + convert_deprecated_attention_blocks = getattr( + self, "_convert_deprecated_attention_blocks", None + ) + if callable(convert_deprecated_attention_blocks): + convert_deprecated_attention_blocks(state_dict) + return super().load_state_dict(state_dict, strict) + + +class VideoAutoencoderKL(diffusers.AutoencoderKL): + """ + We simply inherit the model code from diffusers + """ + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + down_block_types: Tuple[str] = ("DownEncoderBlock3D",), + up_block_types: Tuple[str] = ("UpDecoderBlock3D",), + block_out_channels: Tuple[int] = (64,), + layers_per_block: int = 1, + act_fn: str = "silu", + latent_channels: int = 4, + norm_num_groups: int = 32, + sample_size: int = 32, + scaling_factor: float = 0.18215, + force_upcast: float = True, + attention: bool = True, + temporal_scale_num: int = 2, + slicing_up_num: int = 0, + gradient_checkpoint: bool = False, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "full", + slicing_sample_min_size: int = 32, + use_quant_conv: bool = True, + use_post_quant_conv: bool = True, + *args, + **kwargs, + ): + extra_cond_dim = kwargs.pop("extra_cond_dim") if "extra_cond_dim" in kwargs else None + self.slicing_sample_min_size = slicing_sample_min_size + self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num) + + + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + # [Override] make sure it can be normally initialized + down_block_types=tuple( + [down_block_type.replace("3D", "2D") for down_block_type in down_block_types] + ), + up_block_types=tuple( + [up_block_type.replace("3D", "2D") for up_block_type in up_block_types] + ), + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + act_fn=act_fn, + latent_channels=latent_channels, + norm_num_groups=norm_num_groups, + sample_size=sample_size, + scaling_factor=scaling_factor, + force_upcast=force_upcast, + *args, + **kwargs, + ) + + # pass init params to Encoder + self.encoder = Encoder3D( + in_channels=in_channels, + out_channels=latent_channels, + down_block_types=down_block_types, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + act_fn=act_fn, + norm_num_groups=norm_num_groups, + double_z=True, + extra_cond_dim=extra_cond_dim, + # [Override] add temporal_down_num parameter + temporal_down_num=temporal_scale_num, + gradient_checkpoint=gradient_checkpoint, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + # pass init params to Decoder + self.decoder = Decoder3D( + in_channels=latent_channels, + out_channels=out_channels, + up_block_types=up_block_types, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + norm_num_groups=norm_num_groups, + act_fn=act_fn, + # [Override] add temporal_up_num parameter + temporal_up_num=temporal_scale_num, + slicing_up_num=slicing_up_num, + gradient_checkpoint=gradient_checkpoint, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.quant_conv = ( + init_causal_conv3d( + in_channels=2 * latent_channels, + out_channels=2 * latent_channels, + kernel_size=1, + inflation_mode=inflation_mode, + ) + if use_quant_conv + else None + ) + self.post_quant_conv = ( + init_causal_conv3d( + in_channels=latent_channels, + out_channels=latent_channels, + kernel_size=1, + inflation_mode=inflation_mode, + ) + if use_post_quant_conv + else None + ) + + # A hacky way to remove attention. + if not attention: + self.encoder.mid_block.attentions = torch.nn.ModuleList([None]) + self.decoder.mid_block.attentions = torch.nn.ModuleList([None]) + + @apply_forward_hook + def encode(self, x: torch.FloatTensor, preserve_vram: bool = False, return_dict: bool = True) -> AutoencoderKLOutput: + h = self.slicing_encode(x, preserve_vram) + posterior = DiagonalGaussianDistribution(h) + + if not return_dict: + return (posterior,) + + return AutoencoderKLOutput(latent_dist=posterior) + + @apply_forward_hook + def decode( + self, z: torch.Tensor, preserve_vram: bool = False, return_dict: bool = True + ) -> Union[DecoderOutput, torch.Tensor]: + decoded = self.slicing_decode(z, preserve_vram) + + if not return_dict: + return (decoded,) + + return DecoderOutput(sample=decoded) + + def _encode( + self, x: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED, preserve_vram: bool = False + ) -> torch.Tensor: + _x = x.to(self.device) + _x = causal_conv_slice_inputs(_x, self.slicing_sample_min_size, memory_state=memory_state) + h = self.encoder(_x, memory_state=memory_state, preserve_vram=preserve_vram) + if self.quant_conv is not None: + output = self.quant_conv(h, memory_state=memory_state) + else: + output = h + output = causal_conv_gather_outputs(output) + return output.to(x.device) + + def _decode( + self, z: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED, preserve_vram: bool = False + ) -> torch.Tensor: + _z = z.to(self.device) + _z = causal_conv_slice_inputs(_z, self.slicing_latent_min_size, memory_state=memory_state) + if self.post_quant_conv is not None: + _z = self.post_quant_conv(_z, memory_state=memory_state) + output = self.decoder(_z, memory_state=memory_state, preserve_vram=preserve_vram) + output = causal_conv_gather_outputs(output) + return output.to(z.device) + + def slicing_encode(self, x: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor: + sp_size = get_sequence_parallel_world_size() + if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size * sp_size: + x_slices = x[:, :, 1:].split(split_size=self.slicing_sample_min_size * sp_size, dim=2) + encoded_slices = [ + self._encode( + torch.cat((x[:, :, :1], x_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + preserve_vram=preserve_vram + ) + ] + for x_idx in range(1, len(x_slices)): + encoded_slices.append( + self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE, preserve_vram=preserve_vram) + ) + return torch.cat(encoded_slices, dim=2) + else: + return self._encode(x, preserve_vram=preserve_vram) + + def slicing_decode(self, z: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor: + sp_size = get_sequence_parallel_world_size() + if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size * sp_size: + z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size * sp_size, dim=2) + decoded_slices = [ + self._decode( + torch.cat((z[:, :, :1], z_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + preserve_vram=preserve_vram + ) + ] + for z_idx in range(1, len(z_slices)): + decoded_slices.append( + self._decode(z_slices[z_idx], memory_state=MemoryState.DISABLED, preserve_vram=preserve_vram) + ) + return torch.cat(decoded_slices, dim=2) + else: + return self._decode(z, preserve_vram=preserve_vram) + + def tiled_encode(self, x: torch.Tensor, **kwargs) -> torch.Tensor: + raise NotImplementedError + + def tiled_decode(self, z: torch.Tensor, **kwargs) -> torch.Tensor: + raise NotImplementedError + + def forward( + self, x: torch.FloatTensor, mode: Literal["encode", "decode", "all"] = "all", **kwargs + ): + # x: [b c t h w] + if mode == "encode": + h = self.encode(x) + return h.latent_dist + elif mode == "decode": + h = self.decode(x) + return h.sample + else: + h = self.encode(x) + h = self.decode(h.latent_dist.mode()) + return h.sample + + def load_state_dict(self, state_dict, strict=False): + # Newer version of diffusers changed the model keys, + # causing incompatibility with old checkpoints. + # They provided a method for conversion. + # We call conversion before loading state_dict. + convert_deprecated_attention_blocks = getattr( + self, "_convert_deprecated_attention_blocks", None + ) + if callable(convert_deprecated_attention_blocks): + convert_deprecated_attention_blocks(state_dict) + return super().load_state_dict(state_dict, strict) + + +class VideoAutoencoderKLWrapper(VideoAutoencoderKL): + def __init__( + self, + *args, + spatial_downsample_factor: int, + temporal_downsample_factor: int, + freeze_encoder: bool, + **kwargs, + ): + self.spatial_downsample_factor = spatial_downsample_factor + self.temporal_downsample_factor = temporal_downsample_factor + self.freeze_encoder = freeze_encoder + super().__init__(*args, **kwargs) + + def forward(self, x: torch.FloatTensor) -> CausalAutoencoderOutput: + with torch.no_grad() if self.freeze_encoder else nullcontext(): + z, p = self.encode(x) + x = self.decode(z).sample + return CausalAutoencoderOutput(x, z, p) + + def encode(self, x: torch.FloatTensor, preserve_vram: bool = False) -> CausalEncoderOutput: + if x.ndim == 4: + x = x.unsqueeze(2) + p = super().encode(x, preserve_vram).latent_dist + z = p.sample().squeeze(2) + return CausalEncoderOutput(z, p) + + def decode(self, z: torch.FloatTensor, preserve_vram: bool = False) -> CausalDecoderOutput: + if z.ndim == 4: + z = z.unsqueeze(2) + x = super().decode(z, preserve_vram).sample.squeeze(2) + return CausalDecoderOutput(x) + + def preprocess(self, x: torch.Tensor): + # x should in [B, C, T, H, W], [B, C, H, W] + assert x.ndim == 4 or x.size(2) % 4 == 1 + return x + + def postprocess(self, x: torch.Tensor): + # x should in [B, C, T, H, W], [B, C, H, W] + return x + + def set_causal_slicing( + self, + *, + split_size: Optional[int], + memory_device: _memory_device_t, + ): + assert ( + split_size is None or memory_device is not None + ), "if split_size is set, memory_device must not be None." + if split_size is not None: + self.enable_slicing() + self.slicing_sample_min_size = split_size + self.slicing_latent_min_size = split_size // self.temporal_downsample_factor + else: + self.disable_slicing() + for module in self.modules(): + if isinstance(module, InflatedCausalConv3d): + module.set_memory_device(memory_device) + + def set_memory_limit(self, conv_max_mem: Optional[float], norm_max_mem: Optional[float]): + set_norm_limit(norm_max_mem) + for m in self.modules(): + if isinstance(m, InflatedCausalConv3d): + m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf")) diff --git a/src/models/video_vae_v3_mine_bad/modules/causal_inflation_lib.py b/src/models/video_vae_v3_mine_bad/modules/causal_inflation_lib.py new file mode 100644 index 0000000..df48682 --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/causal_inflation_lib.py @@ -0,0 +1,464 @@ +# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates +# // +# // Licensed under the Apache License, Version 2.0 (the "License"); +# // you may not use this file except in compliance with the License. +# // You may obtain a copy of the License at +# // +# // http://www.apache.org/licenses/LICENSE-2.0 +# // +# // Unless required by applicable law or agreed to in writing, software +# // distributed under the License is distributed on an "AS IS" BASIS, +# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# // See the License for the specific language governing permissions and +# // limitations under the License. + +import math +from contextlib import contextmanager +import time +from typing import List, Optional, Union +import torch +import torch.nn.functional as F +from diffusers.models.normalization import RMSNorm +from einops import rearrange +from torch import Tensor, nn +from torch.nn import Conv3d + +from .context_parallel_lib import cache_send_recv, get_cache_size +from .global_config import get_norm_limit +from .types import MemoryState, _inflation_mode_t, _memory_device_t +from ....common.half_precision_fixes import safe_pad_operation + +# Single GPU inference - no distributed processing needed +#print("Warning: Using single GPU inference mode - distributed features disabled in causal_inflation_lib") + +# Mock distributed functions for single GPU inference +def get_sequence_parallel_group(): + return None + +def get_sequence_parallel_rank(): + return 0 + +def get_sequence_parallel_world_size(): + return 1 + +def get_next_sequence_parallel_rank(): + return 0 + +def get_prev_sequence_parallel_rank(): + return 0 + + +@contextmanager +def ignore_padding(model): + orig_padding = model.padding + model.padding = (0, 0, 0) + try: + yield + finally: + model.padding = orig_padding + + +class InflatedCausalConv3d(Conv3d): + def __init__( + self, + *args, + inflation_mode: _inflation_mode_t, + memory_device: _memory_device_t = "same", + **kwargs, + ): + self.inflation_mode = inflation_mode + self.memory = None + super().__init__(*args, **kwargs) + self.temporal_padding = self.padding[0] + self.memory_device = memory_device + self.padding = (0, *self.padding[1:]) # Remove temporal pad to keep causal. + self.memory_limit = float("inf") + + def set_memory_limit(self, value: float): + self.memory_limit = value + + def set_memory_device(self, memory_device: _memory_device_t): + self.memory_device = memory_device + + def memory_limit_conv( + self, + x, + *, + split_dim=3, + padding=(0, 0, 0, 0, 0, 0), + prev_cache=None, + preserve_vram = False, + ): + # Compatible with no limit. + if math.isinf(self.memory_limit): + if prev_cache is not None: + x = torch.cat([prev_cache, x], dim=split_dim - 1) + return super().forward(x) + + # Compute tensor shape after concat & padding. + shape = torch.tensor(x.size()) + if prev_cache is not None: + shape[split_dim - 1] += prev_cache.size(split_dim - 1) + shape[-3:] += torch.tensor(padding).view(3, 2).sum(-1).flip(0) + memory_occupy = shape.prod() * x.element_size() / 1024**3 # GiB + if memory_occupy < self.memory_limit or split_dim == x.ndim: + if prev_cache is not None: + x = torch.cat([prev_cache, x], dim=split_dim - 1) + x = safe_pad_operation(x, padding, mode='constant', value=0.0) + with ignore_padding(self): + return super().forward(x) + + # Exceed memory limit, splitting tensor + + # Split input (& prev_cache). + num_splits = math.ceil(memory_occupy / self.memory_limit) + size_per_split = x.size(split_dim) // num_splits + split_sizes = [size_per_split] * (num_splits - 1) + split_sizes += [x.size(split_dim) - sum(split_sizes)] + + x = list(x.split(split_sizes, dim=split_dim)) + if prev_cache is not None: + prev_cache = list(prev_cache.split(split_sizes, dim=split_dim)) + if preserve_vram: + torch.cuda.empty_cache() + #print("empty cache 0") + # Loop Fwd. + cache = None + for idx in range(len(x)): + # Concat prev cache from last dim + if prev_cache is not None: + x[idx] = torch.cat([prev_cache[idx], x[idx]], dim=split_dim - 1) + + # Get padding pattern. + lpad_dim = (x[idx].ndim - split_dim - 1) * 2 + rpad_dim = lpad_dim + 1 + padding = list(padding) + padding[lpad_dim] = self.padding[split_dim - 2] if idx == 0 else 0 + padding[rpad_dim] = self.padding[split_dim - 2] if idx == len(x) - 1 else 0 + pad_len = padding[lpad_dim] + padding[rpad_dim] + padding = tuple(padding) + + # Prepare cache for next slice (this dim). + next_cache = None + cache_len = cache.size(split_dim) if cache is not None else 0 + next_catch_size = get_cache_size( + conv_module=self, + input_len=x[idx].size(split_dim) + cache_len, + pad_len=pad_len, + dim=split_dim - 2, + ) + if next_catch_size != 0: + assert next_catch_size <= x[idx].size(split_dim) + next_cache = ( + x[idx].transpose(0, split_dim)[-next_catch_size:].transpose(0, split_dim) + ) + + # Recursive. + x[idx] = self.memory_limit_conv( + x[idx], + split_dim=split_dim + 1, + padding=padding, + prev_cache=cache, + preserve_vram=preserve_vram + ) + + # Update cache. + cache = next_cache + # ADD BY NUMZ + if preserve_vram: + torch.cuda.empty_cache() + #print("empty cache 1") + #time.sleep(2) + try: + output = torch.cat(x, split_dim) + except Exception as e: + print("OOM second chance") + torch.cuda.empty_cache() + time.sleep(2) + output = torch.cat(x, split_dim) + return output + + def forward( + self, + input: Union[Tensor, List[Tensor]], + memory_state: MemoryState = MemoryState.UNSET, + preserve_vram: bool = False, + ) -> Tensor: + assert memory_state != MemoryState.UNSET + if memory_state != MemoryState.ACTIVE: + self.memory = None + if ( + math.isinf(self.memory_limit) + and torch.is_tensor(input) + and get_sequence_parallel_group() is None + ): + return self.basic_forward(input, memory_state) + return self.slicing_forward(input, memory_state, preserve_vram) + + def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET): + mem_size = self.stride[0] - self.kernel_size[0] + if (self.memory is not None) and (memory_state == MemoryState.ACTIVE): + input = extend_head(input, memory=self.memory, times=-1) + else: + input = extend_head(input, times=self.temporal_padding * 2) + memory = ( + input[:, :, mem_size:].detach() + if (mem_size != 0 and memory_state != MemoryState.DISABLED) + else None + ) + if ( + memory_state != MemoryState.DISABLED + and not self.training + and (self.memory_device is not None) + ): + self.memory = memory + if self.memory_device == "cpu" and self.memory is not None: + self.memory = self.memory.to("cpu") + return super().forward(input) + + def slicing_forward( + self, + input: Union[Tensor, List[Tensor]], + memory_state: MemoryState = MemoryState.UNSET, + preserve_vram: bool = False, + ) -> Tensor: + squeeze_out = False + if torch.is_tensor(input): + input = [input] + squeeze_out = True + + cache_size = self.kernel_size[0] - self.stride[0] + cache = cache_send_recv( + input, cache_size=cache_size, memory=self.memory, times=self.temporal_padding * 2 + ) + + # Single GPU inference - simplified memory management + if ( + memory_state in [MemoryState.INITIALIZING, MemoryState.ACTIVE] # use_slicing + and not self.training + and (self.memory_device is not None) + and cache_size != 0 + ): + if cache_size > input[-1].size(2) and cache is not None and len(input) == 1: + input[0] = torch.cat([cache, input[0]], dim=2) + cache = None + if cache_size <= input[-1].size(2): + self.memory = input[-1][:, :, -cache_size:].detach().contiguous() + if self.memory_device == "cpu" and self.memory is not None: + self.memory = self.memory.to("cpu") + + padding = tuple(x for x in reversed(self.padding) for _ in range(2)) + for i in range(len(input)): + # Prepare cache for next input slice. + next_cache = None + cache_size = 0 + if i < len(input) - 1: + cache_len = cache.size(2) if cache is not None else 0 + cache_size = get_cache_size(self, input[i].size(2) + cache_len, pad_len=0) + if cache_size != 0: + if cache_size > input[i].size(2) and cache is not None: + input[i] = torch.cat([cache, input[i]], dim=2) + cache = None + assert cache_size <= input[i].size(2), f"{cache_size} > {input[i].size(2)}" + next_cache = input[i][:, :, -cache_size:] + + # Conv forward for this input slice. + input[i] = self.memory_limit_conv( + input[i], + padding=padding, + prev_cache=cache, + preserve_vram=preserve_vram + ) + + # Update cache. + cache = next_cache + + return input[0] if squeeze_out else input + + def tflops(self, args, kwargs, output) -> float: + if torch.is_tensor(output): + output_numel = output.numel() + elif isinstance(output, list): + output_numel = sum(o.numel() for o in output) + else: + raise NotImplementedError + return (2 * math.prod(self.kernel_size) * self.in_channels * (output_numel / 1e6)) / 1e6 + + def _load_from_state_dict( + self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs + ): + if self.inflation_mode != "none": + state_dict = modify_state_dict( + self, + state_dict, + prefix, + inflate_weight_fn=inflate_weight, + inflate_bias_fn=inflate_bias, + ) + super()._load_from_state_dict( + state_dict, + prefix, + local_metadata, + (strict and self.inflation_mode == "none"), + missing_keys, + unexpected_keys, + error_msgs, + ) + + +def init_causal_conv3d( + *args, + inflation_mode: _inflation_mode_t, + **kwargs, +): + """ + Initialize a Causal-3D convolution layer. + Parameters: + inflation_mode: Listed as below. It's compatible with all the 3D-VAE checkpoints we have. + - none: No inflation will be conducted. + The loading logic of state dict will fall back to default. + - tail / replicate: Refer to the definition of `InflatedCausalConv3d`. + """ + return InflatedCausalConv3d(*args, inflation_mode=inflation_mode, **kwargs) + + +def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor: + input_dtype = x.dtype + if isinstance(norm_layer, (nn.LayerNorm, RMSNorm)): + if x.ndim == 4: + x = rearrange(x, "b c h w -> b h w c") + x = norm_layer(x) + x = rearrange(x, "b h w c -> b c h w") + return x.to(input_dtype) + if x.ndim == 5: + x = rearrange(x, "b c t h w -> b t h w c") + x = norm_layer(x) + x = rearrange(x, "b t h w c -> b c t h w") + return x.to(input_dtype) + if isinstance(norm_layer, (nn.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)): + if x.ndim <= 4: + return norm_layer(x).to(input_dtype) + if x.ndim == 5: + t = x.size(2) + x = rearrange(x, "b c t h w -> (b t) c h w") + memory_occupy = x.numel() * x.element_size() / 1024**3 + if isinstance(norm_layer, nn.GroupNorm) and memory_occupy > get_norm_limit(): + num_chunks = min(4 if x.element_size() == 2 else 2, norm_layer.num_groups) + assert norm_layer.num_groups % num_chunks == 0 + num_groups_per_chunk = norm_layer.num_groups // num_chunks + + x = list(x.chunk(num_chunks, dim=1)) + weights = norm_layer.weight.chunk(num_chunks, dim=0) + biases = norm_layer.bias.chunk(num_chunks, dim=0) + for i, (w, b) in enumerate(zip(weights, biases)): + try: + x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps) + except Exception as e: + print("OOM Second Chance : Group Norm") + torch.cuda.empty_cache() + time.sleep(2) + x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps) + x[i] = x[i].to(input_dtype) + # ADD BY NUMZ + if preserve_vram: + torch.cuda.empty_cache() + x = torch.cat(x, dim=1) + else: + x = norm_layer(x) + x = rearrange(x, "(b t) c h w -> b c t h w", t=t) + return x.to(input_dtype) + raise NotImplementedError + + +def remove_head(tensor: Tensor, times: int = 1) -> Tensor: + """ + Remove duplicated first frame features in the up-sampling process. + """ + # Single GPU inference - always process + if times == 0: + return tensor + return torch.cat(tensors=(tensor[:, :, :1], tensor[:, :, times + 1 :]), dim=2) + + +def extend_head(tensor: Tensor, times: int = 2, memory: Optional[Tensor] = None) -> Tensor: + """ + When memory is None: + - Duplicate first frame features in the down-sampling process. + When memory is not None: + - Concatenate memory features with the input features to keep temporal consistency. + """ + if memory is not None: + return torch.cat((memory.to(tensor), tensor), dim=2) + assert times >= 0, "Invalid input for function 'extend_head'!" + if times == 0: + return tensor + else: + tile_repeat = [1] * tensor.ndim + tile_repeat[2] = times + return torch.cat(tensors=(torch.tile(tensor[:, :, :1], tile_repeat), tensor), dim=2) + + +def inflate_weight(weight_2d: torch.Tensor, weight_3d: torch.Tensor, inflation_mode: str): + """ + Inflate a 2D convolution weight matrix to a 3D one. + Parameters: + weight_2d: The weight matrix of 2D conv to be inflated. + weight_3d: The weight matrix of 3D conv to be initialized. + inflation_mode: the mode of inflation + """ + assert inflation_mode in ["tail", "replicate"] + assert weight_3d.shape[:2] == weight_2d.shape[:2] + with torch.no_grad(): + if inflation_mode == "replicate": + depth = weight_3d.size(2) + weight_3d.copy_(weight_2d.unsqueeze(2).repeat(1, 1, depth, 1, 1) / depth) + else: + weight_3d.fill_(0.0) + weight_3d[:, :, -1].copy_(weight_2d) + return weight_3d + + +def inflate_bias(bias_2d: torch.Tensor, bias_3d: torch.Tensor, inflation_mode: str): + """ + Inflate a 2D convolution bias tensor to a 3D one + Parameters: + bias_2d: The bias tensor of 2D conv to be inflated. + bias_3d: The bias tensor of 3D conv to be initialized. + inflation_mode: Placeholder to align `inflate_weight`. + """ + assert bias_3d.shape == bias_2d.shape + with torch.no_grad(): + bias_3d.copy_(bias_2d) + return bias_3d + + +def modify_state_dict(layer, state_dict, prefix, inflate_weight_fn, inflate_bias_fn): + """ + the main function to inflated 2D parameters to 3D. + """ + weight_name = prefix + "weight" + bias_name = prefix + "bias" + if weight_name in state_dict: + weight_2d = state_dict[weight_name] + if weight_2d.dim() == 4: + # Assuming the 2D weights are 4D tensors (out_channels, in_channels, h, w) + weight_3d = inflate_weight_fn( + weight_2d=weight_2d, + weight_3d=layer.weight, + inflation_mode=layer.inflation_mode, + ) + state_dict[weight_name] = weight_3d + else: + return state_dict + # It's a 3d state dict, should not do inflation on both bias and weight. + if bias_name in state_dict: + bias_2d = state_dict[bias_name] + if bias_2d.dim() == 1: + # Assuming the 2D biases are 1D tensors (out_channels,) + bias_3d = inflate_bias_fn( + bias_2d=bias_2d, + bias_3d=layer.bias, + inflation_mode=layer.inflation_mode, + ) + state_dict[bias_name] = bias_3d + return state_dict diff --git a/src/models/video_vae_v3_mine_bad/modules/context_parallel_lib.py b/src/models/video_vae_v3_mine_bad/modules/context_parallel_lib.py new file mode 100644 index 0000000..663525e --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/context_parallel_lib.py @@ -0,0 +1,67 @@ +# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates +# // +# // Licensed under the Apache License, Version 2.0 (the "License"); +# // you may not use this file except in compliance with the License. +# // You may obtain a copy of the License at +# // +# // http://www.apache.org/licenses/LICENSE-2.0 +# // +# // Unless required by applicable law or agreed to in writing, software +# // distributed under the License is distributed on an "AS IS" BASIS, +# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# // See the License for the specific language governing permissions and +# // limitations under the License. + +from typing import List +import torch +import torch.nn.functional as F +from torch import Tensor + +from .types import MemoryState + +# Single GPU inference - no distributed processing needed +# print("Warning: Using single GPU inference mode - distributed features disabled") + + +def causal_conv_slice_inputs(x, split_size, memory_state): + # Single GPU inference - no slicing needed, return full tensor + return x + + +def causal_conv_gather_outputs(x): + # Single GPU inference - no gathering needed, return tensor as is + return x + + +def get_output_len(conv_module, input_len, pad_len, dim=0): + dilated_kernerl_size = conv_module.dilation[dim] * (conv_module.kernel_size[dim] - 1) + 1 + output_len = (input_len + pad_len - dilated_kernerl_size) // conv_module.stride[dim] + 1 + return output_len + + +def get_cache_size(conv_module, input_len, pad_len, dim=0): + dilated_kernerl_size = conv_module.dilation[dim] * (conv_module.kernel_size[dim] - 1) + 1 + output_len = (input_len + pad_len - dilated_kernerl_size) // conv_module.stride[dim] + 1 + remain_len = ( + input_len + pad_len - ((output_len - 1) * conv_module.stride[dim] + dilated_kernerl_size) + ) + overlap_len = dilated_kernerl_size - conv_module.stride[dim] + cache_len = overlap_len + remain_len # >= 0 + + assert output_len > 0 + return cache_len + + +def cache_send_recv(tensor: List[Tensor], cache_size, times, memory=None): + # Single GPU inference - simplified cache handling + recv_buffer = None + + # Handle memory buffer for single GPU case + if memory is not None: + recv_buffer = memory.to(tensor[0]) + elif times > 0: + tile_repeat = [1] * tensor[0].ndim + tile_repeat[2] = times + recv_buffer = torch.tile(tensor[0][:, :, :1], tile_repeat) + + return recv_buffer diff --git a/src/models/video_vae_v3_mine_bad/modules/global_config.py b/src/models/video_vae_v3_mine_bad/modules/global_config.py new file mode 100644 index 0000000..8631175 --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/global_config.py @@ -0,0 +1,28 @@ +# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates +# // +# // Licensed under the Apache License, Version 2.0 (the "License"); +# // you may not use this file except in compliance with the License. +# // You may obtain a copy of the License at +# // +# // http://www.apache.org/licenses/LICENSE-2.0 +# // +# // Unless required by applicable law or agreed to in writing, software +# // distributed under the License is distributed on an "AS IS" BASIS, +# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# // See the License for the specific language governing permissions and +# // limitations under the License. + +from typing import Optional + +_NORM_LIMIT = float("inf") + + +def get_norm_limit(): + return _NORM_LIMIT + + +def set_norm_limit(value: Optional[float] = None): + global _NORM_LIMIT + if value is None: + value = float("inf") + _NORM_LIMIT = value diff --git a/src/models/video_vae_v3_mine_bad/modules/inflated_layers.py b/src/models/video_vae_v3_mine_bad/modules/inflated_layers.py new file mode 100644 index 0000000..4b8e6df --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/inflated_layers.py @@ -0,0 +1,106 @@ +# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates +# // +# // Licensed under the Apache License, Version 2.0 (the "License"); +# // you may not use this file except in compliance with the License. +# // You may obtain a copy of the License at +# // +# // http://www.apache.org/licenses/LICENSE-2.0 +# // +# // Unless required by applicable law or agreed to in writing, software +# // distributed under the License is distributed on an "AS IS" BASIS, +# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# // See the License for the specific language governing permissions and +# // limitations under the License. + +from functools import partial +from typing import Literal, Optional +from torch import Tensor +from torch.nn import Conv3d + +from .inflated_lib import ( + MemoryState, + extend_head, + inflate_bias, + inflate_weight, + modify_state_dict, +) + +_inflation_mode_t = Literal["none", "tail", "replicate"] +_memory_device_t = Optional[Literal["cpu", "same"]] + + +class InflatedCausalConv3d(Conv3d): + def __init__( + self, + *args, + inflation_mode: _inflation_mode_t, + memory_device: _memory_device_t = "same", + **kwargs, + ): + self.inflation_mode = inflation_mode + self.memory = None + super().__init__(*args, **kwargs) + self.temporal_padding = self.padding[0] + self.memory_device = memory_device + self.padding = (0, *self.padding[1:]) # Remove temporal pad to keep causal. + + def set_memory_device(self, memory_device: _memory_device_t): + self.memory_device = memory_device + + def forward(self, input: Tensor, memory_state: MemoryState = MemoryState.DISABLED) -> Tensor: + mem_size = self.stride[0] - self.kernel_size[0] + if (self.memory is not None) and (memory_state == MemoryState.ACTIVE): + input = extend_head(input, memory=self.memory) + else: + input = extend_head(input, times=self.temporal_padding * 2) + memory = ( + input[:, :, mem_size:].detach() + if (mem_size != 0 and memory_state != MemoryState.DISABLED) + else None + ) + if ( + memory_state != MemoryState.DISABLED + and not self.training + and (self.memory_device is not None) + ): + self.memory = memory + if self.memory_device == "cpu" and self.memory is not None: + self.memory = self.memory.to("cpu") + return super().forward(input) + + def _load_from_state_dict( + self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs + ): + if self.inflation_mode != "none": + state_dict = modify_state_dict( + self, + state_dict, + prefix, + inflate_weight_fn=partial(inflate_weight, position="tail"), + inflate_bias_fn=partial(inflate_bias, position="tail"), + ) + super()._load_from_state_dict( + state_dict, + prefix, + local_metadata, + (strict and self.inflation_mode == "none"), + missing_keys, + unexpected_keys, + error_msgs, + ) + + +def init_causal_conv3d( + *args, + inflation_mode: _inflation_mode_t, + **kwargs, +): + """ + Initialize a Causal-3D convolution layer. + Parameters: + inflation_mode: Listed as below. It's compatible with all the 3D-VAE checkpoints we have. + - none: No inflation will be conducted. + The loading logic of state dict will fall back to default. + - tail / replicate: Refer to the definition of `InflatedCausalConv3d`. + """ + return InflatedCausalConv3d(*args, inflation_mode=inflation_mode, **kwargs) diff --git a/src/models/video_vae_v3_mine_bad/modules/inflated_lib.py b/src/models/video_vae_v3_mine_bad/modules/inflated_lib.py new file mode 100644 index 0000000..486c63c --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/inflated_lib.py @@ -0,0 +1,156 @@ +# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates +# // +# // Licensed under the Apache License, Version 2.0 (the "License"); +# // you may not use this file except in compliance with the License. +# // You may obtain a copy of the License at +# // +# // http://www.apache.org/licenses/LICENSE-2.0 +# // +# // Unless required by applicable law or agreed to in writing, software +# // distributed under the License is distributed on an "AS IS" BASIS, +# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# // See the License for the specific language governing permissions and +# // limitations under the License. + +from enum import Enum +from typing import Optional +import numpy as np +import torch +from diffusers.models.normalization import RMSNorm +from einops import rearrange +from torch import Tensor, nn + +from ....common.logger import get_logger + +logger = get_logger(__name__) + + +class MemoryState(Enum): + """ + State[Disabled]: No memory bank will be enabled. + State[Initializing]: The model is handling the first clip, + need to reset / initialize the memory bank. + State[Active]: There has been some data in the memory bank. + """ + + DISABLED = 0 + INITIALIZING = 1 + ACTIVE = 2 + + +def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor: + if isinstance(norm_layer, (nn.LayerNorm, RMSNorm)): + if x.ndim == 4: + x = rearrange(x, "b c h w -> b h w c") + x = norm_layer(x) + x = rearrange(x, "b h w c -> b c h w") + return x + if x.ndim == 5: + x = rearrange(x, "b c t h w -> b t h w c") + x = norm_layer(x) + x = rearrange(x, "b t h w c -> b c t h w") + return x + if isinstance(norm_layer, (nn.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)): + if x.ndim <= 4: + return norm_layer(x) + if x.ndim == 5: + t = x.size(2) + x = rearrange(x, "b c t h w -> (b t) c h w") + x = norm_layer(x) + x = rearrange(x, "(b t) c h w -> b c t h w", t=t) + return x + raise NotImplementedError + + +def remove_head(tensor: Tensor, times: int = 1) -> Tensor: + """ + Remove duplicated first frame features in the up-sampling process. + """ + if times == 0: + return tensor + return torch.cat(tensors=(tensor[:, :, :1], tensor[:, :, times + 1 :]), dim=2) + + +def extend_head( + tensor: Tensor, times: Optional[int] = 2, memory: Optional[Tensor] = None +) -> Tensor: + """ + When memory is None: + - Duplicate first frame features in the down-sampling process. + When memory is not None: + - Concatenate memory features with the input features to keep temporal consistency. + """ + if times == 0: + return tensor + if memory is not None: + return torch.cat((memory.to(tensor), tensor), dim=2) + else: + tile_repeat = np.ones(tensor.ndim).astype(int) + tile_repeat[2] = times + return torch.cat(tensors=(torch.tile(tensor[:, :, :1], list(tile_repeat)), tensor), dim=2) + + +def inflate_weight(weight_2d: torch.Tensor, weight_3d: torch.Tensor, inflation_mode: str): + """ + Inflate a 2D convolution weight matrix to a 3D one. + Parameters: + weight_2d: The weight matrix of 2D conv to be inflated. + weight_3d: The weight matrix of 3D conv to be initialized. + inflation_mode: the mode of inflation + """ + assert inflation_mode in ["constant", "replicate"] + assert weight_3d.shape[:2] == weight_2d.shape[:2] + with torch.no_grad(): + if inflation_mode == "replicate": + depth = weight_3d.size(2) + weight_3d.copy_(weight_2d.unsqueeze(2).repeat(1, 1, depth, 1, 1) / depth) + else: + weight_3d.fill_(0.0) + weight_3d[:, :, -1].copy_(weight_2d) + return weight_3d + + +def inflate_bias(bias_2d: torch.Tensor, bias_3d: torch.Tensor, inflation_mode: str): + """ + Inflate a 2D convolution bias tensor to a 3D one + Parameters: + bias_2d: The bias tensor of 2D conv to be inflated. + bias_3d: The bias tensor of 3D conv to be initialized. + inflation_mode: Placeholder to align `inflate_weight`. + """ + assert bias_3d.shape == bias_2d.shape + with torch.no_grad(): + bias_3d.copy_(bias_2d) + return bias_3d + + +def modify_state_dict(layer, state_dict, prefix, inflate_weight_fn, inflate_bias_fn): + """ + the main function to inflated 2D parameters to 3D. + """ + weight_name = prefix + "weight" + bias_name = prefix + "bias" + if weight_name in state_dict: + weight_2d = state_dict[weight_name] + if weight_2d.dim() == 4: + # Assuming the 2D weights are 4D tensors (out_channels, in_channels, h, w) + weight_3d = inflate_weight_fn( + weight_2d=weight_2d, + weight_3d=layer.weight, + inflation_mode=layer.inflation_mode, + ) + state_dict[weight_name] = weight_3d + else: + return state_dict + # It's a 3d state dict, should not do inflation on both bias and weight. + if bias_name in state_dict: + bias_2d = state_dict[bias_name] + if bias_2d.dim() == 1: + # Assuming the 2D biases are 1D tensors (out_channels,) + bias_3d = inflate_bias_fn( + bias_2d=bias_2d, + bias_3d=layer.bias, + inflation_mode=layer.inflation_mode, + ) + state_dict[bias_name] = bias_3d + return state_dict diff --git a/src/models/video_vae_v3_mine_bad/modules/types.py b/src/models/video_vae_v3_mine_bad/modules/types.py new file mode 100644 index 0000000..5a030d2 --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/types.py @@ -0,0 +1,76 @@ +# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates +# // +# // Licensed under the Apache License, Version 2.0 (the "License"); +# // you may not use this file except in compliance with the License. +# // You may obtain a copy of the License at +# // +# // http://www.apache.org/licenses/LICENSE-2.0 +# // +# // Unless required by applicable law or agreed to in writing, software +# // distributed under the License is distributed on an "AS IS" BASIS, +# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# // See the License for the specific language governing permissions and +# // limitations under the License. + +from enum import Enum +from typing import Dict, Literal, NamedTuple, Optional +import torch + +_receptive_field_t = Literal["half", "full"] +_inflation_mode_t = Literal["none", "tail", "replicate"] +_memory_device_t = Optional[Literal["cpu", "same"]] +_gradient_checkpointing_t = Optional[Literal["half", "full"]] +_selective_checkpointing_t = Optional[Literal["coarse", "fine"]] + +class DiagonalGaussianDistribution: + def __init__(self, mean: torch.Tensor, logvar: torch.Tensor): + self.mean = mean + self.logvar = torch.clamp(logvar, -30.0, 20.0) + self.std = torch.exp(0.5 * self.logvar) + self.var = torch.exp(self.logvar) + + def mode(self) -> torch.Tensor: + return self.mean + + def sample(self) -> torch.FloatTensor: + return self.mean + self.std * torch.randn_like(self.mean) + + def kl(self) -> torch.Tensor: + return 0.5 * torch.sum( + self.mean**2 + self.var - 1.0 - self.logvar, + dim=list(range(1, self.mean.ndim)), + ) + +class MemoryState(Enum): + """ + State[Disabled]: No memory bank will be enabled. + State[Initializing]: The model is handling the first clip, need to reset the memory bank. + State[Active]: There has been some data in the memory bank. + State[Unset]: Error state, indicating users didn't pass correct memory state in. + """ + + DISABLED = 0 + INITIALIZING = 1 + ACTIVE = 2 + UNSET = 3 + + +class QuantizerOutput(NamedTuple): + latent: torch.Tensor + extra_loss: torch.Tensor + statistics: Dict[str, torch.Tensor] + + +class CausalAutoencoderOutput(NamedTuple): + sample: torch.Tensor + latent: torch.Tensor + posterior: Optional[DiagonalGaussianDistribution] + + +class CausalEncoderOutput(NamedTuple): + latent: torch.Tensor + posterior: Optional[DiagonalGaussianDistribution] + + +class CausalDecoderOutput(NamedTuple): + sample: torch.Tensor diff --git a/src/models/video_vae_v3_mine_bad/modules/video_vae.py b/src/models/video_vae_v3_mine_bad/modules/video_vae.py new file mode 100644 index 0000000..daa1fc0 --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/modules/video_vae.py @@ -0,0 +1,956 @@ +# Copyright (c) 2023 HuggingFace Team +# Copyright (c) 2025 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache License, Version 2.0 (the "License") +# +# This file has been modified by ByteDance Ltd. and/or its affiliates. on 1st June 2025 +# +# Original file was released under Apache License, Version 2.0 (the "License"), with the full license text +# available at http://www.apache.org/licenses/LICENSE-2.0. +# +# This modified file is released under the same license. + +from contextlib import nullcontext +from typing import Optional, Tuple, Literal, Callable, Union + +import torch +import torch.nn as nn +import torch.nn.functional as F +from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution +from einops import rearrange +from ....common.half_precision_fixes import safe_pad_operation + +from ....common.distributed.advanced import get_sequence_parallel_world_size +from ....common.logger import get_logger +from .causal_inflation_lib import ( + InflatedCausalConv3d, + causal_norm_wrapper, + init_causal_conv3d, + remove_head, +) +from .context_parallel_lib import ( + causal_conv_gather_outputs, + causal_conv_slice_inputs, +) +from .global_config import set_norm_limit +from .types import ( + CausalAutoencoderOutput, + CausalDecoderOutput, + CausalEncoderOutput, + MemoryState, + _inflation_mode_t, + _memory_device_t, + _receptive_field_t, + _selective_checkpointing_t, +) + +logger = get_logger(__name__) # pylint: disable=invalid-name + +# Fake func, no checkpointing is required for inference +def gradient_checkpointing(module: Union[Callable, nn.Module], *args, enabled: bool, **kwargs): + return module(*args, **kwargs) + +class ResnetBlock2D(nn.Module): + r""" + A Resnet block. + + Parameters: + in_channels (`int`): The number of channels in the input. + out_channels (`int`, *optional*, default to be `None`): + The number of output channels for the first conv2d layer. + If None, same as `in_channels`. + dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use. + """ + + def __init__( + self, *, in_channels: int, out_channels: Optional[int] = None, dropout: float = 0.0 + ): + super().__init__() + self.in_channels = in_channels + out_channels = in_channels if out_channels is None else out_channels + self.out_channels = out_channels + + self.nonlinearity = nn.SiLU() + + self.norm1 = torch.nn.GroupNorm( + num_groups=32, num_channels=in_channels, eps=1e-6, affine=True + ) + + self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) + + self.norm2 = torch.nn.GroupNorm( + num_groups=32, num_channels=out_channels, eps=1e-6, affine=True + ) + + self.dropout = torch.nn.Dropout(dropout) + self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) + + self.use_in_shortcut = self.in_channels != out_channels + + self.conv_shortcut = None + if self.use_in_shortcut: + self.conv_shortcut = nn.Conv2d( + in_channels, out_channels, kernel_size=1, stride=1, padding=0 + ) + + def forward(self, input_tensor: torch.Tensor) -> torch.Tensor: + hidden = input_tensor + + hidden = self.norm1(hidden) + hidden = self.nonlinearity(hidden) + hidden = self.conv1(hidden) + + hidden = self.norm2(hidden) + hidden = self.nonlinearity(hidden) + hidden = self.dropout(hidden) + hidden = self.conv2(hidden) + + if self.conv_shortcut is not None: + input_tensor = self.conv_shortcut(input_tensor) + + output_tensor = input_tensor + hidden + + return output_tensor + +class Upsample3D(nn.Module): + """A 3D upsampling layer.""" + + def __init__( + self, + channels: int, + inflation_mode: _inflation_mode_t = "tail", + temporal_up: bool = False, + spatial_up: bool = True, + slicing: bool = False, + ): + super().__init__() + self.channels = channels + self.conv = init_causal_conv3d( + self.channels, self.channels, kernel_size=3, padding=1, inflation_mode=inflation_mode + ) + + self.temporal_up = temporal_up + self.spatial_up = spatial_up + self.temporal_ratio = 2 if temporal_up else 1 + self.spatial_ratio = 2 if spatial_up else 1 + self.slicing = slicing + + upscale_ratio = (self.spatial_ratio**2) * self.temporal_ratio + self.upscale_conv = nn.Conv3d( + self.channels, self.channels * upscale_ratio, kernel_size=1, padding=0 + ) + identity = ( + torch.eye(self.channels).repeat(upscale_ratio, 1).reshape_as(self.upscale_conv.weight) + ) + + self.upscale_conv.weight.data.copy_(identity) + nn.init.zeros_(self.upscale_conv.bias) + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state: MemoryState, + ) -> torch.FloatTensor: + return gradient_checkpointing( + self.custom_forward, + hidden_states, + memory_state, + enabled=self.training and self.gradient_checkpointing, + ) + + def custom_forward( + self, + hidden_states: torch.FloatTensor, + memory_state: MemoryState, + ) -> torch.FloatTensor: + assert hidden_states.shape[1] == self.channels + + if self.slicing: + split_size = hidden_states.size(2) // 2 + hidden_states = list( + hidden_states.split([split_size, hidden_states.size(2) - split_size], dim=2) + ) + else: + hidden_states = [hidden_states] + + for i in range(len(hidden_states)): + hidden_states[i] = self.upscale_conv(hidden_states[i]) + hidden_states[i] = rearrange( + hidden_states[i], + "b (x y z c) f h w -> b c (f z) (h x) (w y)", + x=self.spatial_ratio, + y=self.spatial_ratio, + z=self.temporal_ratio, + ) + + # [Overridden] For causal temporal conv + if self.temporal_up and memory_state != MemoryState.ACTIVE: + hidden_states[0] = remove_head(hidden_states[0]) + + if self.slicing: + hidden_states = self.conv(hidden_states, memory_state=memory_state) + return torch.cat(hidden_states, dim=2) + else: + return self.conv(hidden_states[0], memory_state=memory_state) + + +class Downsample3D(nn.Module): + """A 3D downsampling layer.""" + + def __init__( + self, + channels: int, + inflation_mode: _inflation_mode_t = "tail", + temporal_down: bool = False, + spatial_down: bool = True, + ): + super().__init__() + self.channels = channels + self.temporal_down = temporal_down + self.spatial_down = spatial_down + + self.temporal_ratio = 2 if temporal_down else 1 + self.spatial_ratio = 2 if spatial_down else 1 + + self.temporal_kernel = 3 if temporal_down else 1 + self.spatial_kernel = 3 if spatial_down else 1 + + self.conv = init_causal_conv3d( + self.channels, + self.channels, + kernel_size=(self.temporal_kernel, self.spatial_kernel, self.spatial_kernel), + stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), + padding=((1 if self.temporal_down else 0), 0, 0), + inflation_mode=inflation_mode, + ) + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state: MemoryState, + ) -> torch.FloatTensor: + return gradient_checkpointing( + self.custom_forward, + hidden_states, + memory_state, + enabled=self.training and self.gradient_checkpointing, + ) + + def custom_forward( + self, + hidden_states: torch.FloatTensor, + memory_state: MemoryState, + ) -> torch.FloatTensor: + + assert hidden_states.shape[1] == self.channels + + if self.spatial_down: + hidden_states = safe_pad_operation(hidden_states, (0, 1, 0, 1), mode="constant", value=0) + + hidden_states = self.conv(hidden_states, memory_state=memory_state) + return hidden_states + + +class ResnetBlock3D(ResnetBlock2D): + def __init__( + self, + *args, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + **kwargs, + ): + super().__init__(*args, **kwargs) + self.conv1 = init_causal_conv3d( + self.in_channels, + self.out_channels, + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.conv2 = init_causal_conv3d( + self.out_channels, + self.out_channels, + kernel_size=(1, 3, 3) if time_receptive_field == "half" else (3, 3, 3), + stride=1, + padding=(0, 1, 1) if time_receptive_field == "half" else (1, 1, 1), + inflation_mode=inflation_mode, + ) + + if self.use_in_shortcut: + self.conv_shortcut = init_causal_conv3d( + self.in_channels, + self.out_channels, + kernel_size=1, + stride=1, + padding=0, + bias=(self.conv_shortcut.bias is not None), + inflation_mode=inflation_mode, + ) + self.gradient_checkpointing = False + + def forward(self, input_tensor: torch.Tensor, memory_state: MemoryState = MemoryState.UNSET): + return gradient_checkpointing( + self.custom_forward, + input_tensor, + memory_state, + enabled=self.training and self.gradient_checkpointing, + ) + + def custom_forward( + self, input_tensor: torch.Tensor, memory_state: MemoryState = MemoryState.UNSET + ): + assert memory_state != MemoryState.UNSET + hidden_states = input_tensor + + hidden_states = causal_norm_wrapper(self.norm1, hidden_states) + hidden_states = self.nonlinearity(hidden_states) + hidden_states = self.conv1(hidden_states, memory_state=memory_state) + + hidden_states = causal_norm_wrapper(self.norm2, hidden_states) + hidden_states = self.nonlinearity(hidden_states) + hidden_states = self.dropout(hidden_states) + hidden_states = self.conv2(hidden_states, memory_state=memory_state) + + if self.conv_shortcut is not None: + input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state) + + output_tensor = input_tensor + hidden_states + + return output_tensor + + +class DownEncoderBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + add_downsample: bool = True, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_down: bool = True, + spatial_down: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + ResnetBlock3D( + in_channels=in_channels, + out_channels=out_channels, + dropout=dropout, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + self.downsamplers = None + if add_downsample: + # Todo: Refactor this line before V5 Image VAE Training. + self.downsamplers = nn.ModuleList( + [ + Downsample3D( + channels=out_channels, + inflation_mode=inflation_mode, + temporal_down=temporal_down, + spatial_down=spatial_down, + ) + ] + ) + + def forward( + self, hidden_states: torch.FloatTensor, memory_state: MemoryState + ) -> torch.FloatTensor: + for resnet in self.resnets: + hidden_states = resnet(hidden_states, memory_state=memory_state) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states, memory_state=memory_state) + + return hidden_states + + +class UpDecoderBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + add_upsample: bool = True, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up: bool = True, + spatial_up: bool = True, + slicing: bool = False, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + input_channels = in_channels if i == 0 else out_channels + + resnets.append( + ResnetBlock3D( + in_channels=input_channels, + out_channels=out_channels, + dropout=dropout, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + self.upsamplers = None + # Todo: Refactor this line before V5 Image VAE Training. + if add_upsample: + self.upsamplers = nn.ModuleList( + [ + Upsample3D( + channels=out_channels, + inflation_mode=inflation_mode, + temporal_up=temporal_up, + spatial_up=spatial_up, + slicing=slicing, + ) + ] + ) + + def forward( + self, hidden_states: torch.FloatTensor, memory_state: MemoryState + ) -> torch.FloatTensor: + for resnet in self.resnets: + hidden_states = resnet(hidden_states, memory_state=memory_state) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, memory_state=memory_state) + + return hidden_states + + +class UNetMidBlock3D(nn.Module): + def __init__( + self, + channels: int, + dropout: float = 0.0, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + self.resnets = nn.ModuleList( + [ + ResnetBlock3D( + in_channels=channels, + out_channels=channels, + dropout=dropout, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ), + ResnetBlock3D( + in_channels=channels, + out_channels=channels, + dropout=dropout, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ), + ] + ) + + def forward(self, hidden_states: torch.Tensor, memory_state: MemoryState): + for resnet in self.resnets: + hidden_states = resnet(hidden_states, memory_state) + return hidden_states + + +class Encoder3D(nn.Module): + r""" + The `Encoder` layer of a variational autoencoder that encodes + its input into a latent representation. + """ + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + double_z: bool = True, + temporal_down_num: int = 2, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), + ): + super().__init__() + self.layers_per_block = layers_per_block + + self.temporal_down_num = temporal_down_num + + self.conv_in = init_causal_conv3d( + in_channels, + block_out_channels[0], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.down_blocks = nn.ModuleList([]) + + # down + output_channel = block_out_channels[0] + for i in range(len(block_out_channels)): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + is_temporal_down_block = i >= len(block_out_channels) - self.temporal_down_num - 1 + # Note: take the last one + + down_block = DownEncoderBlock3D( + num_layers=self.layers_per_block, + in_channels=input_channel, + out_channels=output_channel, + add_downsample=not is_final_block, + temporal_down=is_temporal_down_block, + spatial_down=True, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.down_blocks.append(down_block) + + # mid + self.mid_block = UNetMidBlock3D( + channels=block_out_channels[-1], + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + # out + self.conv_norm_out = nn.GroupNorm( + num_channels=block_out_channels[-1], num_groups=32, eps=1e-6 + ) + self.conv_act = nn.SiLU() + + conv_out_channels = 2 * out_channels if double_z else out_channels + self.conv_out = init_causal_conv3d( + block_out_channels[-1], conv_out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + assert len(selective_checkpointing) == len(self.down_blocks) + self.set_gradient_checkpointing(selective_checkpointing) + + def set_gradient_checkpointing(self, checkpointing_types): + gradient_checkpointing = [] + for down_block, sac_type in zip(self.down_blocks, checkpointing_types): + if sac_type == "coarse": + gradient_checkpointing.append(True) + elif sac_type == "fine": + for n, m in down_block.named_modules(): + if hasattr(m, "gradient_checkpointing"): + m.gradient_checkpointing = True + logger.debug(f"set gradient_checkpointing: {n}") + gradient_checkpointing.append(False) + else: + gradient_checkpointing.append(False) + self.gradient_checkpointing = gradient_checkpointing + logger.info(f"[Encoder3D] gradient_checkpointing: {checkpointing_types}") + + def forward(self, sample: torch.FloatTensor, memory_state: MemoryState) -> torch.FloatTensor: + r"""The forward method of the `Encoder` class.""" + sample = self.conv_in(sample, memory_state=memory_state) + # down + for down_block, sac in zip(self.down_blocks, self.gradient_checkpointing): + sample = gradient_checkpointing( + down_block, + sample, + memory_state=memory_state, + enabled=self.training and sac, + ) + + # middle + sample = self.mid_block(sample, memory_state=memory_state) + + # post-process + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state) + + return sample + + +class Decoder3D(nn.Module): + r""" + The `Decoder` layer of a variational autoencoder that + decodes its latent representation into an output sample. + """ + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up_num: int = 2, + slicing_up_num: int = 0, + selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_up_num = temporal_up_num + + self.conv_in = init_causal_conv3d( + in_channels, + block_out_channels[-1], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.up_blocks = nn.ModuleList([]) + + # mid + self.mid_block = UNetMidBlock3D( + channels=block_out_channels[-1], + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + # up + reversed_block_out_channels = list(reversed(block_out_channels)) + output_channel = reversed_block_out_channels[0] + for i in range(len(reversed_block_out_channels)): + prev_output_channel = output_channel + output_channel = reversed_block_out_channels[i] + + is_final_block = i == len(block_out_channels) - 1 + is_temporal_up_block = i < self.temporal_up_num + is_slicing_up_block = i >= len(block_out_channels) - slicing_up_num + # Note: Keep symmetric + + up_block = UpDecoderBlock3D( + num_layers=self.layers_per_block + 1, + in_channels=prev_output_channel, + out_channels=output_channel, + add_upsample=not is_final_block, + temporal_up=is_temporal_up_block, + slicing=is_slicing_up_block, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.up_blocks.append(up_block) + + # out + self.conv_norm_out = nn.GroupNorm( + num_channels=block_out_channels[0], num_groups=32, eps=1e-6 + ) + self.conv_act = nn.SiLU() + self.conv_out = init_causal_conv3d( + block_out_channels[0], out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + assert len(selective_checkpointing) == len(self.up_blocks) + self.set_gradient_checkpointing(selective_checkpointing) + + def set_gradient_checkpointing(self, checkpointing_types): + gradient_checkpointing = [] + for up_block, sac_type in zip(self.up_blocks, checkpointing_types): + if sac_type == "coarse": + gradient_checkpointing.append(True) + elif sac_type == "fine": + for n, m in up_block.named_modules(): + if hasattr(m, "gradient_checkpointing"): + m.gradient_checkpointing = True + logger.debug(f"set gradient_checkpointing: {n}") + gradient_checkpointing.append(False) + else: + gradient_checkpointing.append(False) + self.gradient_checkpointing = gradient_checkpointing + logger.info(f"[Decoder3D] gradient_checkpointing: {checkpointing_types}") + + def forward(self, sample: torch.FloatTensor, memory_state: MemoryState) -> torch.FloatTensor: + r"""The forward method of the `Decoder` class.""" + + sample = self.conv_in(sample, memory_state=memory_state) + + # middle + sample = self.mid_block(sample, memory_state=memory_state) + + # up + for up_block, sac in zip(self.up_blocks, self.gradient_checkpointing): + sample = gradient_checkpointing( + up_block, + sample, + memory_state=memory_state, + enabled=self.training and sac, + ) + + # post-process + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state) + + return sample + + +class VideoAutoencoderKL(nn.Module): + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + block_out_channels: Tuple[int] = (64,), + layers_per_block: int = 1, + latent_channels: int = 4, + use_quant_conv: bool = True, + use_post_quant_conv: bool = True, + enc_selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), + dec_selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), + temporal_scale_num: int = 3, + slicing_up_num: int = 0, + inflation_mode: _inflation_mode_t = "tail", + time_receptive_field: _receptive_field_t = "half", + slicing_sample_min_size: int = None, + spatial_downsample_factor: int = 16, + temporal_downsample_factor: int = 8, + freeze_encoder: bool = False, + ): + super().__init__() + self.spatial_downsample_factor = spatial_downsample_factor + self.temporal_downsample_factor = temporal_downsample_factor + self.freeze_encoder = freeze_encoder + if slicing_sample_min_size is None: + slicing_sample_min_size = temporal_downsample_factor + self.slicing_sample_min_size = slicing_sample_min_size + self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num) + + # pass init params to Encoder + self.encoder = Encoder3D( + in_channels=in_channels, + out_channels=latent_channels, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + double_z=True, + temporal_down_num=temporal_scale_num, + selective_checkpointing=enc_selective_checkpointing, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + # pass init params to Decoder + self.decoder = Decoder3D( + in_channels=latent_channels, + out_channels=out_channels, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + # [Override] add temporal_up_num parameter + temporal_up_num=temporal_scale_num, + slicing_up_num=slicing_up_num, + selective_checkpointing=dec_selective_checkpointing, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.quant_conv = ( + init_causal_conv3d( + in_channels=2 * latent_channels, + out_channels=2 * latent_channels, + kernel_size=1, + inflation_mode=inflation_mode, + ) + if use_quant_conv + else None + ) + self.post_quant_conv = ( + init_causal_conv3d( + in_channels=latent_channels, + out_channels=latent_channels, + kernel_size=1, + inflation_mode=inflation_mode, + ) + if use_post_quant_conv + else None + ) + + self.use_slicing = False + + def enable_slicing(self): + self.use_slicing = True + + def disable_slicing(self): + self.use_slicing = False + + def encode(self, x: torch.FloatTensor) -> CausalEncoderOutput: + if x.ndim == 4: + x = x.unsqueeze(2) + h = self.slicing_encode(x) + p = DiagonalGaussianDistribution(h) + z = p.sample() + return CausalEncoderOutput(z, p) + + def decode(self, z: torch.FloatTensor) -> CausalDecoderOutput: + if z.ndim == 4: + z = z.unsqueeze(2) + x = self.slicing_decode(z) + return CausalDecoderOutput(x) + + def _encode(self, x: torch.Tensor, memory_state: MemoryState) -> torch.Tensor: + x = causal_conv_slice_inputs(x, self.slicing_sample_min_size, memory_state=memory_state) + h = self.encoder(x, memory_state=memory_state) + h = self.quant_conv(h, memory_state=memory_state) if self.quant_conv is not None else h + h = causal_conv_gather_outputs(h) + return h + + def _decode(self, z: torch.Tensor, memory_state: MemoryState) -> torch.Tensor: + z = causal_conv_slice_inputs(z, self.slicing_latent_min_size, memory_state=memory_state) + z = ( + self.post_quant_conv(z, memory_state=memory_state) + if self.post_quant_conv is not None + else z + ) + x = self.decoder(z, memory_state=memory_state) + x = causal_conv_gather_outputs(x) + return x + + def slicing_encode(self, x: torch.Tensor) -> torch.Tensor: + sp_size = get_sequence_parallel_world_size() + if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size * sp_size: + x_slices = x[:, :, 1:].split(split_size=self.slicing_sample_min_size * sp_size, dim=2) + encoded_slices = [ + self._encode( + torch.cat((x[:, :, :1], x_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + ) + ] + for x_idx in range(1, len(x_slices)): + encoded_slices.append( + self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE) + ) + return torch.cat(encoded_slices, dim=2) + else: + return self._encode(x, memory_state=MemoryState.DISABLED) + + def slicing_decode(self, z: torch.Tensor) -> torch.Tensor: + sp_size = get_sequence_parallel_world_size() + if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size * sp_size: + z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size * sp_size, dim=2) + decoded_slices = [ + self._decode( + torch.cat((z[:, :, :1], z_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + ) + ] + for z_idx in range(1, len(z_slices)): + decoded_slices.append( + self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE) + ) + return torch.cat(decoded_slices, dim=2) + else: + return self._decode(z, memory_state=MemoryState.DISABLED) + + def forward(self, x: torch.FloatTensor) -> CausalAutoencoderOutput: + with torch.no_grad() if self.freeze_encoder else nullcontext(): + z, p = self.encode(x) + x = self.decode(z).sample + return CausalAutoencoderOutput(x, z, p) + + def preprocess(self, x: torch.Tensor): + # x should in [B, C, T, H, W], [B, C, H, W] + assert x.ndim == 4 or x.size(2) % self.temporal_downsample_factor == 1 + return x + + def postprocess(self, x: torch.Tensor): + # x should in [B, C, T, H, W], [B, C, H, W] + return x + + def set_causal_slicing( + self, + *, + split_size: Optional[int], + memory_device: _memory_device_t, + ): + assert ( + split_size is None or memory_device is not None + ), "if split_size is set, memory_device must not be None." + if split_size is not None: + self.enable_slicing() + self.slicing_sample_min_size = split_size + self.slicing_latent_min_size = split_size // self.temporal_downsample_factor + else: + self.disable_slicing() + for module in self.modules(): + if isinstance(module, InflatedCausalConv3d): + module.set_memory_device(memory_device) + + def set_memory_limit(self, conv_max_mem: Optional[float], norm_max_mem: Optional[float]): + set_norm_limit(norm_max_mem) + for m in self.modules(): + if isinstance(m, InflatedCausalConv3d): + m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf")) + + +class VideoAutoencoderKLWrapper(VideoAutoencoderKL): + def __init__( + self, *args, spatial_downsample_factor: int, temporal_downsample_factor: int, **kwargs + ): + self.spatial_downsample_factor = spatial_downsample_factor + self.temporal_downsample_factor = temporal_downsample_factor + super().__init__(*args, **kwargs) + + def forward(self, x) -> CausalAutoencoderOutput: + z, _, p = self.encode(x) + x, _ = self.decode(z) + return CausalAutoencoderOutput(x, z, None, p) + + def encode(self, x) -> CausalEncoderOutput: + if x.ndim == 4: + x = x.unsqueeze(2) + p = super().encode(x).latent_dist + z = p.sample().squeeze(2) + return CausalEncoderOutput(z, None, p) + + def decode(self, z) -> CausalDecoderOutput: + if z.ndim == 4: + z = z.unsqueeze(2) + x = super().decode(z).sample.squeeze(2) + return CausalDecoderOutput(x, None) + + def preprocess(self, x): + # x should in [B, C, T, H, W], [B, C, H, W] + assert x.ndim == 4 or x.size(2) % 4 == 1 + return x + + def postprocess(self, x): + # x should in [B, C, T, H, W], [B, C, H, W] + return x + + def set_causal_slicing( + self, + *, + split_size: Optional[int], + memory_device: Optional[Literal["cpu", "same"]], + ): + assert ( + split_size is None or memory_device is not None + ), "if split_size is set, memory_device must not be None." + if split_size is not None: + self.enable_slicing() + else: + self.disable_slicing() + self.slicing_sample_min_size = split_size + if split_size is not None: + self.slicing_latent_min_size = split_size // self.temporal_downsample_factor + for module in self.modules(): + if isinstance(module, InflatedCausalConv3d): + module.set_memory_device(memory_device) \ No newline at end of file diff --git a/src/models/video_vae_v3_mine_bad/s8_c16_t4_inflation_sd3.yaml b/src/models/video_vae_v3_mine_bad/s8_c16_t4_inflation_sd3.yaml new file mode 100644 index 0000000..9491226 --- /dev/null +++ b/src/models/video_vae_v3_mine_bad/s8_c16_t4_inflation_sd3.yaml @@ -0,0 +1,28 @@ +act_fn: silu +block_out_channels: + - 128 + - 256 + - 512 + - 512 +down_block_types: + - DownEncoderBlock3D + - DownEncoderBlock3D + - DownEncoderBlock3D + - DownEncoderBlock3D +in_channels: 3 +latent_channels: 16 +layers_per_block: 2 +norm_num_groups: 32 +out_channels: 3 +slicing_sample_min_size: 4 +temporal_scale_num: 2 +inflation_mode: pad +up_block_types: + - UpDecoderBlock3D + - UpDecoderBlock3D + - UpDecoderBlock3D + - UpDecoderBlock3D +spatial_downsample_factor: 8 +temporal_downsample_factor: 4 +use_quant_conv: False +use_post_quant_conv: False diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 674581f..4246adb 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -10,8 +10,8 @@ import torch import gc import time from typing import Tuple, Optional -from common.cache import Cache -from models.dit_v2.rope import RotaryEmbeddingBase +from src.common.cache import Cache +from src.models.dit_v2.rope import RotaryEmbeddingBase def get_basic_vram_info(): diff --git a/src/utils/color_fix.py b/src/utils/color_fix.py index a037481..fe6aee7 100644 --- a/src/utils/color_fix.py +++ b/src/utils/color_fix.py @@ -2,7 +2,7 @@ import torch from PIL import Image from torch import Tensor from torch.nn import functional as F -from common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation +from src.common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation from torchvision.transforms import ToTensor, ToPILImage def adain_color_fix(target: Image, source: Image):