VAE cleanup and adjustments, add node to allow using native VAE

This commit is contained in:
kijai
2025-08-06 17:28:50 +03:00
parent 69dd4689bf
commit 2a21c4f6ae
3 changed files with 110 additions and 124 deletions
+48 -44
View File
@@ -3214,8 +3214,6 @@ class WanVideoDecode:
if drop_last:
latents = latents[:, :, :-1]
#if is_looped:
# latents = torch.cat([latents[:, :, :warmup_latent_count],latents], dim=2)
if type(vae).__name__ == "TAEHV":
images = vae.decode_video(latents.permute(0, 2, 1, 3, 4))[0].permute(1, 0, 2, 3)
images = torch.clamp(images, 0.0, 1.0)
@@ -3227,34 +3225,30 @@ class WanVideoDecode:
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
vae.model.clear_cache()
images = images.cpu()
images = images.cpu().float()
if normalization == "minmax":
images = (images - images.min()) / (images.max() - images.min())
images.sub_(images.min()).div_(images.max() - images.min())
else:
images = torch.clamp(images, -1.0, 1.0)
images = (images + 1.0) / 2.0
images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if is_looped:
#images = images[:, warmup_latent_count * 4:]
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2)
temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0]
temp_images = (temp_images - temp_images.min()) / (temp_images.max() - temp_images.min())
images = torch.cat([temp_images[:, 9:].to(images), images[:, 5:]], dim=1)
if end_image is not None:
#end_image = (end_image - end_image.min()) / (end_image.max() - end_image.min())
#image[:, -1] = end_image[:, 0].to(image) #not sure about this
images = images[:, 0:-1]
vae.model.clear_cache()
vae.to(offload_device)
mm.soft_empty_cache()
images = torch.clamp(images, 0.0, 1.0)
images = images.permute(1, 2, 3, 0).float()
images.clamp_(0.0, 1.0)
return (images,)
return (images.permute(1, 2, 3, 0),)
#region VideoEncode
class WanVideoEncode:
@@ -3295,35 +3289,7 @@ class WanVideoEncode:
if image.shape[-1] == 4:
image = image[..., :3]
image = image.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W
# empty_frame_indices = []
# for i in range(image.shape[2]):
# if is_image_black(image[:, :, i]):
# empty_frame_indices.append(i)
# empty_frame_indices = []
# for i in range(image.shape[2]):
# if is_image_black(image[:, :, i]):
# empty_frame_indices.append(i)
# empty_latent_indices = []
# if empty_frame_indices:
# frames_per_latent = 4
# num_frames = image.shape[2]
# # Special mapping: latent 0 = [0], latent 1 = [1,2,3,4], latent 2 = [5,6,7,8], ...
# latent_frame_ranges = []
# latent_frame_ranges.append([0])
# for i in range(1, math.ceil((num_frames - 1) / frames_per_latent) + 1):
# start = 1 + (i - 1) * frames_per_latent
# end = min(start + frames_per_latent, num_frames)
# latent_frame_ranges.append(list(range(start, end)))
# for latent_idx, latent_frames in enumerate(latent_frame_ranges):
# print(f"latent {latent_idx}: frames {latent_frames}")
# if latent_frames and set(latent_frames).issubset(empty_frame_indices):
# empty_latent_indices.append(latent_idx)
# if empty_latent_indices:
# log.info(f"Empty frames {empty_frame_indices} map to latents {empty_latent_indices}")
image = image.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W
if noise_aug_strength > 0.0:
image = add_noise_to_reference_video(image, ratio=noise_aug_strength)
@@ -3342,9 +3308,6 @@ class WanVideoEncode:
if mask is None:
vae.to(offload_device)
else:
#latent_mask = mask.clone().to(vae.dtype).to(device) * 2.0 - 1.0
#latent_mask = latent_mask.unsqueeze(0).unsqueeze(0).repeat(1, 3, 1, 1, 1)
#latent_mask = vae.encode(latent_mask, device=device, tiled=enable_vae_tiling, tile_size=(tile_x, tile_y), tile_stride=(tile_stride_x, tile_stride_y))
target_h, target_w = latents.shape[3:]
mask = torch.nn.functional.interpolate(
@@ -3362,6 +3325,45 @@ class WanVideoEncode:
return ({"samples": latents, "mask": latent_mask},)
class WanVideoLatentReScale:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"samples": ("LATENT",),
"direction": (["comfy_to_wrapper", "wrapper_to_comfy"], {"tooltip": "Direction to rescale latents, from comfy to wrapper or vice versa"}),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "encode"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding. Can be used to "
def encode(self, samples, direction):
samples = samples.copy()
latents = samples["samples"]
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
]
mean = torch.tensor(mean).view(1, latents.shape[1], 1, 1, 1)
std = torch.tensor(std).view(1, latents.shape[1], 1, 1, 1)
inv_std = (1.0 / std).view(1, latents.shape[1], 1, 1, 1)
if direction == "comfy_to_wrapper":
latents = (latents - mean.to(latents)) * inv_std.to(latents)
elif direction == "wrapper_to_comfy":
latents = latents / inv_std.to(latents) + mean.to(latents)
samples["samples"] = latents
return (samples,)
NODE_CLASS_MAPPINGS = {
"WanVideoSampler": WanVideoSampler,
"WanVideoDecode": WanVideoDecode,
@@ -3390,6 +3392,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoBlockList": WanVideoBlockList,
"WanVideoTextEncodeCached": WanVideoTextEncodeCached,
"WanVideoAddExtraLatent": WanVideoAddExtraLatent,
"WanVideoLatentReScale": WanVideoLatentReScale,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -3420,4 +3423,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoBlockList": "WanVideo Block List",
"WanVideoTextEncodeCached": "WanVideo TextEncode Cached",
"WanVideoAddExtraLatent": "WanVideo Add Extra Latent",
"WanVideoLatentReScale": "WanVideo Latent ReScale",
}
+10 -10
View File
@@ -1,18 +1,21 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from ...utils import log
# Flash Attention imports
try:
import flash_attn_interface
FLASH_ATTN_3_AVAILABLE = True
except ModuleNotFoundError:
except Exception as e:
FLASH_ATTN_3_AVAILABLE = False
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except ModuleNotFoundError:
except Exception as e:
FLASH_ATTN_2_AVAILABLE = False
# Sage Attention imports
try:
from sageattention import sageattn
@torch.compiler.disable()
@@ -22,11 +25,11 @@ try:
else:
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
except Exception as e:
print(f"Warning: Could not load sageattention: {str(e)}")
log.warning(f"Warning: Could not load sageattention: {str(e)}")
if isinstance(e, ModuleNotFoundError):
print("sageattention package is not installed")
log.warning("sageattention package is not installed, sageattention will not be available")
elif isinstance(e, ImportError) and "DLL" in str(e):
print("sageattention DLL loading error")
log.warning("sageattention DLL loading error, sageattention will not be available")
sageattn_func = None
try:
@@ -34,7 +37,6 @@ try:
except:
SAGE3_AVAILABLE = False
import warnings
__all__ = [
'flash_attention',
@@ -107,9 +109,7 @@ def flash_attention(
q = q * q_scale
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
warnings.warn(
'Flash attention 3 is not available, use flash attention 2 instead.'
)
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
# apply attention
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
+52 -70
View File
@@ -958,7 +958,9 @@ class VideoVAE_(nn.Module):
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0,):
dropout=0.0,
mean=None,
inv_std=None):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -967,6 +969,8 @@ class VideoVAE_(nn.Module):
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
self.temperal_upsample = temperal_downsample[::-1]
self.mean = mean
self.inv_std = inv_std
# modules
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
@@ -984,7 +988,7 @@ class VideoVAE_(nn.Module):
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
def encode_2(self, x, scale):
def encode_2(self, x):
self.clear_cache()
## cache
t = x.shape[2]
@@ -1008,18 +1012,14 @@ class VideoVAE_(nn.Module):
out = torch.cat([out, out_], 2)
out_head = out[:, :, :iter_ - 1, :, :]
out_tail = out[:, :, -1, :, :].unsqueeze(2)
mu, log_var = torch.cat([self.conv1(out_head), self.conv1(out_tail)], dim=2).chunk(2, dim=1)
if isinstance(scale[0], torch.Tensor):
scale = [s.to(dtype=mu.dtype, device=mu.device) for s in scale]
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
scale = scale.to(dtype=mu.dtype, device=mu.device)
mu = (mu - scale[0]) * scale[1]
mu = torch.cat([self.conv1(out_head), self.conv1(out_tail)], dim=2).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
return mu
def encode(self, x, scale):
def encode(self, x):
self.clear_cache()
## cache
t = x.shape[2]
@@ -1036,28 +1036,20 @@ class VideoVAE_(nn.Module):
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):
scale = [s.to(dtype=mu.dtype, device=mu.device) for s in scale]
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
scale = scale.to(dtype=mu.dtype, device=mu.device)
mu = (mu - scale[0]) * scale[1]
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
return mu
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
def decode_2(self, z, scale):
def decode_2(self, z):
self.clear_cache()
# z: [b,c,t,h,w]
if isinstance(scale[0], torch.Tensor):
scale = [s.to(dtype=z.dtype, device=z.device) for s in scale]
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
1, self.z_dim, 1, 1, 1)
else:
scale = scale.to(dtype=z.dtype, device=z.device)
z = z / scale[1] + scale[0]
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
z_head=z[:,:,:-1,:,:]
z_tail=z[:,:,-1,:,:].unsqueeze(2)
@@ -1082,17 +1074,13 @@ class VideoVAE_(nn.Module):
def decode(self, z, scale):
def decode(self, z):
self.clear_cache()
# z: [b,c,t,h,w]
pbar = ProgressBar(z.shape[2])
if isinstance(scale[0], torch.Tensor):
scale = [s.to(dtype=z.dtype, device=z.device) for s in scale]
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
1, self.z_dim, 1, 1, 1)
else:
scale = scale.to(dtype=z.dtype, device=z.device)
z = z / scale[1] + scale[0]
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
x = self.conv2(z)
for i in range(iter_):
@@ -1146,13 +1134,12 @@ class WanVideoVAE(nn.Module):
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)
self.std = torch.tensor(std)
self.scale = [self.mean, 1.0 / self.std]
self.mean = torch.tensor(mean).view(1, z_dim, 1, 1, 1)
self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1)
self.z_dim = z_dim
# init model
self.model = VideoVAE_(z_dim=z_dim).eval().requires_grad_(False)
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std).eval().requires_grad_(False)
self.upsampling_factor = 8
@@ -1202,7 +1189,7 @@ class WanVideoVAE(nn.Module):
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
hidden_states_batch = self.model.decode(hidden_states_batch, self.scale).to(data_device)
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
mask = self.build_mask(
hidden_states_batch,
@@ -1264,9 +1251,9 @@ class WanVideoVAE(nn.Module):
for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):
hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device)
if end_:
hidden_states_batch = self.model.encode_2(hidden_states_batch, self.scale).to(data_device)
hidden_states_batch = self.model.encode_2(hidden_states_batch).to(data_device)
else:
hidden_states_batch = self.model.encode(hidden_states_batch, self.scale).to(data_device)
hidden_states_batch = self.model.encode(hidden_states_batch).to(data_device)
mask = self.build_mask(
hidden_states_batch,
@@ -1298,26 +1285,26 @@ class WanVideoVAE(nn.Module):
def single_encode(self, video, device):
video = video.to(device)
x = self.model.encode(video, self.scale)
x = self.model.encode(video)
return x.float()
def single_decode(self, hidden_state, device):
hidden_state = hidden_state.to(device)
video = self.model.decode(hidden_state, self.scale)
return video.float().clamp_(-1, 1)
video = self.model.decode(hidden_state)
return video
def double_encode(self, video, device):
print('double_encode')
video = video.to(device)
x = self.model.encode_2(video, self.scale)
x = self.model.encode_2(video)
return x.float()
def double_decode(self, hidden_state, device):
print('double_decode')
hidden_state = hidden_state.to(device)
video = self.model.decode_2(hidden_state, self.scale)
return video.float().clamp_(-1, 1)
video = self.model.decode_2(hidden_state)
return video
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None):
videos = [video.to("cpu") for video in videos]
@@ -1384,7 +1371,9 @@ class VideoVAE38_(VideoVAE_):
attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0,
dtype=torch.bfloat16):
dtype=torch.bfloat16,
mean=None,
inv_std=None):
super(VideoVAE_, self).__init__()
self.dim = dim
self.z_dim = z_dim
@@ -1394,6 +1383,8 @@ class VideoVAE38_(VideoVAE_):
self.temperal_downsample = temperal_downsample
self.temperal_upsample = temperal_downsample[::-1]
self.dtype = dtype
self.mean = mean
self.inv_std = inv_std
# modules
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
@@ -1404,7 +1395,7 @@ class VideoVAE38_(VideoVAE_):
attn_scales, self.temperal_upsample, dropout)
def encode(self, x, scale):
def encode(self, x):
self.clear_cache()
x = patchify(x, patch_size=2)
t = x.shape[2]
@@ -1420,27 +1411,19 @@ class VideoVAE38_(VideoVAE_):
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):
scale = [s.to(dtype=mu.dtype, device=mu.device) for s in scale]
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
scale = scale.to(dtype=mu.dtype, device=mu.device)
mu = (mu - scale[0]) * scale[1]
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
self.clear_cache()
return mu
def decode(self, z, scale):
def decode(self, z):
self.clear_cache()
if isinstance(scale[0], torch.Tensor):
scale = [s.to(dtype=z.dtype, device=z.device) for s in scale]
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
1, self.z_dim, 1, 1, 1)
else:
scale = scale.to(dtype=z.dtype, device=z.device)
z = z / scale[1] + scale[0]
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
x = self.conv2(z)
for i in range(iter_):
@@ -1481,12 +1464,11 @@ class WanVideoVAE38(WanVideoVAE):
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.mean = torch.tensor(mean)
self.std = torch.tensor(std)
self.scale = [self.mean, 1.0 / self.std]
self.mean = torch.tensor(mean).view(1, z_dim, 1, 1, 1)
self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1)
self.dtype = dtype
self.z_dim = z_dim
# init model
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype).eval().requires_grad_(False)
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std).eval().requires_grad_(False)
self.upsampling_factor = 16