Add support for Wan 2.2 previews

This commit is contained in:
blepping
2025-08-29 17:57:18 -06:00
parent b58c9637c1
commit 9ea61a2df3
6 changed files with 109 additions and 51 deletions
+2
View File
@@ -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
+5
View File
@@ -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.
+47
View File
@@ -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")
+4 -33
View File
@@ -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,
)
+22 -9
View File
@@ -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 :]
+29 -9
View File
@@ -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)