Fix scaling for light TAE

This commit is contained in:
kijai
2025-12-05 14:03:33 +02:00
parent f85ea5a48a
commit 9e005fab90
2 changed files with 24 additions and 16 deletions
+15 -14
View File
@@ -2174,7 +2174,7 @@ class WanVideoDecode:
"tile_stride_y": ("INT", {"default": 128, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride height in pixels. Smaller values use less VRAM but will introduce more seams."}),
},
"optional": {
"normalization": (["default", "minmax"], {"advanced": True}),
"normalization": (["default", "minmax", "none"], {"advanced": True}),
}
}
@@ -2217,23 +2217,24 @@ class WanVideoDecode:
if drop_last:
latents = latents[:, :, :-1]
if type(vae).__name__ == "TAEHV":
if type(vae).__name__ == "TAEHV":
images = vae.decode_video(latents.permute(0, 2, 1, 3, 4), cond=flashvsr_LQ_images.to(vae.dtype) if flashvsr_LQ_images is not None else None)[0].permute(1, 0, 2, 3)
images = torch.clamp(images, 0.0, 1.0)
images = images.permute(1, 2, 3, 0).cpu().float()
return (images,)
else:
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]
images = images.cpu().float()
if normalization == "minmax":
images.sub_(images.min()).div_(images.max() - images.min())
else:
images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if normalization != "none":
if normalization == "minmax":
images.sub_(images.min()).div_(images.max() - images.min())
else:
images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if is_looped:
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]
@@ -2244,7 +2245,7 @@ class WanVideoDecode:
if end_image is not None:
images = images[:, 0:-1]
vae.to(offload_device)
mm.soft_empty_cache()
@@ -2295,7 +2296,7 @@ class WanVideoEncodeLatentBatch:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, 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))
else:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling)
if latent_strength != 1.0:
latent *= latent_strength
latent_list.append(latent.squeeze(0).cpu())
@@ -2355,7 +2356,7 @@ class WanVideoEncode:
latents = latents.permute(0, 2, 1, 3, 4)
else:
latents = vae.encode(image * 2.0 - 1.0, device=device, 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))
vae.to(offload_device)
if latent_strength != 1.0:
latents *= latent_strength
@@ -2364,7 +2365,7 @@ class WanVideoEncode:
log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}")
mm.soft_empty_cache()
return ({"samples": latents, "noise_mask": mask},)
NODE_CLASS_MAPPINGS = {
+9 -2
View File
@@ -8,6 +8,7 @@ import torch.nn as nn
import torch.nn.functional as F
from tqdm.auto import tqdm
from collections import namedtuple
from ..wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
@@ -146,7 +147,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
return x
class TAEHV(nn.Module):
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16):
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16, model_name="taehv"):
"""Initialize pretrained TAEHV from the given checkpoint.
Arg:
@@ -161,6 +162,7 @@ class TAEHV(nn.Module):
if self.latent_channels == 48:
self.patch_size = 2
self.dtype = dtype
self.model_name = model_name
self.encoder = nn.Sequential(
conv(self.image_channels*self.patch_size**2, 64), nn.ReLU(inplace=True),
@@ -180,8 +182,11 @@ class TAEHV(nn.Module):
)
if state_dict is not None:
self.load_state_dict(self.patch_tgrow_layers(state_dict))
self.parallel = parallel
orig_vae = WanVideoVAE38() if self.latent_channels == 48 else WanVideoVAE()
self.mean = orig_vae.mean.to(dtype).movedim(1, 2)
self.inv_std = orig_vae.inv_std.to(dtype).movedim(1, 2)
def patch_tgrow_layers(self, sd):
"""Patch TGrow layers to use a smaller kernel if needed.
@@ -221,6 +226,8 @@ class TAEHV(nn.Module):
if False, frames will be processed sequentially.
Returns NTCHW RGB tensor with ~[0, 1] values.
"""
if "light" in self.model_name.lower():
x = x / self.inv_std.to(x) + self.mean.to(x)
x = apply_model_with_memblocks(self.decoder, x, self.parallel, show_progress_bar)
if self.patch_size > 1: x = F.pixel_shuffle(x, self.patch_size)
return x[:, self.frames_to_trim:]