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
+2 -1
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}),
}
}
@@ -2228,6 +2228,7 @@ class WanVideoDecode:
images = images.cpu().float()
if normalization != "none":
if normalization == "minmax":
images.sub_(images.min()).div_(images.max() - images.min())
else:
+8 -1
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),
@@ -182,6 +184,9 @@ class TAEHV(nn.Module):
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:]