fix import
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+7
-7
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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():
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user