Fix scaling for light TAE
This commit is contained in:
@@ -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}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2228,6 +2228,7 @@ class WanVideoDecode:
|
|||||||
|
|
||||||
images = images.cpu().float()
|
images = images.cpu().float()
|
||||||
|
|
||||||
|
if normalization != "none":
|
||||||
if normalization == "minmax":
|
if normalization == "minmax":
|
||||||
images.sub_(images.min()).div_(images.max() - images.min())
|
images.sub_(images.min()).div_(images.max() - images.min())
|
||||||
else:
|
else:
|
||||||
|
|||||||
+8
-1
@@ -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),
|
||||||
@@ -182,6 +184,9 @@ class TAEHV(nn.Module):
|
|||||||
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:]
|
||||||
|
|||||||
Reference in New Issue
Block a user