From 9ea61a2df34f33b24cdbaad28670ae263ab0c7e9 Mon Sep 17 00:00:00 2001 From: blepping Date: Fri, 29 Aug 2025 17:57:18 -0600 Subject: [PATCH] Add support for Wan 2.2 previews --- README.md | 2 ++ changelog.md | 5 ++++ py/better_previews/base.py | 47 +++++++++++++++++++++++++++++++++ py/better_previews/previewer.py | 37 +++----------------------- py/better_previews/tae_vid.py | 31 +++++++++++++++------- py/nodes/taevid.py | 38 +++++++++++++++++++------- 6 files changed, 109 insertions(+), 51 deletions(-) create mode 100644 py/better_previews/base.py diff --git a/README.md b/README.md index 43424ce..b75aa8f 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,7 @@ For recent user-visible changes, please see the [ChangeLog](changelog.md). * Better TAESD previews (see below). * Visual previews for some audio models (currently only ACE-Steps). +* Multi-frame video previews for most common video models (Wan 2.2, 2.1, Hunyuan, etc). See [the section on video encode/decode](#blehtaevideoencode-and-blehtaevideodecode). * Allow setting seed, timestep range and step interval for HyperTile (look for the [`BlehHyperTile`](#blehhypertile) node). * Allow applying Kohya Deep Shrink to multiple blocks, also allow gradually fading out the downscale factor (look for the [`BlehDeepShrink`](#blehdeepshrink) node). * Allow discarding penultimate sigma (look for the `BlehDiscardPenultimateSigma` node). This can be useful if you find certain samplers are ruining your image by spewing a bunch of noise into it at the very end (usually only an issue with `dpm2 a` or SDE samplers). @@ -284,6 +285,7 @@ Fast video latent encoding/decoding with models from madebyollin (same person th You will need to download the models and put them in `models/vae_approx`. Don't change the names. +* **WAN 2.2**: https://github.com/madebyollin/taehv/blob/main/taew2_2.pth * **WAN 2.1**: https://github.com/madebyollin/taehv/blob/main/taew2_1.pth * **Hunyean**: https://github.com/madebyollin/taehv/blob/main/taehv.pth * **Mochi**: https://github.com/madebyollin/taem1/blob/main/taem1.pth diff --git a/changelog.md b/changelog.md index 58b83c7..d8202f7 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,11 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20250829 + +* Added support for Wan 2.2 video previews. +* The `BlehTAEVideoEncode` should work for encoding single images or numbers of frames that aren't a multiple of the video latent temporal compression size. The input will be padded with the last frame. + ## 20250809 This set of changes involves refactoring parts of the previewer. Please create an issue if you experience problems. diff --git a/py/better_previews/base.py b/py/better_previews/base.py new file mode 100644 index 0000000..04b3897 --- /dev/null +++ b/py/better_previews/base.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, NamedTuple + +from comfy import latent_formats + +if TYPE_CHECKING: + from pathlib import Path + + +class VideoModelInfo(NamedTuple): + latent_format: latent_formats.LatentFormat + fps: int = 24 + temporal_compression: int = 8 + patch_size: int = 1 + tae_model: str | Path | None = None + + +VIDEO_FORMATS = { + "mochi": VideoModelInfo( + latent_formats.Mochi, + temporal_compression=6, + tae_model="taem1.pth", + ), + "hunyuanvideo": VideoModelInfo( + latent_formats.HunyuanVideo, + temporal_compression=4, + tae_model="taehv.pth", + ), + "cosmos1cv8x8x8": VideoModelInfo(latent_formats.Cosmos1CV8x8x8), + "wan21": VideoModelInfo( + latent_formats.Wan21, + fps=16, + temporal_compression=4, + tae_model="taew2_1.pth", + ), + "wan22": VideoModelInfo( + latent_formats.Wan22, + fps=24, + temporal_compression=4, + patch_size=2, + tae_model="taew2_2.pth", + ), +} + + +__all__ = ("VIDEO_FORMATS", "VideoModelInfo") diff --git a/py/better_previews/previewer.py b/py/better_previews/previewer.py index b3cf618..384403e 100644 --- a/py/better_previews/previewer.py +++ b/py/better_previews/previewer.py @@ -2,12 +2,11 @@ from __future__ import annotations import math from time import time -from typing import TYPE_CHECKING, NamedTuple +from typing import TYPE_CHECKING import folder_paths import latent_preview import torch -from comfy import latent_formats from comfy.cli_args import LatentPreviewMethod from comfy.cli_args import args as comfy_args from comfy.model_management import device_supports_non_blocking, vae_dtype @@ -16,12 +15,12 @@ from PIL import Image from tqdm import tqdm from ..settings import SETTINGS # noqa: TID252 +from .base import VIDEO_FORMATS, VideoModelInfo from .tae_vid import TAEVid if TYPE_CHECKING: - from pathlib import Path - import numpy as np + from comfy import latent_formats _ORIG_PREVIEWER = latent_preview.TAESDPreviewerImpl _ORIG_GET_PREVIEWER = latent_preview.get_previewer @@ -82,34 +81,6 @@ def normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)): ) -class VideoModelInfo(NamedTuple): - latent_format: latent_formats.LatentFormat - fps: int = 24 - temporal_compression: int = 8 - tae_model: str | Path | None = None - - -VIDEO_FORMATS = { - "mochi": VideoModelInfo( - latent_formats.Mochi, - temporal_compression=6, - tae_model="taem1.pth", - ), - "hunyuanvideo": VideoModelInfo( - latent_formats.HunyuanVideo, - temporal_compression=4, - tae_model="taehv.pth", - ), - "cosmos1cv8x8x8": VideoModelInfo(latent_formats.Cosmos1CV8x8x8), - "wan21": VideoModelInfo( - latent_formats.Wan21, - fps=16, - temporal_compression=4, - tae_model="taew2_1.pth", - ), -} - - class ImageWrapper: def __init__(self, frames: tuple, frame_duration: int): self._frames = frames @@ -653,7 +624,7 @@ def bleh_get_previewer( tae_model = ( TAEVid( checkpoint_path=tae_model_path, - latent_channels=latent_format.latent_channels, + vmi=vid_info, device=torch.device("cpu"), decoder_time_upscale=decoder_time_upscale, ) diff --git a/py/better_previews/tae_vid.py b/py/better_previews/tae_vid.py index 9ae451b..325f2f0 100644 --- a/py/better_previews/tae_vid.py +++ b/py/better_previews/tae_vid.py @@ -14,6 +14,8 @@ if TYPE_CHECKING: from collections.abc import Iterable from pathlib import Path + from .base import VideoModelInfo + F = torch.nn.functional @@ -184,22 +186,26 @@ class TAEVidContext: class TAEVid(nn.Module): temporal_upscale_blocks = 2 spatial_upscale_blocks = 3 + _nf = (256, 128, 64, 64) def __init__( self, *, checkpoint_path: str | Path, - latent_channels: int, + vmi: VideoModelInfo, image_channels: int = 3, device="cpu", decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), ): + n_f = self._nf super().__init__() - self.latent_channels = latent_channels + self.vmi = vmi + self.latent_channels = vmi.latent_format.latent_channels self.image_channels = image_channels + self.patch_size = vmi.patch_size self.encoder = nn.Sequential( - conv(image_channels, 64), + conv(image_channels * self.patch_size**2, 64), nn.ReLU(inplace=True), TPool(64, 2), conv(64, 64, stride=2, bias=False), @@ -216,13 +222,12 @@ class TAEVid(nn.Module): MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), - conv(64, latent_channels), + conv(64, vmi.latent_format.latent_channels), ) - n_f = (256, 128, 64, 64) self.frames_to_trim = 2 ** sum(decoder_time_upscale) - 1 self.decoder = nn.Sequential( Clamp(), - conv(latent_channels, n_f[0]), + conv(vmi.latent_format.latent_channels, n_f[0]), nn.ReLU(inplace=True), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), @@ -243,7 +248,7 @@ class TAEVid(nn.Module): TGrow(n_f[2], 2 if decoder_time_upscale[1] else 1), conv(n_f[2], n_f[3], bias=False), nn.ReLU(inplace=True), - conv(n_f[3], image_channels), + conv(n_f[3], image_channels * self.patch_size**2), ) if checkpoint_path is None: return @@ -299,9 +304,17 @@ class TAEVid(nn.Module): show_progress=False, ) -> torch.Tensor: model = self.decoder if decode else self.encoder + if not decode and self.vmi.patch_size > 1: + x = F.pixel_unshuffle(x, self.patch_size) if parallel: - return self.apply_parallel(x, model, show_progress=show_progress) - return TAEVidContext(model).apply(x, show_progress=show_progress) + result = self.apply_parallel(x, model, show_progress=show_progress) + else: + result = TAEVidContext(model).apply(x, show_progress=show_progress) + return ( + result + if not decode or self.vmi.patch_size < 2 + else F.pixel_shuffle(result, self.patch_size) + ) def decode(self, *args: list, **kwargs: dict) -> torch.Tensor: return self.apply(*args, decode=True, **kwargs)[:, self.frames_to_trim :] diff --git a/py/nodes/taevid.py b/py/nodes/taevid.py index b1a7945..0892342 100644 --- a/py/nodes/taevid.py +++ b/py/nodes/taevid.py @@ -1,8 +1,9 @@ # ruff: noqa: TID252 -import torch # noqa: I001 +import math import folder_paths +import torch from comfy import model_management from ..better_previews.previewer import VIDEO_FORMATS, VideoModelInfo @@ -17,7 +18,7 @@ class TAEVideoNodeBase: def INPUT_TYPES(cls) -> dict: return { "required": { - "latent_type": (("wan21", "hunyuanvideo", "mochi"),), + "latent_type": (("wan21", "wan22", "hunyuanvideo", "mochi"),), "parallel_mode": ( "BOOLEAN", { @@ -40,6 +41,8 @@ class TAEVideoNodeBase: if tae_model_path is None: if latent_type == "wan21": model_src = "taew2_1.pth from https://github.com/madebyollin/taehv" + elif latent_type == "wan22": + model_src = "taew2_2.pth from https://github.com/madebyollin/taehv" elif latent_type == "hunyuanvideo": model_src = "taehv.pth from https://github.com/madebyollin/taehv" else: @@ -49,11 +52,7 @@ class TAEVideoNodeBase: device = model_management.vae_device() dtype = model_management.vae_dtype(device=device) return ( - TAEVid( - checkpoint_path=tae_model_path, - latent_channels=vmi.latent_format.latent_channels, - device=device, - ).to(device), + TAEVid(checkpoint_path=tae_model_path, vmi=vmi, device=device).to(device), device, dtype, vmi, @@ -115,10 +114,31 @@ class TAEVideoEncode(TAEVideoNodeBase): def go(cls, *, image: torch.Tensor, latent_type: str, parallel_mode: bool) -> tuple: model, device, dtype, vmi = cls.get_taevid_model(latent_type) image = image.detach().to(device=device, dtype=dtype, copy=True) - if image.ndim == 4: + if image.ndim < 5: image = image.unsqueeze(0) + if image.ndim < 5: + image = image.unsqueeze(0) + if image.ndim != 5: + raise ValueError("Unexpected input image dimensions") + frames = image.shape[1] + add_frames = ( + math.ceil(frames / vmi.temporal_compression) * vmi.temporal_compression + - frames + ) + if add_frames > 0: + image = torch.cat( + ( + image, + image[:, frames - 1 :, ...].expand( + image.shape[0], + add_frames, + *image.shape[2:], + ), + ), + dim=1, + ) latent = model.encode( - image.movedim(-1, 2), + image[..., :3].movedim(-1, 2), parallel=parallel_mode, show_progress=True, ).transpose(1, 2)