This commit is contained in:
smthemex
2025-11-13 19:10:23 +08:00
committed by GitHub
parent 62dddad89e
commit 06831dde8f
5 changed files with 2335 additions and 0 deletions
+29
View File
@@ -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
+296
View File
@@ -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 :]
+485
View File
@@ -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
+1309
View File
File diff suppressed because it is too large Load Diff
+216
View File
@@ -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)