diff --git a/Chrono/memory_profiler.py b/Chrono/memory_profiler.py new file mode 100644 index 0000000..9e7058a --- /dev/null +++ b/Chrono/memory_profiler.py @@ -0,0 +1,29 @@ +import torch +from loguru import logger + + +def peak_memory_decorator(func): + def wrapper(*args, **kwargs): + # 检查是否在分布式环境中 + rank_info = "" + if torch.distributed.is_available() and torch.distributed.is_initialized(): + rank = torch.distributed.get_rank() + rank_info = f"Rank {rank} - " + + # 如果使用GPU,重置显存统计 + if torch.cuda.is_available(): + torch.cuda.reset_peak_memory_stats() + + # 执行目标函数 + result = func(*args, **kwargs) + + # 获取峰值显存 + if torch.cuda.is_available(): + peak_memory = torch.cuda.max_memory_allocated() / (1024**3) # 转换为GB + logger.info(f"{rank_info}Function '{func.__qualname__}' Peak Memory: {peak_memory:.2f} GB") + else: + logger.info(f"{rank_info}Function '{func.__qualname__}' executed without GPU.") + + return result + + return wrapper diff --git a/Chrono/tae.py b/Chrono/tae.py new file mode 100644 index 0000000..7b357cc --- /dev/null +++ b/Chrono/tae.py @@ -0,0 +1,296 @@ +#!/usr/bin/env python3 +""" +Tiny AutoEncoder for Hunyuan Video +(DNN for encoding / decoding videos to Hunyuan Video's latent space) +""" + +import os +from collections import namedtuple + +import torch +import torch.nn as nn +import torch.nn.functional as F +from safetensors.torch import load_file +from tqdm.auto import tqdm + +DecoderResult = namedtuple("DecoderResult", ("frame", "memory")) +TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index")) + + +def conv(n_in, n_out, **kwargs): + return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs) + + +class Clamp(nn.Module): + def forward(self, x): + return torch.tanh(x / 3) * 3 + + +class MemBlock(nn.Module): + def __init__(self, n_in, n_out): + super().__init__() + self.conv = nn.Sequential(conv(n_in * 2, n_out), nn.ReLU(inplace=True), conv(n_out, n_out), nn.ReLU(inplace=True), conv(n_out, n_out)) + self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity() + self.act = nn.ReLU(inplace=True) + + def forward(self, x, past): + return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x)) + + +class TPool(nn.Module): + def __init__(self, n_f, stride): + super().__init__() + self.stride = stride + self.conv = nn.Conv2d(n_f * stride, n_f, 1, bias=False) + + def forward(self, x): + _NT, C, H, W = x.shape + return self.conv(x.reshape(-1, self.stride * C, H, W)) + + +class TGrow(nn.Module): + def __init__(self, n_f, stride): + super().__init__() + self.stride = stride + self.conv = nn.Conv2d(n_f, n_f * stride, 1, bias=False) + + def forward(self, x): + _NT, C, H, W = x.shape + x = self.conv(x) + return x.reshape(-1, C, H, W) + + +def apply_model_with_memblocks(model, x, parallel, show_progress_bar): + """ + Apply a sequential model with memblocks to the given input. + Args: + - model: nn.Sequential of blocks to apply + - x: input data, of dimensions NTCHW + - parallel: if True, parallelize over timesteps (fast but uses O(T) memory) + if False, each timestep will be processed sequentially (slow but uses O(1) memory) + - show_progress_bar: if True, enables tqdm progressbar display + + Returns NTCHW tensor of output data. + """ + assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor" + N, T, C, H, W = x.shape + if parallel: + x = x.reshape(N * T, C, H, W) + # parallel over input timesteps, iterate over blocks + for b in tqdm(model, disable=not show_progress_bar): + if isinstance(b, MemBlock): + NT, C, H, W = x.shape + T = NT // N + _x = x.reshape(N, T, C, H, W) + mem = F.pad(_x, (0, 0, 0, 0, 0, 0, 1, 0), value=0)[:, :T].reshape(x.shape) + x = b(x, mem) + else: + x = b(x) + NT, C, H, W = x.shape + T = NT // N + x = x.view(N, T, C, H, W) + else: + # TODO(oboerbohan): at least on macos this still gradually uses more memory during decode... + # need to fix :( + out = [] + # iterate over input timesteps and also iterate over blocks. + # because of the cursed TPool/TGrow blocks, this is not a nested loop, + # it's actually a ***graph traversal*** problem! so let's make a queue + work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))] + # in addition to manually managing our queue, we also need to manually manage our progressbar. + # we'll update it for every source node that we consume. + progress_bar = tqdm(range(T), disable=not show_progress_bar) + # we'll also need a separate addressable memory per node as well + mem = [None] * len(model) + while work_queue: + xt, i = work_queue.pop(0) + if i == 0: + # new source node consumed + progress_bar.update(1) + if i == len(model): + # reached end of the graph, append result to output list + out.append(xt) + else: + # fetch the block to process + b = model[i] + if isinstance(b, MemBlock): + # mem blocks are simple since we're visiting the graph in causal order + if mem[i] is None: + xt_new = b(xt, xt * 0) + mem[i] = xt + else: + xt_new = b(xt, mem[i]) + mem[i].copy_(xt) # inplace might reduce mysterious pytorch memory allocations? doesn't help though + # add successor to work queue + work_queue.insert(0, TWorkItem(xt_new, i + 1)) + elif isinstance(b, TPool): + # pool blocks are miserable + if mem[i] is None: + mem[i] = [] # pool memory is itself a queue of inputs to pool + mem[i].append(xt) + if len(mem[i]) > b.stride: + # pool mem is in invalid state, we should have pooled before this + raise ValueError("???") + elif len(mem[i]) < b.stride: + # pool mem is not yet full, go back to processing the work queue + pass + else: + # pool mem is ready, run the pool block + N, C, H, W = xt.shape + xt = b(torch.cat(mem[i], 1).view(N * b.stride, C, H, W)) + # reset the pool mem + mem[i] = [] + # add successor to work queue + work_queue.insert(0, TWorkItem(xt, i + 1)) + elif isinstance(b, TGrow): + xt = b(xt) + NT, C, H, W = xt.shape + # each tgrow has multiple successor nodes + for xt_next in reversed(xt.view(N, b.stride * C, H, W).chunk(b.stride, 1)): + # add successor to work queue + work_queue.insert(0, TWorkItem(xt_next, i + 1)) + else: + # normal block with no funny business + xt = b(xt) + # add successor to work queue + work_queue.insert(0, TWorkItem(xt, i + 1)) + progress_bar.close() + x = torch.stack(out, 1) + return x + + +class TAEHV(nn.Module): + def __init__(self, checkpoint_path="taehv.pth", decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), patch_size=1, latent_channels=16, model_type="wan21"): + """Initialize pretrained TAEHV from the given checkpoint. + + Arg: + checkpoint_path: path to weight file to load. taehv.pth for Hunyuan, taew2_1.pth for Wan 2.1. + decoder_time_upscale: whether temporal upsampling is enabled for each block. upsampling can be disabled for a cheaper preview. + decoder_space_upscale: whether spatial upsampling is enabled for each block. upsampling can be disabled for a cheaper preview. + patch_size: input/output pixelshuffle patch-size for this model. + latent_channels: number of latent channels (z dim) for this model. + """ + super().__init__() + self.patch_size = patch_size + self.latent_channels = latent_channels + self.image_channels = 3 + self.is_cogvideox = checkpoint_path is not None and "taecvx" in checkpoint_path + # if checkpoint_path is not None and "taew2_2" in checkpoint_path: + # self.patch_size, self.latent_channels = 2, 48 + + if model_type == "wan22": + self.patch_size, self.latent_channels = 2, 48 + self.encoder = nn.Sequential( + conv(self.image_channels * self.patch_size**2, 64), + nn.ReLU(inplace=True), + TPool(64, 2), + conv(64, 64, stride=2, bias=False), + MemBlock(64, 64), + MemBlock(64, 64), + MemBlock(64, 64), + TPool(64, 2), + conv(64, 64, stride=2, bias=False), + MemBlock(64, 64), + MemBlock(64, 64), + MemBlock(64, 64), + TPool(64, 1), + conv(64, 64, stride=2, bias=False), + MemBlock(64, 64), + MemBlock(64, 64), + MemBlock(64, 64), + conv(64, self.latent_channels), + ) + n_f = [256, 128, 64, 64] + self.frames_to_trim = 2 ** sum(decoder_time_upscale) - 1 + self.decoder = nn.Sequential( + Clamp(), + conv(self.latent_channels, n_f[0]), + nn.ReLU(inplace=True), + MemBlock(n_f[0], n_f[0]), + MemBlock(n_f[0], n_f[0]), + MemBlock(n_f[0], n_f[0]), + nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1), + TGrow(n_f[0], 1), + conv(n_f[0], n_f[1], bias=False), + MemBlock(n_f[1], n_f[1]), + MemBlock(n_f[1], n_f[1]), + MemBlock(n_f[1], n_f[1]), + nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1), + TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1), + conv(n_f[1], n_f[2], bias=False), + MemBlock(n_f[2], n_f[2]), + MemBlock(n_f[2], n_f[2]), + MemBlock(n_f[2], n_f[2]), + nn.Upsample(scale_factor=2 if decoder_space_upscale[2] else 1), + TGrow(n_f[2], 2 if decoder_time_upscale[1] else 1), + conv(n_f[2], n_f[3], bias=False), + nn.ReLU(inplace=True), + conv(n_f[3], self.image_channels * self.patch_size**2), + ) + if checkpoint_path is not None: + ext = os.path.splitext(checkpoint_path)[1].lower() + + if ext == ".pth": + state_dict = torch.load(checkpoint_path, map_location="cpu", weights_only=True) + elif ext == ".safetensors": + state_dict = load_file(checkpoint_path, device="cpu") + else: + raise ValueError(f"Unsupported checkpoint format: {ext}. Supported formats: .pth, .safetensors") + + self.load_state_dict(self.patch_tgrow_layers(state_dict)) + + def patch_tgrow_layers(self, sd): + """Patch TGrow layers to use a smaller kernel if needed. + + Args: + sd: state dict to patch + """ + new_sd = self.state_dict() + for i, layer in enumerate(self.decoder): + if isinstance(layer, TGrow): + key = f"decoder.{i}.conv.weight" + if sd[key].shape[0] > new_sd[key].shape[0]: + # take the last-timestep output channels + sd[key] = sd[key][-new_sd[key].shape[0] :] + return sd + + def encode_video(self, x, parallel=True, show_progress_bar=True): + """Encode a sequence of frames. + + Args: + x: input NTCHW RGB (C=3) tensor with values in [0, 1]. + parallel: if True, all frames will be processed at once. + (this is faster but may require more memory). + if False, frames will be processed sequentially. + Returns NTCHW latent tensor with ~Gaussian values. + """ + if self.patch_size > 1: + x = F.pixel_unshuffle(x, self.patch_size) + if x.shape[1] % 4 != 0: + # pad at end to multiple of 4 + n_pad = 4 - x.shape[1] % 4 + padding = x[:, -1:].repeat_interleave(n_pad, dim=1) + x = torch.cat([x, padding], 1) + return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar) + + def decode_video(self, x, parallel=True, show_progress_bar=True): + """Decode a sequence of frames. + + Args: + x: input NTCHW latent (C=12) tensor with ~Gaussian values. + parallel: if True, all frames will be processed at once. + (this is faster but may require more memory). + if False, frames will be processed sequentially. + Returns NTCHW RGB tensor with ~[0, 1] values. + """ + skip_trim = self.is_cogvideox and x.shape[1] % 2 == 0 + x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar) + x = x.clamp_(0, 1) + if self.patch_size > 1: + x = F.pixel_shuffle(x, self.patch_size) + if skip_trim: + # skip trimming for cogvideox to make frame counts match. + # this still doesn't have correct temporal alignment for certain frame counts + # (cogvideox seems to pad at the start?), but for multiple-of-4 it's fine. + return x + return x[:, self.frames_to_trim :] diff --git a/Chrono/utils.py b/Chrono/utils.py new file mode 100644 index 0000000..4b0d491 --- /dev/null +++ b/Chrono/utils.py @@ -0,0 +1,485 @@ +import os +import random +import subprocess +from typing import Optional + +import imageio +import imageio_ffmpeg as ffmpeg +import numpy as np +import safetensors +import torch +import torch.distributed as dist +import torchvision +from einops import rearrange +from loguru import logger + + +def seed_all(seed): + random.seed(seed) + os.environ["PYTHONHASHSEED"] = str(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + torch.backends.cudnn.benchmark = False + torch.backends.cudnn.deterministic = True + + +def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=1, fps=24): + """save videos by video tensor + copy from https://github.com/guoyww/AnimateDiff/blob/e92bd5671ba62c0d774a32951453e328018b7c5b/animatediff/utils/util.py#L61 + + Args: + videos (torch.Tensor): video tensor predicted by the model + path (str): path to save video + rescale (bool, optional): rescale the video tensor from [-1, 1] to . Defaults to False. + n_rows (int, optional): Defaults to 1. + fps (int, optional): video save fps. Defaults to 8. + """ + videos = rearrange(videos, "b c t h w -> t b c h w") + outputs = [] + for x in videos: + x = torchvision.utils.make_grid(x, nrow=n_rows) + x = x.transpose(0, 1).transpose(1, 2).squeeze(-1) + if rescale: + x = (x + 1.0) / 2.0 # -1,1 -> 0,1 + x = torch.clamp(x, 0, 1) + x = (x * 255).numpy().astype(np.uint8) + outputs.append(x) + + os.makedirs(os.path.dirname(path), exist_ok=True) + imageio.mimsave(path, outputs, fps=fps) + + +def cache_video( + tensor, + save_file: str, + fps=30, + suffix=".mp4", + nrow=8, + normalize=True, + value_range=(-1, 1), + retry=5, +): + save_dir = os.path.dirname(save_file) + try: + if not os.path.exists(save_dir): + os.makedirs(save_dir, exist_ok=True) + except Exception as e: + logger.error(f"Failed to create directory: {save_dir}, error: {e}") + return None + + cache_file = save_file + + # save to cache + error = None + for _ in range(retry): + try: + # preprocess + tensor = tensor.clamp(min(value_range), max(value_range)) # type: ignore + tensor = torch.stack( + [torchvision.utils.make_grid(u, nrow=nrow, normalize=normalize, value_range=value_range) for u in tensor.unbind(2)], + dim=1, + ).permute(1, 2, 3, 0) + tensor = (tensor * 255).type(torch.uint8).cpu() + + # write video + writer = imageio.get_writer(cache_file, fps=fps, codec="libx264", quality=8) + for frame in tensor.numpy(): + writer.append_data(frame) + writer.close() + del tensor + torch.cuda.empty_cache() + return cache_file + except Exception as e: + error = e + continue + else: + logger.info(f"cache_video failed, error: {error}", flush=True) + return None + + +def vae_to_comfyui_image(vae_output: torch.Tensor) -> torch.Tensor: + """ + Convert VAE decoder output to ComfyUI Image format + + Args: + vae_output: VAE decoder output tensor, typically in range [-1, 1] + Shape: [B, C, T, H, W] or [B, C, H, W] + + Returns: + ComfyUI Image tensor in range [0, 1] + Shape: [B, H, W, C] for single frame or [B*T, H, W, C] for video + """ + # Handle video tensor (5D) vs image tensor (4D) + if vae_output.dim() == 5: + # Video tensor: [B, C, T, H, W] + B, C, T, H, W = vae_output.shape + # Reshape to [B*T, C, H, W] for processing + vae_output = vae_output.permute(0, 2, 1, 3, 4).reshape(B * T, C, H, W) + + # Normalize from [-1, 1] to [0, 1] + images = (vae_output + 1) / 2 + + # Clamp values to [0, 1] + images = torch.clamp(images, 0, 1) + + # Convert from [B, C, H, W] to [B, H, W, C] + images = images.permute(0, 2, 3, 1).cpu() + + return images + + +def vae_to_comfyui_image_inplace(vae_output: torch.Tensor) -> torch.Tensor: + """ + Convert VAE decoder output to ComfyUI Image format (inplace operation) + + Args: + vae_output: VAE decoder output tensor, typically in range [-1, 1] + Shape: [B, C, T, H, W] or [B, C, H, W] + WARNING: This tensor will be modified in-place! + + Returns: + ComfyUI Image tensor in range [0, 1] + Shape: [B, H, W, C] for single frame or [B*T, H, W, C] for video + Note: The returned tensor is the same object as input (modified in-place) + """ + # Handle video tensor (5D) vs image tensor (4D) + if vae_output.dim() == 5: + # Video tensor: [B, C, T, H, W] + B, C, T, H, W = vae_output.shape + # Reshape to [B*T, C, H, W] for processing (inplace view) + vae_output = vae_output.permute(0, 2, 1, 3, 4).contiguous().view(B * T, C, H, W) + + # Normalize from [-1, 1] to [0, 1] (inplace) + vae_output.add_(1).div_(2) + + # Clamp values to [0, 1] (inplace) + vae_output.clamp_(0, 1) + + # Convert from [B, C, H, W] to [B, H, W, C] and move to CPU + vae_output = vae_output.permute(0, 2, 3, 1).cpu() + + return vae_output + + +def save_to_video( + images: torch.Tensor, + output_path: str, + fps: float = 24.0, + method: str = "imageio", + lossless: bool = False, + output_pix_fmt: Optional[str] = "yuv420p", +) -> None: + """ + Save ComfyUI Image tensor to video file + + Args: + images: ComfyUI Image tensor [N, H, W, C] in range [0, 1] + output_path: Path to save the video + fps: Frames per second + method: Save method - "imageio" or "ffmpeg" + lossless: Whether to use lossless encoding (ffmpeg method only) + output_pix_fmt: Pixel format for output (ffmpeg method only) + """ + assert images.dim() == 4 and images.shape[-1] == 3, "Input must be [N, H, W, C] with C=3" + + # Ensure output directory exists + os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True) + + if method == "imageio": + # Convert to uint8 + # frames = (images * 255).cpu().numpy().astype(np.uint8) + frames = (images * 255).to(torch.uint8).cpu().numpy() + imageio.mimsave(output_path, frames, fps=fps) # type: ignore + + elif method == "ffmpeg": + # Convert to numpy and scale to [0, 255] + # frames = (images * 255).cpu().numpy().clip(0, 255).astype(np.uint8) + frames = (images * 255).clamp(0, 255).to(torch.uint8).cpu().numpy() + + # Convert RGB to BGR for OpenCV/FFmpeg + frames = frames[..., ::-1].copy() + + N, height, width, _ = frames.shape + + # Ensure even dimensions for x264 + width += width % 2 + height += height % 2 + + # Get ffmpeg executable from imageio_ffmpeg + ffmpeg_exe = ffmpeg.get_ffmpeg_exe() + + if lossless: + command = [ + ffmpeg_exe, + "-y", # Overwrite output file if it exists + "-f", + "rawvideo", + "-s", + f"{int(width)}x{int(height)}", + "-pix_fmt", + "bgr24", + "-r", + f"{fps}", + "-loglevel", + "error", + "-threads", + "4", + "-i", + "-", # Input from pipe + "-vcodec", + "libx264rgb", + "-crf", + "0", + "-an", # No audio + output_path, + ] + else: + command = [ + ffmpeg_exe, + "-y", # Overwrite output file if it exists + "-f", + "rawvideo", + "-s", + f"{int(width)}x{int(height)}", + "-pix_fmt", + "bgr24", + "-r", + f"{fps}", + "-loglevel", + "error", + "-threads", + "4", + "-i", + "-", # Input from pipe + "-vcodec", + "libx264", + "-pix_fmt", + output_pix_fmt, + "-an", # No audio + output_path, + ] + + # Run FFmpeg + process = subprocess.Popen( + command, + stdin=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + + if process.stdin is None: + raise BrokenPipeError("No stdin buffer received.") + + # Write frames to FFmpeg + for frame in frames: + # Pad frame if needed + if frame.shape[0] < height or frame.shape[1] < width: + padded = np.zeros((height, width, 3), dtype=np.uint8) + padded[: frame.shape[0], : frame.shape[1]] = frame + frame = padded + process.stdin.write(frame.tobytes()) + + process.stdin.close() + process.wait() + + if process.returncode != 0: + error_output = process.stderr.read().decode() if process.stderr else "Unknown error" + raise RuntimeError(f"FFmpeg failed with error: {error_output}") + + else: + raise ValueError(f"Unknown save method: {method}") + + +def remove_substrings_from_keys(original_dict, substr): + new_dict = {} + for key, value in original_dict.items(): + new_dict[key.replace(substr, "")] = value + return new_dict + + +def find_torch_model_path(config, ckpt_config_key=None, filename=None, subdir=["original", "fp8", "int8", "distill_models", "distill_fp8", "distill_int8"]): + if ckpt_config_key and config.get(ckpt_config_key, None) is not None: + return config.get(ckpt_config_key) + + paths_to_check = [ + os.path.join(config["model_path"], filename), + ] + if isinstance(subdir, list): + for sub in subdir: + paths_to_check.insert(0, os.path.join(config["model_path"], sub, filename)) + else: + paths_to_check.insert(0, os.path.join(config["model_path"], subdir, filename)) + + for path in paths_to_check: + if os.path.exists(path): + return path + raise FileNotFoundError(f"PyTorch model file '{filename}' not found.\nPlease download the model from https://huggingface.co/lightx2v/ or specify the model path in the configuration file.") + + +def load_safetensors(in_path, remove_key=None, include_keys=None): + """加载safetensors文件或目录,支持按key包含筛选或排除""" + include_keys = include_keys or [] + if os.path.isdir(in_path): + return load_safetensors_from_dir(in_path, remove_key, include_keys) + elif os.path.isfile(in_path): + return load_safetensors_from_path(in_path, remove_key, include_keys) + else: + raise ValueError(f"{in_path} does not exist") + + +def load_safetensors_from_path(in_path, remove_key=None, include_keys=None): + """从单个safetensors文件加载权重,支持按key筛选""" + include_keys = include_keys or [] + tensors = {} + with safetensors.safe_open(in_path, framework="pt", device="cpu") as f: + for key in f.keys(): + # 优先处理include_keys:如果非空,只保留包含任意指定key的条目 + if include_keys: + if any(inc_key in key for inc_key in include_keys): + tensors[key] = f.get_tensor(key) + # 否则使用remove_key排除 + else: + if not (remove_key and remove_key in key): + tensors[key] = f.get_tensor(key) + return tensors + + +def load_safetensors_from_dir(in_dir, remove_key=None, include_keys=None): + """从目录加载所有safetensors文件,支持按key筛选""" + include_keys = include_keys or [] + tensors = {} + safetensors_files = os.listdir(in_dir) + safetensors_files = [f for f in safetensors_files if f.endswith(".safetensors")] + for f in safetensors_files: + tensors.update(load_safetensors_from_path(os.path.join(in_dir, f), remove_key, include_keys)) + return tensors + + +def load_pt_safetensors(in_path, remove_key=None, include_keys=None): + """加载pt/pth或safetensors权重,支持按key筛选""" + include_keys = include_keys or [] + ext = os.path.splitext(in_path)[-1] + if ext in (".pt", ".pth", ".tar"): + state_dict = torch.load(in_path, map_location="cpu", weights_only=True) + # 处理筛选逻辑 + keys_to_keep = [] + for key in state_dict.keys(): + if include_keys: + if any(inc_key in key for inc_key in include_keys): + keys_to_keep.append(key) + else: + if not (remove_key and remove_key in key): + keys_to_keep.append(key) + # 只保留符合条件的key + state_dict = {k: state_dict[k] for k in keys_to_keep} + else: + state_dict = load_safetensors(in_path, remove_key, include_keys) + return state_dict + + +def load_weights(checkpoint_path, cpu_offload=False, remove_key=None, load_from_rank0=False, include_keys=None): + if not dist.is_initialized() or not load_from_rank0: + # Single GPU mode + logger.info(f"Loading weights from {checkpoint_path}") + cpu_weight_dict = load_pt_safetensors(checkpoint_path, remove_key, include_keys) + return cpu_weight_dict + + # Multi-GPU mode + is_weight_loader = False + current_rank = dist.get_rank() + if current_rank == 0: + is_weight_loader = True + + cpu_weight_dict = {} + if is_weight_loader: + logger.info(f"Loading weights from {checkpoint_path}") + cpu_weight_dict = load_pt_safetensors(checkpoint_path, remove_key) + + meta_dict = {} + if is_weight_loader: + for key, tensor in cpu_weight_dict.items(): + meta_dict[key] = {"shape": tensor.shape, "dtype": tensor.dtype} + + obj_list = [meta_dict] if is_weight_loader else [None] + + src_global_rank = 0 + dist.broadcast_object_list(obj_list, src=src_global_rank) + synced_meta_dict = obj_list[0] + + if cpu_offload: + target_device = "cpu" + distributed_weight_dict = {key: torch.empty(meta["shape"], dtype=meta["dtype"], device=target_device) for key, meta in synced_meta_dict.items()} + dist.barrier() + else: + target_device = torch.device(f"cuda:{current_rank}") + distributed_weight_dict = {key: torch.empty(meta["shape"], dtype=meta["dtype"], device=target_device) for key, meta in synced_meta_dict.items()} + dist.barrier(device_ids=[torch.cuda.current_device()]) + + for key in sorted(synced_meta_dict.keys()): + tensor_to_broadcast = distributed_weight_dict[key] + if is_weight_loader: + tensor_to_broadcast.copy_(cpu_weight_dict[key], non_blocking=True) + + if cpu_offload: + if is_weight_loader: + gpu_tensor = tensor_to_broadcast.cuda() + dist.broadcast(gpu_tensor, src=src_global_rank) + tensor_to_broadcast.copy_(gpu_tensor.cpu(), non_blocking=True) + del gpu_tensor + torch.cuda.empty_cache() + else: + gpu_tensor = torch.empty_like(tensor_to_broadcast, device="cuda") + dist.broadcast(gpu_tensor, src=src_global_rank) + tensor_to_broadcast.copy_(gpu_tensor.cpu(), non_blocking=True) + del gpu_tensor + torch.cuda.empty_cache() + else: + dist.broadcast(tensor_to_broadcast, src=src_global_rank) + + if is_weight_loader: + del cpu_weight_dict + + if cpu_offload: + torch.cuda.empty_cache() + + logger.info(f"Weights distributed across {dist.get_world_size()} devices on {target_device}") + return distributed_weight_dict + + +def masks_like(tensor, zero=False, generator=None, p=0.2, prev_len=1): + assert isinstance(tensor, torch.Tensor) + out = torch.ones_like(tensor) + if zero: + if generator is not None: + random_num = torch.rand(1, generator=generator, device=generator.device).item() + if random_num < p: + out[:, :prev_len] = torch.zeros_like(out[:, :prev_len]) + else: + out[:, :prev_len] = torch.zeros_like(out[:, :prev_len]) + return out + + +def best_output_size(w, h, dw, dh, expected_area): + # float output size + ratio = w / h + ow = (expected_area * ratio) ** 0.5 + oh = expected_area / ow + + # process width first + ow1 = int(ow // dw * dw) + oh1 = int(expected_area / ow1 // dh * dh) + assert ow1 % dw == 0 and oh1 % dh == 0 and ow1 * oh1 <= expected_area + ratio1 = ow1 / oh1 + + # process height first + oh2 = int(oh // dh * dh) + ow2 = int(expected_area / oh2 // dw * dw) + assert oh2 % dh == 0 and ow2 % dw == 0 and ow2 * oh2 <= expected_area + ratio2 = ow2 / oh2 + + # compare ratios + if max(ratio / ratio1, ratio1 / ratio) < max(ratio / ratio2, ratio2 / ratio): + return ow1, oh1 + else: + return ow2, oh2 diff --git a/Chrono/vae.py b/Chrono/vae.py new file mode 100644 index 0000000..5cd35e3 --- /dev/null +++ b/Chrono/vae.py @@ -0,0 +1,1309 @@ +# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange +from loguru import logger + +from .utils import load_weights + +__all__ = [ + "WanVAE", +] + +CACHE_T = 2 + + +class CausalConv3d(nn.Conv3d): + """ + Causal 3d convolusion. + """ + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._padding = ( + self.padding[2], + self.padding[2], + self.padding[1], + self.padding[1], + 2 * self.padding[0], + 0, + ) + self.padding = (0, 0, 0) + + def forward(self, x, cache_x=None): + padding = list(self._padding) + if cache_x is not None and self._padding[4] > 0: + cache_x = cache_x.to(x.device) + x = torch.cat([cache_x, x], dim=2) + padding[4] -= cache_x.shape[2] + x = F.pad(x, padding) + + return super().forward(x) + + +class RMS_norm(nn.Module): + def __init__(self, dim, channel_first=True, images=True, bias=False): + super().__init__() + broadcastable_dims = (1, 1, 1) if not images else (1, 1) + shape = (dim, *broadcastable_dims) if channel_first else (dim,) + + self.channel_first = channel_first + self.scale = dim**0.5 + self.gamma = nn.Parameter(torch.ones(shape)) + self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 + + def forward(self, x): + return F.normalize(x, dim=(1 if self.channel_first else -1)) * self.scale * self.gamma + self.bias + + +class Upsample(nn.Upsample): + def forward(self, x): + """ + Fix bfloat16 support for nearest neighbor interpolation. + """ + return super().forward(x) + + +class Resample(nn.Module): + def __init__(self, dim, mode): + assert mode in ( + "none", + "upsample2d", + "upsample3d", + "downsample2d", + "downsample3d", + ) + super().__init__() + self.dim = dim + self.mode = mode + + # layers + if mode == "upsample2d": + self.resample = nn.Sequential( + Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), + nn.Conv2d(dim, dim // 2, 3, padding=1), + ) + elif mode == "upsample3d": + self.resample = nn.Sequential( + Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), + nn.Conv2d(dim, dim // 2, 3, padding=1), + ) + self.time_conv = CausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0)) + + elif mode == "downsample2d": + self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))) + elif mode == "downsample3d": + self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))) + self.time_conv = CausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0)) + + else: + self.resample = nn.Identity() + + def forward(self, x, feat_cache=None, feat_idx=[0]): + b, c, t, h, w = x.size() + if self.mode == "upsample3d": + if feat_cache is not None: + idx = feat_idx[0] + if feat_cache[idx] is None: + feat_cache[idx] = "Rep" + feat_idx[0] += 1 + else: + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] != "Rep": + # cache last frame of last two chunk + cache_x = torch.cat( + [ + feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] == "Rep": + cache_x = torch.cat( + [torch.zeros_like(cache_x).to(cache_x.device), cache_x], + dim=2, + ) + if feat_cache[idx] == "Rep": + x = self.time_conv(x) + else: + x = self.time_conv(x, feat_cache[idx]) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + + x = x.reshape(b, 2, c, t, h, w) + x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) + x = x.reshape(b, c, t * 2, h, w) + t = x.shape[2] + x = rearrange(x, "b c t h w -> (b t) c h w") + x = self.resample(x) + x = rearrange(x, "(b t) c h w -> b c t h w", t=t) + + if self.mode == "downsample3d": + if feat_cache is not None: + idx = feat_idx[0] + if feat_cache[idx] is None: + feat_cache[idx] = x.clone() + feat_idx[0] += 1 + else: + cache_x = x[:, :, -1:, :, :].clone() + # if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx]!='Rep': + # # cache last frame of last two chunk + # cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2) + + x = self.time_conv(torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2)) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + return x + + def init_weight(self, conv): + conv_weight = conv.weight + nn.init.zeros_(conv_weight) + c1, c2, t, h, w = conv_weight.size() + one_matrix = torch.eye(c1, c2) + init_matrix = one_matrix + nn.init.zeros_(conv_weight) + # conv_weight.data[:,:,-1,1,1] = init_matrix * 0.5 + conv_weight.data[:, :, 1, 0, 0] = init_matrix # * 0.5 + conv.weight.data.copy_(conv_weight) + nn.init.zeros_(conv.bias.data) + + def init_weight2(self, conv): + conv_weight = conv.weight.data + nn.init.zeros_(conv_weight) + c1, c2, t, h, w = conv_weight.size() + init_matrix = torch.eye(c1 // 2, c2) + # init_matrix = repeat(init_matrix, 'o ... -> (o 2) ...').permute(1,0,2).contiguous().reshape(c1,c2) + conv_weight[: c1 // 2, :, -1, 0, 0] = init_matrix + conv_weight[c1 // 2 :, :, -1, 0, 0] = init_matrix + conv.weight.data.copy_(conv_weight) + nn.init.zeros_(conv.bias.data) + + +class ResidualBlock(nn.Module): + def __init__(self, in_dim, out_dim, dropout=0.0): + super().__init__() + self.in_dim = in_dim + self.out_dim = out_dim + + # layers + self.residual = nn.Sequential( + RMS_norm(in_dim, images=False), + nn.SiLU(), + CausalConv3d(in_dim, out_dim, 3, padding=1), + RMS_norm(out_dim, images=False), + nn.SiLU(), + nn.Dropout(dropout), + CausalConv3d(out_dim, out_dim, 3, padding=1), + ) + self.shortcut = CausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() + + def forward(self, x, feat_cache=None, feat_idx=[0]): + h = self.shortcut(x) + for layer in self.residual: + if isinstance(layer, CausalConv3d) and feat_cache is not None: + idx = feat_idx[0] + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and feat_cache[idx] is not None: + # cache last frame of last two chunk + cache_x = torch.cat( + [ + feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + x = layer(x, feat_cache[idx]) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + else: + x = layer(x) + return x + h + + +class AttentionBlock(nn.Module): + """ + Causal self-attention with a single head. + """ + + def __init__(self, dim): + super().__init__() + self.dim = dim + + # layers + self.norm = RMS_norm(dim) + self.to_qkv = nn.Conv2d(dim, dim * 3, 1) + self.proj = nn.Conv2d(dim, dim, 1) + + # zero out the last layer params + nn.init.zeros_(self.proj.weight) + + def forward(self, x): + identity = x + b, c, t, h, w = x.size() + x = rearrange(x, "b c t h w -> (b t) c h w") + x = self.norm(x) + # compute query, key, value + q, k, v = self.to_qkv(x).reshape(b * t, 1, c * 3, -1).permute(0, 1, 3, 2).contiguous().chunk(3, dim=-1) + + # apply attention + x = F.scaled_dot_product_attention( + q, + k, + v, + ) + x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w) + + # output + x = self.proj(x) + x = rearrange(x, "(b t) c h w-> b c t h w", t=t) + return x + identity + + +class Encoder3d(nn.Module): + def __init__(self, dim=128, z_dim=4, dim_mult=[1, 2, 4, 4], num_res_blocks=2, attn_scales=[], temperal_downsample=[True, True, False], dropout=0.0, pruning_rate=0.0): + super().__init__() + self.dim = dim + self.z_dim = z_dim + self.dim_mult = dim_mult + self.num_res_blocks = num_res_blocks + self.attn_scales = attn_scales + self.temperal_downsample = temperal_downsample + + # dimensions + dims = [dim * u for u in [1] + dim_mult] + dims = [int(d * (1 - pruning_rate)) for d in dims] + scale = 1.0 + + # init block + self.conv1 = CausalConv3d(3, dims[0], 3, padding=1) + + # downsample blocks + downsamples = [] + for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): + # residual (+attention) blocks + for _ in range(num_res_blocks): + downsamples.append(ResidualBlock(in_dim, out_dim, dropout)) + if scale in attn_scales: + downsamples.append(AttentionBlock(out_dim)) + in_dim = out_dim + + # downsample block + if i != len(dim_mult) - 1: + mode = "downsample3d" if temperal_downsample[i] else "downsample2d" + downsamples.append(Resample(out_dim, mode=mode)) + scale /= 2.0 + self.downsamples = nn.Sequential(*downsamples) + + # middle blocks + self.middle = nn.Sequential( + ResidualBlock(out_dim, out_dim, dropout), + AttentionBlock(out_dim), + ResidualBlock(out_dim, out_dim, dropout), + ) + + # output blocks + self.head = nn.Sequential( + RMS_norm(out_dim, images=False), + nn.SiLU(), + CausalConv3d(out_dim, z_dim, 3, padding=1), + ) + + def forward(self, x, feat_cache=None, feat_idx=[0]): + if feat_cache is not None: + idx = feat_idx[0] + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and feat_cache[idx] is not None: + # cache last frame of last two chunk + cache_x = torch.cat( + [ + feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + x = self.conv1(x, feat_cache[idx]) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + else: + x = self.conv1(x) + + ## downsamples + for layer in self.downsamples: + if feat_cache is not None: + x = layer(x, feat_cache, feat_idx) + else: + x = layer(x) + + ## middle + for layer in self.middle: + if isinstance(layer, ResidualBlock) and feat_cache is not None: + x = layer(x, feat_cache, feat_idx) + else: + x = layer(x) + + ## head + for layer in self.head: + if isinstance(layer, CausalConv3d) and feat_cache is not None: + idx = feat_idx[0] + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and feat_cache[idx] is not None: + # cache last frame of last two chunk + cache_x = torch.cat( + [ + feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + x = layer(x, feat_cache[idx]) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + else: + x = layer(x) + return x + + +class Decoder3d(nn.Module): + def __init__(self, dim=128, z_dim=4, dim_mult=[1, 2, 4, 4], num_res_blocks=2, attn_scales=[], temperal_upsample=[False, True, True], dropout=0.0, pruning_rate=0.0): + super().__init__() + self.dim = dim + self.z_dim = z_dim + self.dim_mult = dim_mult + self.num_res_blocks = num_res_blocks + self.attn_scales = attn_scales + self.temperal_upsample = temperal_upsample + + # dimensions + dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] + dims = [int(d * (1 - pruning_rate)) for d in dims] + + scale = 1.0 / 2 ** (len(dim_mult) - 2) + + # init block + self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1) + + # middle blocks + self.middle = nn.Sequential( + ResidualBlock(dims[0], dims[0], dropout), + AttentionBlock(dims[0]), + ResidualBlock(dims[0], dims[0], dropout), + ) + + # upsample blocks + upsamples = [] + for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): + # residual (+attention) blocks + if i == 1 or i == 2 or i == 3: + in_dim = in_dim // 2 + for _ in range(num_res_blocks + 1): + upsamples.append(ResidualBlock(in_dim, out_dim, dropout)) + if scale in attn_scales: + upsamples.append(AttentionBlock(out_dim)) + in_dim = out_dim + + # upsample block + if i != len(dim_mult) - 1: + mode = "upsample3d" if temperal_upsample[i] else "upsample2d" + upsamples.append(Resample(out_dim, mode=mode)) + scale *= 2.0 + self.upsamples = nn.Sequential(*upsamples) + + # output blocks + self.head = nn.Sequential( + RMS_norm(out_dim, images=False), + nn.SiLU(), + CausalConv3d(out_dim, 3, 3, padding=1), + ) + + def forward(self, x, feat_cache=None, feat_idx=[0]): + ## conv1 + if feat_cache is not None: + idx = feat_idx[0] + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and feat_cache[idx] is not None: + # cache last frame of last two chunk + cache_x = torch.cat( + [ + feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + x = self.conv1(x, feat_cache[idx]) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + else: + x = self.conv1(x) + + ## middle + for layer in self.middle: + if isinstance(layer, ResidualBlock) and feat_cache is not None: + x = layer(x, feat_cache, feat_idx) + else: + x = layer(x) + + ## upsamples + for layer in self.upsamples: + if feat_cache is not None: + x = layer(x, feat_cache, feat_idx) + else: + x = layer(x) + + ## head + for layer in self.head: + if isinstance(layer, CausalConv3d) and feat_cache is not None: + idx = feat_idx[0] + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and feat_cache[idx] is not None: + # cache last frame of last two chunk + cache_x = torch.cat( + [ + feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + x = layer(x, feat_cache[idx]) + feat_cache[idx] = cache_x + feat_idx[0] += 1 + else: + x = layer(x) + return x + + +def count_conv3d(model): + count = 0 + for m in model.modules(): + if isinstance(m, CausalConv3d): + count += 1 + return count + + +class WanVAE_(nn.Module): + def __init__(self, dim=128, z_dim=4, dim_mult=[1, 2, 4, 4], num_res_blocks=2, attn_scales=[], temperal_downsample=[True, True, False], dropout=0.0, pruning_rate=0.0): + super().__init__() + self.dim = dim + self.z_dim = z_dim + self.dim_mult = dim_mult + self.num_res_blocks = num_res_blocks + self.attn_scales = attn_scales + self.temperal_downsample = temperal_downsample + self.temperal_upsample = temperal_downsample[::-1] + self.spatial_compression_ratio = 2 ** len(self.temperal_downsample) + + # The minimal tile height and width for spatial tiling to be used + self.tile_sample_min_height = 256 + self.tile_sample_min_width = 256 + + # The minimal distance between two spatial tiles + self.tile_sample_stride_height = 192 + self.tile_sample_stride_width = 192 + # modules + self.encoder = Encoder3d( + dim, + z_dim * 2, + dim_mult, + num_res_blocks, + attn_scales, + self.temperal_downsample, + dropout, + pruning_rate, + ) + self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1) + self.conv2 = CausalConv3d(z_dim, z_dim, 1) + self.decoder = Decoder3d( + dim, + z_dim, + dim_mult, + num_res_blocks, + attn_scales, + self.temperal_upsample, + dropout, + pruning_rate, + ) + + def forward(self, x): + mu, log_var = self.encode(x) + z = self.reparameterize(mu, log_var) + x_recon = self.decode(z) + return x_recon, mu, log_var + + def blend_v(self, a, b, blend_extent): + blend_extent = min(a.shape[-2], b.shape[-2], blend_extent) + for y in range(blend_extent): + b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (y / blend_extent) + return b + + def blend_h(self, a, b, blend_extent): + blend_extent = min(a.shape[-1], b.shape[-1], blend_extent) + for x in range(blend_extent): + b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (x / blend_extent) + return b + + def tiled_encode(self, x, scale): + _, _, num_frames, height, width = x.shape + latent_height = height // self.spatial_compression_ratio + latent_width = width // self.spatial_compression_ratio + + tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio + tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio + tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio + tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio + + blend_height = tile_latent_min_height - tile_latent_stride_height + blend_width = tile_latent_min_width - tile_latent_stride_width + + # Split x into overlapping tiles and encode them separately. + # The tiles have an overlap to avoid seams between tiles. + rows = [] + for i in range(0, height, self.tile_sample_stride_height): + row = [] + for j in range(0, width, self.tile_sample_stride_width): + self.clear_cache() + time = [] + frame_range = 1 + (num_frames - 1) // 4 + for k in range(frame_range): + self._enc_conv_idx = [0] + if k == 0: + tile = x[:, :, :1, i : i + self.tile_sample_min_height, j : j + self.tile_sample_min_width] + else: + tile = x[ + :, + :, + 1 + 4 * (k - 1) : 1 + 4 * k, + i : i + self.tile_sample_min_height, + j : j + self.tile_sample_min_width, + ] + tile = self.encoder(tile, feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) + mu, log_var = self.conv1(tile).chunk(2, dim=1) + if isinstance(scale[0], torch.Tensor): + mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(1, self.z_dim, 1, 1, 1) + else: + mu = (mu - scale[0]) * scale[1] + + time.append(mu) + + row.append(torch.cat(time, dim=2)) + rows.append(row) + self.clear_cache() + + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + # blend the above tile and the left tile + # to the current tile and add the current tile to the result row + if i > 0: + tile = self.blend_v(rows[i - 1][j], tile, blend_height) + if j > 0: + tile = self.blend_h(row[j - 1], tile, blend_width) + result_row.append(tile[:, :, :, :tile_latent_stride_height, :tile_latent_stride_width]) + result_rows.append(torch.cat(result_row, dim=-1)) + + enc = torch.cat(result_rows, dim=3)[:, :, :, :latent_height, :latent_width] + return enc + + def tiled_decode(self, z, scale): + if isinstance(scale[0], torch.Tensor): + z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(1, self.z_dim, 1, 1, 1) + else: + z = z / scale[1] + scale[0] + + _, _, num_frames, height, width = z.shape + sample_height = height * self.spatial_compression_ratio + sample_width = width * self.spatial_compression_ratio + + tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio + tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio + tile_latent_stride_height = self.tile_sample_stride_height // self.spatial_compression_ratio + tile_latent_stride_width = self.tile_sample_stride_width // self.spatial_compression_ratio + + blend_height = self.tile_sample_min_height - self.tile_sample_stride_height + blend_width = self.tile_sample_min_width - self.tile_sample_stride_width + + # Split z into overlapping tiles and decode them separately. + # The tiles have an overlap to avoid seams between tiles. + rows = [] + for i in range(0, height, tile_latent_stride_height): + row = [] + for j in range(0, width, tile_latent_stride_width): + self.clear_cache() + time = [] + for k in range(num_frames): + self._conv_idx = [0] + tile = z[:, :, k : k + 1, i : i + tile_latent_min_height, j : j + tile_latent_min_width] + tile = self.conv2(tile) + decoded = self.decoder(tile, feat_cache=self._feat_map, feat_idx=self._conv_idx) + time.append(decoded) + row.append(torch.cat(time, dim=2)) + rows.append(row) + self.clear_cache() + + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + # blend the above tile and the left tile + # to the current tile and add the current tile to the result row + if i > 0: + tile = self.blend_v(rows[i - 1][j], tile, blend_height) + if j > 0: + tile = self.blend_h(row[j - 1], tile, blend_width) + result_row.append(tile[:, :, :, : self.tile_sample_stride_height, : self.tile_sample_stride_width]) + result_rows.append(torch.cat(result_row, dim=-1)) + + dec = torch.cat(result_rows, dim=3)[:, :, :, :sample_height, :sample_width] + + return dec + + def encode(self, x, scale, return_mu=False): + self.clear_cache() + ## cache + t = x.shape[2] + iter_ = 1 + (t - 1) // 4 + for i in range(iter_): + self._enc_conv_idx = [0] + if i == 0: + out = self.encoder( + x[:, :, :1, :, :], + feat_cache=self._enc_feat_map, + feat_idx=self._enc_conv_idx, + ) + else: + out_ = self.encoder( + x[:, :, 1 + 4 * (i - 1) : 1 + 4 * i, :, :], + feat_cache=self._enc_feat_map, + feat_idx=self._enc_conv_idx, + ) + out = torch.cat([out, out_], 2) + mu, log_var = self.conv1(out).chunk(2, dim=1) + if isinstance(scale[0], torch.Tensor): + mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(1, self.z_dim, 1, 1, 1) + else: + mu = (mu - scale[0]) * scale[1] + + self.clear_cache() + if return_mu: + return mu, log_var + else: + return mu + + def decode(self, z, scale): + self.clear_cache() + + # z: [b,c,t,h,w] + if isinstance(scale[0], torch.Tensor): + z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(1, self.z_dim, 1, 1, 1) + else: + z = z / scale[1] + scale[0] + iter_ = z.shape[2] + x = self.conv2(z) + for i in range(iter_): + self._conv_idx = [0] + if i == 0: + out = self.decoder( + x[:, :, i : i + 1, :, :], + feat_cache=self._feat_map, + feat_idx=self._conv_idx, + ) + else: + out_ = self.decoder( + x[:, :, i : i + 1, :, :], + feat_cache=self._feat_map, + feat_idx=self._conv_idx, + ) + out = torch.cat([out, out_], 2) + + self.clear_cache() + return out + + def reparameterize(self, mu, log_var): + std = torch.exp(0.5 * log_var) + eps = torch.randn_like(std) + return eps * std + mu + + def sample(self, imgs, deterministic=False, scale=[0, 1]): + mu, log_var = self.encode(imgs, scale, return_mu=True) + if deterministic: + return mu + std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0)) + return mu + std * torch.randn_like(std), mu, log_var + + def clear_cache(self): + self._conv_num = count_conv3d(self.decoder) + self._conv_idx = [0] + self._feat_map = [None] * self._conv_num + # cache encode + self._enc_conv_num = count_conv3d(self.encoder) + self._enc_conv_idx = [0] + self._enc_feat_map = [None] * self._enc_conv_num + + def encode_video(self, x, scale=[0, 1]): + assert x.ndim == 5 # NTCHW + assert x.shape[2] % 3 == 0 + x = x.transpose(1, 2) + y = x.mul(2).sub_(1) + y, mu, log_var = self.sample(y, scale=scale) + return y.transpose(1, 2).to(x), mu, log_var + + def decode_video(self, x, scale=[0, 1]): + assert x.ndim == 5 # NTCHW + assert x.shape[2] % self.z_dim == 0 + x = x.transpose(1, 2) + # B, C, T, H, W + y = x + y = self.decode(y, scale).clamp_(-1, 1) + y = y.mul_(0.5).add_(0.5).clamp_(0, 1) # NCTHW + return y.transpose(1, 2).to(x) + + +def _video_vae(pretrained_path=None, z_dim=None, device="cpu", cpu_offload=False, dtype=torch.float, load_from_rank0=False, pruning_rate=0.0, **kwargs): + """ + Autoencoder3d adapted from Stable Diffusion 1.x, 2.x and XL. + """ + # params + cfg = dict( + dim=96, + z_dim=z_dim, + dim_mult=[1, 2, 4, 4], + num_res_blocks=2, + attn_scales=[], + temperal_downsample=[False, True, True], + dropout=0.0, + pruning_rate=pruning_rate, + ) + cfg.update(**kwargs) + + # init model + with torch.device("meta"): + model = WanVAE_(**cfg) + + # load checkpoint + weights_dict = load_weights(pretrained_path, cpu_offload=cpu_offload, load_from_rank0=load_from_rank0) + for k in weights_dict.keys(): + if weights_dict[k].dtype != dtype: + weights_dict[k] = weights_dict[k].to(dtype) + model.load_state_dict(weights_dict, assign=True) + + return model + + +class WanVAE: + def __init__( + self, + z_dim=16, + vae_path="cache/vae_step_411000.pth", + dtype=torch.float, + device="cuda", + parallel=False, + use_tiling=False, + cpu_offload=False, + use_2d_split=True, + load_from_rank0=False, + use_lightvae=False, + ): + self.dtype = dtype + self.device = device + self.parallel = parallel + self.use_tiling = use_tiling + self.cpu_offload = cpu_offload + self.use_2d_split = use_2d_split + if use_lightvae: + pruning_rate = 0.75 # 0.75 + else: + pruning_rate = 0.0 + + mean = [ + -0.7571, + -0.7089, + -0.9113, + 0.1075, + -0.1745, + 0.9653, + -0.1517, + 1.5508, + 0.4134, + -0.0715, + 0.5517, + -0.3632, + -0.1922, + -0.9497, + 0.2503, + -0.2921, + ] + std = [ + 2.8184, + 1.4541, + 2.3275, + 2.6558, + 1.2196, + 1.7708, + 2.6052, + 2.0743, + 3.2687, + 2.1526, + 2.8652, + 1.5579, + 1.6382, + 1.1253, + 2.8251, + 1.9160, + ] + self.mean = torch.tensor(mean, dtype=dtype, device=device) + self.inv_std = 1.0 / torch.tensor(std, dtype=dtype, device=device) + self.scale = [self.mean, self.inv_std] + + # (height, width, world_size) -> (world_size_h, world_size_w) + self.grid_table = { + # world_size = 2 + (60, 104, 2): (1, 2), + (68, 120, 2): (1, 2), + (90, 160, 2): (1, 2), + (60, 60, 2): (1, 2), + (72, 72, 2): (1, 2), + (88, 88, 2): (1, 2), + (120, 120, 2): (1, 2), + (104, 60, 2): (2, 1), + (120, 68, 2): (2, 1), + (160, 90, 2): (2, 1), + # world_size = 4 + (60, 104, 4): (2, 2), + (68, 120, 4): (2, 2), + (90, 160, 4): (2, 2), + (60, 60, 4): (2, 2), + (72, 72, 4): (2, 2), + (88, 88, 4): (2, 2), + (120, 120, 4): (2, 2), + (104, 60, 4): (2, 2), + (120, 68, 4): (2, 2), + (160, 90, 4): (2, 2), + # world_size = 8 + (60, 104, 8): (2, 4), + (68, 120, 8): (2, 4), + (90, 160, 8): (2, 4), + (60, 60, 8): (2, 4), + (72, 72, 8): (2, 4), + (88, 88, 8): (2, 4), + (120, 120, 8): (2, 4), + (104, 60, 8): (4, 2), + (120, 68, 8): (4, 2), + (160, 90, 8): (4, 2), + } + + # init model + self.model = ( + _video_vae(pretrained_path=vae_path, z_dim=z_dim, cpu_offload=cpu_offload, dtype=dtype, load_from_rank0=load_from_rank0, pruning_rate=pruning_rate) + .eval() + .requires_grad_(False) + .to(device) + .to(dtype) + ) + + def _calculate_2d_grid(self, latent_height, latent_width, world_size): + if (latent_height, latent_width, world_size) in self.grid_table: + best_h, best_w = self.grid_table[(latent_height, latent_width, world_size)] + # logger.info(f"Vae using cached 2D grid: {best_h}x{best_w} grid for {latent_height}x{latent_width} latent") + return best_h, best_w + + best_h, best_w = 1, world_size + min_aspect_diff = float("inf") + + for h in range(1, world_size + 1): + if world_size % h == 0: + w = world_size // h + if latent_height % h == 0 and latent_width % w == 0: + # Calculate how close this grid is to square + aspect_diff = abs((latent_height / h) - (latent_width / w)) + if aspect_diff < min_aspect_diff: + min_aspect_diff = aspect_diff + best_h, best_w = h, w + # logger.info(f"Vae using 2D grid & Update cache: {best_h}x{best_w} grid for {latent_height}x{latent_width} latent") + self.grid_table[(latent_height, latent_width, world_size)] = (best_h, best_w) + return best_h, best_w + + def current_device(self): + return next(self.model.parameters()).device + + def to_cpu(self): + self.model.encoder = self.model.encoder.to("cpu") + self.model.decoder = self.model.decoder.to("cpu") + self.model = self.model.to("cpu") + self.mean = self.mean.cpu() + self.inv_std = self.inv_std.cpu() + self.scale = [self.mean, self.inv_std] + + def to_cuda(self): + self.model.encoder = self.model.encoder.to("cuda") + self.model.decoder = self.model.decoder.to("cuda") + self.model = self.model.to("cuda") + self.mean = self.mean.cuda() + self.inv_std = self.inv_std.cuda() + self.scale = [self.mean, self.inv_std] + + def encode_dist(self, video, world_size, cur_rank, split_dim): + spatial_ratio = 8 + + if split_dim == 3: + total_latent_len = video.shape[3] // spatial_ratio + elif split_dim == 4: + total_latent_len = video.shape[4] // spatial_ratio + else: + raise ValueError(f"Unsupported split_dim: {split_dim}") + + splited_chunk_len = total_latent_len // world_size + padding_size = 1 + + video_chunk_len = splited_chunk_len * spatial_ratio + video_padding_len = padding_size * spatial_ratio + + if cur_rank == 0: + if split_dim == 3: + video_chunk = video[:, :, :, : video_chunk_len + 2 * video_padding_len, :].contiguous() + elif split_dim == 4: + video_chunk = video[:, :, :, :, : video_chunk_len + 2 * video_padding_len].contiguous() + elif cur_rank == world_size - 1: + if split_dim == 3: + video_chunk = video[:, :, :, -(video_chunk_len + 2 * video_padding_len) :, :].contiguous() + elif split_dim == 4: + video_chunk = video[:, :, :, :, -(video_chunk_len + 2 * video_padding_len) :].contiguous() + else: + start_idx = cur_rank * video_chunk_len - video_padding_len + end_idx = (cur_rank + 1) * video_chunk_len + video_padding_len + if split_dim == 3: + video_chunk = video[:, :, :, start_idx:end_idx, :].contiguous() + elif split_dim == 4: + video_chunk = video[:, :, :, :, start_idx:end_idx].contiguous() + + if self.use_tiling: + encoded_chunk = self.model.tiled_encode(video_chunk, self.scale) + else: + encoded_chunk = self.model.encode(video_chunk, self.scale) + + if cur_rank == 0: + if split_dim == 3: + encoded_chunk = encoded_chunk[:, :, :, :splited_chunk_len, :].contiguous() + elif split_dim == 4: + encoded_chunk = encoded_chunk[:, :, :, :, :splited_chunk_len].contiguous() + elif cur_rank == world_size - 1: + if split_dim == 3: + encoded_chunk = encoded_chunk[:, :, :, -splited_chunk_len:, :].contiguous() + elif split_dim == 4: + encoded_chunk = encoded_chunk[:, :, :, :, -splited_chunk_len:].contiguous() + else: + if split_dim == 3: + encoded_chunk = encoded_chunk[:, :, :, padding_size:-padding_size, :].contiguous() + elif split_dim == 4: + encoded_chunk = encoded_chunk[:, :, :, :, padding_size:-padding_size].contiguous() + + full_encoded = [torch.empty_like(encoded_chunk) for _ in range(world_size)] + dist.all_gather(full_encoded, encoded_chunk) + + torch.cuda.synchronize() + + encoded = torch.cat(full_encoded, dim=split_dim) + + return encoded.squeeze(0) + + def encode_dist_2d(self, video, world_size_h, world_size_w, cur_rank_h, cur_rank_w): + spatial_ratio = 8 + + # Calculate chunk sizes for both dimensions + total_latent_h = video.shape[3] // spatial_ratio + total_latent_w = video.shape[4] // spatial_ratio + + chunk_h = total_latent_h // world_size_h + chunk_w = total_latent_w // world_size_w + + padding_size = 1 + video_chunk_h = chunk_h * spatial_ratio + video_chunk_w = chunk_w * spatial_ratio + video_padding_h = padding_size * spatial_ratio + video_padding_w = padding_size * spatial_ratio + + # Calculate H dimension slice + if cur_rank_h == 0: + h_start = 0 + h_end = video_chunk_h + 2 * video_padding_h + elif cur_rank_h == world_size_h - 1: + h_start = video.shape[3] - (video_chunk_h + 2 * video_padding_h) + h_end = video.shape[3] + else: + h_start = cur_rank_h * video_chunk_h - video_padding_h + h_end = (cur_rank_h + 1) * video_chunk_h + video_padding_h + + # Calculate W dimension slice + if cur_rank_w == 0: + w_start = 0 + w_end = video_chunk_w + 2 * video_padding_w + elif cur_rank_w == world_size_w - 1: + w_start = video.shape[4] - (video_chunk_w + 2 * video_padding_w) + w_end = video.shape[4] + else: + w_start = cur_rank_w * video_chunk_w - video_padding_w + w_end = (cur_rank_w + 1) * video_chunk_w + video_padding_w + + # Extract the video chunk for this process + video_chunk = video[:, :, :, h_start:h_end, w_start:w_end].contiguous() + + # Encode the chunk + if self.use_tiling: + encoded_chunk = self.model.tiled_encode(video_chunk, self.scale) + else: + encoded_chunk = self.model.encode(video_chunk, self.scale) + + # Remove padding from encoded chunk + if cur_rank_h == 0: + encoded_h_start = 0 + encoded_h_end = chunk_h + elif cur_rank_h == world_size_h - 1: + encoded_h_start = encoded_chunk.shape[3] - chunk_h + encoded_h_end = encoded_chunk.shape[3] + else: + encoded_h_start = padding_size + encoded_h_end = encoded_chunk.shape[3] - padding_size + + if cur_rank_w == 0: + encoded_w_start = 0 + encoded_w_end = chunk_w + elif cur_rank_w == world_size_w - 1: + encoded_w_start = encoded_chunk.shape[4] - chunk_w + encoded_w_end = encoded_chunk.shape[4] + else: + encoded_w_start = padding_size + encoded_w_end = encoded_chunk.shape[4] - padding_size + + encoded_chunk = encoded_chunk[:, :, :, encoded_h_start:encoded_h_end, encoded_w_start:encoded_w_end].contiguous() + + # Gather all chunks + total_processes = world_size_h * world_size_w + full_encoded = [torch.empty_like(encoded_chunk) for _ in range(total_processes)] + + dist.all_gather(full_encoded, encoded_chunk) + + torch.cuda.synchronize() + + # Reconstruct the full encoded tensor + encoded_rows = [] + for h_idx in range(world_size_h): + encoded_cols = [] + for w_idx in range(world_size_w): + process_idx = h_idx * world_size_w + w_idx + encoded_cols.append(full_encoded[process_idx]) + encoded_rows.append(torch.cat(encoded_cols, dim=4)) + + encoded = torch.cat(encoded_rows, dim=3) + + return encoded.squeeze(0) + + def encode(self, video): + """ + video: one video with shape [1, C, T, H, W]. + """ + if self.cpu_offload: + self.to_cuda() + + if self.parallel: + world_size = dist.get_world_size() + cur_rank = dist.get_rank() + height, width = video.shape[3], video.shape[4] + + if self.use_2d_split: + world_size_h, world_size_w = self._calculate_2d_grid(height // 8, width // 8, world_size) + cur_rank_h = cur_rank // world_size_w + cur_rank_w = cur_rank % world_size_w + out = self.encode_dist_2d(video, world_size_h, world_size_w, cur_rank_h, cur_rank_w) + else: + # Original 1D splitting logic + if width % world_size == 0: + out = self.encode_dist(video, world_size, cur_rank, split_dim=4) + elif height % world_size == 0: + out = self.encode_dist(video, world_size, cur_rank, split_dim=3) + else: + logger.info("Fall back to naive encode mode") + if self.use_tiling: + out = self.model.tiled_encode(video, self.scale).squeeze(0) + else: + out = self.model.encode(video, self.scale).squeeze(0) + else: + if self.use_tiling: + out = self.model.tiled_encode(video, self.scale).squeeze(0) + else: + out = self.model.encode(video, self.scale).squeeze(0) + + if self.cpu_offload: + self.to_cpu() + return out + + def decode_dist(self, zs, world_size, cur_rank, split_dim): + splited_total_len = zs.shape[split_dim] + splited_chunk_len = splited_total_len // world_size + padding_size = 1 + + if cur_rank == 0: + if split_dim == 2: + zs = zs[:, :, : splited_chunk_len + 2 * padding_size, :].contiguous() + elif split_dim == 3: + zs = zs[:, :, :, : splited_chunk_len + 2 * padding_size].contiguous() + elif cur_rank == world_size - 1: + if split_dim == 2: + zs = zs[:, :, -(splited_chunk_len + 2 * padding_size) :, :].contiguous() + elif split_dim == 3: + zs = zs[:, :, :, -(splited_chunk_len + 2 * padding_size) :].contiguous() + else: + if split_dim == 2: + zs = zs[:, :, cur_rank * splited_chunk_len - padding_size : (cur_rank + 1) * splited_chunk_len + padding_size, :].contiguous() + elif split_dim == 3: + zs = zs[:, :, :, cur_rank * splited_chunk_len - padding_size : (cur_rank + 1) * splited_chunk_len + padding_size].contiguous() + + decode_func = self.model.tiled_decode if self.use_tiling else self.model.decode + images = decode_func(zs.unsqueeze(0), self.scale).clamp_(-1, 1) + + if cur_rank == 0: + if split_dim == 2: + images = images[:, :, :, : splited_chunk_len * 8, :].contiguous() + elif split_dim == 3: + images = images[:, :, :, :, : splited_chunk_len * 8].contiguous() + elif cur_rank == world_size - 1: + if split_dim == 2: + images = images[:, :, :, -splited_chunk_len * 8 :, :].contiguous() + elif split_dim == 3: + images = images[:, :, :, :, -splited_chunk_len * 8 :].contiguous() + else: + if split_dim == 2: + images = images[:, :, :, 8 * padding_size : -8 * padding_size, :].contiguous() + elif split_dim == 3: + images = images[:, :, :, :, 8 * padding_size : -8 * padding_size].contiguous() + + full_images = [torch.empty_like(images) for _ in range(world_size)] + dist.all_gather(full_images, images) + + torch.cuda.synchronize() + + images = torch.cat(full_images, dim=split_dim + 1) + + return images + + def decode_dist_2d(self, zs, world_size_h, world_size_w, cur_rank_h, cur_rank_w): + total_h = zs.shape[2] + total_w = zs.shape[3] + + chunk_h = total_h // world_size_h + chunk_w = total_w // world_size_w + + padding_size = 1 + + # Calculate H dimension slice + if cur_rank_h == 0: + h_start = 0 + h_end = chunk_h + 2 * padding_size + elif cur_rank_h == world_size_h - 1: + h_start = total_h - (chunk_h + 2 * padding_size) + h_end = total_h + else: + h_start = cur_rank_h * chunk_h - padding_size + h_end = (cur_rank_h + 1) * chunk_h + padding_size + + # Calculate W dimension slice + if cur_rank_w == 0: + w_start = 0 + w_end = chunk_w + 2 * padding_size + elif cur_rank_w == world_size_w - 1: + w_start = total_w - (chunk_w + 2 * padding_size) + w_end = total_w + else: + w_start = cur_rank_w * chunk_w - padding_size + w_end = (cur_rank_w + 1) * chunk_w + padding_size + + # Extract the latent chunk for this process + zs_chunk = zs[:, :, h_start:h_end, w_start:w_end].contiguous() + + # Decode the chunk + decode_func = self.model.tiled_decode if self.use_tiling else self.model.decode + images_chunk = decode_func(zs_chunk.unsqueeze(0), self.scale).clamp_(-1, 1) + + # Remove padding from decoded chunk + spatial_ratio = 8 + if cur_rank_h == 0: + decoded_h_start = 0 + decoded_h_end = chunk_h * spatial_ratio + elif cur_rank_h == world_size_h - 1: + decoded_h_start = images_chunk.shape[3] - chunk_h * spatial_ratio + decoded_h_end = images_chunk.shape[3] + else: + decoded_h_start = padding_size * spatial_ratio + decoded_h_end = images_chunk.shape[3] - padding_size * spatial_ratio + + if cur_rank_w == 0: + decoded_w_start = 0 + decoded_w_end = chunk_w * spatial_ratio + elif cur_rank_w == world_size_w - 1: + decoded_w_start = images_chunk.shape[4] - chunk_w * spatial_ratio + decoded_w_end = images_chunk.shape[4] + else: + decoded_w_start = padding_size * spatial_ratio + decoded_w_end = images_chunk.shape[4] - padding_size * spatial_ratio + + images_chunk = images_chunk[:, :, :, decoded_h_start:decoded_h_end, decoded_w_start:decoded_w_end].contiguous() + + # Gather all chunks + total_processes = world_size_h * world_size_w + full_images = [torch.empty_like(images_chunk) for _ in range(total_processes)] + + dist.all_gather(full_images, images_chunk) + + torch.cuda.synchronize() + + # Reconstruct the full image tensor + image_rows = [] + for h_idx in range(world_size_h): + image_cols = [] + for w_idx in range(world_size_w): + process_idx = h_idx * world_size_w + w_idx + image_cols.append(full_images[process_idx]) + image_rows.append(torch.cat(image_cols, dim=4)) + + images = torch.cat(image_rows, dim=3) + + return images + + def decode(self, zs): + if self.cpu_offload: + self.to_cuda() + + if self.parallel: + world_size = dist.get_world_size() + cur_rank = dist.get_rank() + latent_height, latent_width = zs.shape[2], zs.shape[3] + + if self.use_2d_split: + world_size_h, world_size_w = self._calculate_2d_grid(latent_height, latent_width, world_size) + cur_rank_h = cur_rank // world_size_w + cur_rank_w = cur_rank % world_size_w + images = self.decode_dist_2d(zs, world_size_h, world_size_w, cur_rank_h, cur_rank_w) + else: + # Original 1D splitting logic + if latent_width % world_size == 0: + images = self.decode_dist(zs, world_size, cur_rank, split_dim=3) + elif latent_height % world_size == 0: + images = self.decode_dist(zs, world_size, cur_rank, split_dim=2) + else: + logger.info("Fall back to naive decode mode") + images = self.model.decode(zs.unsqueeze(0), self.scale).clamp_(-1, 1) + else: + decode_func = self.model.tiled_decode if self.use_tiling else self.model.decode + images = decode_func(zs.unsqueeze(0), self.scale).clamp_(-1, 1) + + if self.cpu_offload: + images = images.cpu() + self.to_cpu() + + return images + + def encode_video(self, vid): + return self.model.encode_video(vid) + + def decode_video(self, vid_enc): + return self.model.decode_video(vid_enc) diff --git a/Chrono/vae_tiny.py b/Chrono/vae_tiny.py new file mode 100644 index 0000000..1fb59a4 --- /dev/null +++ b/Chrono/vae_tiny.py @@ -0,0 +1,216 @@ +import torch +import torch.nn as nn + +from .tae import TAEHV +from .memory_profiler import peak_memory_decorator + + +class DotDict(dict): + __getattr__ = dict.__getitem__ + __setattr__ = dict.__setitem__ + + +class WanVAE_tiny(nn.Module): + def __init__(self, vae_path="taew2_1.pth", dtype=torch.bfloat16, device="cuda", need_scaled=False): + super().__init__() + self.dtype = dtype + self.device = torch.device("cuda") + self.taehv = TAEHV(vae_path).to(self.dtype) + self.temperal_downsample = [True, True, False] + self.need_scaled = need_scaled + + if self.need_scaled: + self.latents_mean = [ + -0.7571, + -0.7089, + -0.9113, + 0.1075, + -0.1745, + 0.9653, + -0.1517, + 1.5508, + 0.4134, + -0.0715, + 0.5517, + -0.3632, + -0.1922, + -0.9497, + 0.2503, + -0.2921, + ] + + self.latents_std = [ + 2.8184, + 1.4541, + 2.3275, + 2.6558, + 1.2196, + 1.7708, + 2.6052, + 2.0743, + 3.2687, + 2.1526, + 2.8652, + 1.5579, + 1.6382, + 1.1253, + 2.8251, + 1.9160, + ] + + self.z_dim = 16 + + @peak_memory_decorator + @torch.no_grad() + def decode(self, latents): + latents = latents.unsqueeze(0) + + if self.need_scaled: + latents_mean = torch.tensor(self.latents_mean).view(1, self.z_dim, 1, 1, 1).to(latents.device, latents.dtype) + latents_std = 1.0 / torch.tensor(self.latents_std).view(1, self.z_dim, 1, 1, 1).to(latents.device, latents.dtype) + latents = latents / latents_std + latents_mean + + # low-memory, set parallel=True for faster + higher memory + return self.taehv.decode_video(latents.transpose(1, 2).to(self.dtype), parallel=False).transpose(1, 2).mul_(2).sub_(1) + + @torch.no_grad() + def encode_video(self, vid): + return self.taehv.encode_video(vid) + + @torch.no_grad() + def decode_video(self, vid_enc): + return self.taehv.decode_video(vid_enc) + + +class Wan2_2_VAE_tiny(nn.Module): + def __init__(self, vae_path="taew2_2.pth", dtype=torch.bfloat16, device="cuda", need_scaled=False): + super().__init__() + self.dtype = dtype + self.device = torch.device("cuda") + self.taehv = TAEHV(vae_path, model_type="wan22").to(self.dtype) + self.need_scaled = need_scaled + if self.need_scaled: + self.latents_mean = [ + -0.2289, + -0.0052, + -0.1323, + -0.2339, + -0.2799, + 0.0174, + 0.1838, + 0.1557, + -0.1382, + 0.0542, + 0.2813, + 0.0891, + 0.1570, + -0.0098, + 0.0375, + -0.1825, + -0.2246, + -0.1207, + -0.0698, + 0.5109, + 0.2665, + -0.2108, + -0.2158, + 0.2502, + -0.2055, + -0.0322, + 0.1109, + 0.1567, + -0.0729, + 0.0899, + -0.2799, + -0.1230, + -0.0313, + -0.1649, + 0.0117, + 0.0723, + -0.2839, + -0.2083, + -0.0520, + 0.3748, + 0.0152, + 0.1957, + 0.1433, + -0.2944, + 0.3573, + -0.0548, + -0.1681, + -0.0667, + ] + + self.latents_std = [ + 0.4765, + 1.0364, + 0.4514, + 1.1677, + 0.5313, + 0.4990, + 0.4818, + 0.5013, + 0.8158, + 1.0344, + 0.5894, + 1.0901, + 0.6885, + 0.6165, + 0.8454, + 0.4978, + 0.5759, + 0.3523, + 0.7135, + 0.6804, + 0.5833, + 1.4146, + 0.8986, + 0.5659, + 0.7069, + 0.5338, + 0.4889, + 0.4917, + 0.4069, + 0.4999, + 0.6866, + 0.4093, + 0.5709, + 0.6065, + 0.6415, + 0.4944, + 0.5726, + 1.2042, + 0.5458, + 1.6887, + 0.3971, + 1.0600, + 0.3943, + 0.5537, + 0.5444, + 0.4089, + 0.7468, + 0.7744, + ] + + self.z_dim = 48 + + @peak_memory_decorator + @torch.no_grad() + def decode(self, latents): + latents = latents.unsqueeze(0) + + if self.need_scaled: + latents_mean = torch.tensor(self.latents_mean).view(1, self.z_dim, 1, 1, 1).to(latents.device, latents.dtype) + latents_std = 1.0 / torch.tensor(self.latents_std).view(1, self.z_dim, 1, 1, 1).to(latents.device, latents.dtype) + latents = latents / latents_std + latents_mean + + # low-memory, set parallel=True for faster + higher memory + return self.taehv.decode_video(latents.transpose(1, 2).to(self.dtype), parallel=False).transpose(1, 2).mul_(2).sub_(1) + + @torch.no_grad() + def encode_video(self, vid): + return self.taehv.encode_video(vid) + + @torch.no_grad() + def decode_video(self, vid_enc): + return self.taehv.decode_video(vid_enc)