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."}), "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": { "optional": {
"normalization": (["default", "minmax"], {"advanced": True}), "normalization": (["default", "minmax", "none"], {"advanced": True}),
} }
} }
@@ -2217,23 +2217,24 @@ class WanVideoDecode:
if drop_last: if drop_last:
latents = latents[:, :, :-1] 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 = 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 = torch.clamp(images, 0.0, 1.0)
images = images.permute(1, 2, 3, 0).cpu().float() images = images.permute(1, 2, 3, 0).cpu().float()
return (images,) return (images,)
else: 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 = 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() images = images.cpu().float()
if normalization == "minmax": if normalization != "none":
images.sub_(images.min()).div_(images.max() - images.min()) if normalization == "minmax":
else: images.sub_(images.min()).div_(images.max() - images.min())
images.clamp_(-1.0, 1.0) else:
images.add_(1.0).div_(2.0) images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if is_looped: if is_looped:
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2) 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 = 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: if end_image is not None:
images = images[:, 0:-1] images = images[:, 0:-1]
vae.to(offload_device) vae.to(offload_device)
mm.soft_empty_cache() 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)) 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: else:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling) 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: if latent_strength != 1.0:
latent *= latent_strength latent *= latent_strength
latent_list.append(latent.squeeze(0).cpu()) latent_list.append(latent.squeeze(0).cpu())
@@ -2355,7 +2356,7 @@ class WanVideoEncode:
latents = latents.permute(0, 2, 1, 3, 4) latents = latents.permute(0, 2, 1, 3, 4)
else: 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)) 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) vae.to(offload_device)
if latent_strength != 1.0: if latent_strength != 1.0:
latents *= latent_strength latents *= latent_strength
@@ -2364,7 +2365,7 @@ class WanVideoEncode:
log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}") log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}")
mm.soft_empty_cache() mm.soft_empty_cache()
return ({"samples": latents, "noise_mask": mask},) return ({"samples": latents, "noise_mask": mask},)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
+9 -2
View File
@@ -8,6 +8,7 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from tqdm.auto import tqdm from tqdm.auto import tqdm
from collections import namedtuple from collections import namedtuple
from ..wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
DecoderResult = namedtuple("DecoderResult", ("frame", "memory")) DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index")) TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
@@ -146,7 +147,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
return x return x
class TAEHV(nn.Module): 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. """Initialize pretrained TAEHV from the given checkpoint.
Arg: Arg:
@@ -161,6 +162,7 @@ class TAEHV(nn.Module):
if self.latent_channels == 48: if self.latent_channels == 48:
self.patch_size = 2 self.patch_size = 2
self.dtype = dtype self.dtype = dtype
self.model_name = model_name
self.encoder = nn.Sequential( self.encoder = nn.Sequential(
conv(self.image_channels*self.patch_size**2, 64), nn.ReLU(inplace=True), 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: if state_dict is not None:
self.load_state_dict(self.patch_tgrow_layers(state_dict)) self.load_state_dict(self.patch_tgrow_layers(state_dict))
self.parallel = parallel 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): def patch_tgrow_layers(self, sd):
"""Patch TGrow layers to use a smaller kernel if needed. """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. if False, frames will be processed sequentially.
Returns NTCHW RGB tensor with ~[0, 1] values. 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) 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) if self.patch_size > 1: x = F.pixel_shuffle(x, self.patch_size)
return x[:, self.frames_to_trim:] return x[:, self.frames_to_trim:]