Add support for Wan 2.2 previews
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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")
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user