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