diff --git a/nodes.py b/nodes.py index 0ef571c..cee21ea 100644 --- a/nodes.py +++ b/nodes.py @@ -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 = { diff --git a/taehv/taehv.py b/taehv/taehv.py index c0a8fc8..287817c 100644 --- a/taehv/taehv.py +++ b/taehv/taehv.py @@ -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:]