Add basic TAE video encode/decode nodes

Tweak some other TAE video stuff
This commit is contained in:
blepping
2025-03-20 09:14:25 -06:00
parent 1324dfb355
commit 78520b7568
6 changed files with 185 additions and 19 deletions
+3
View File
@@ -14,6 +14,7 @@ from .py.nodes import (
refinerAfter,
sageAttention,
samplers,
taevid,
)
samplers.add_sampler_presets()
@@ -39,6 +40,8 @@ NODE_CLASS_MAPPINGS = {
"BlehCast": misc.BlehCast,
"BlehSetSigmas": misc.BlehSetSigmas,
"BlehEnsurePreviwer": misc.BlehEnsurePreviewer,
"BlehTAEVideoDecode": taevid.TAEVideoDecode,
"BlehTAEVideoEncode": taevid.TAEVideoEncode,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+1 -1
View File
@@ -59,7 +59,7 @@ betterTaesdPreviews:
blacklist_formats: []
# Controls whether video previewing uses parallel mode (faster, requires more memory).
video_parallel: true
video_parallel: false
# Maximum frames to include in a preview. -1 means no limit.
# Frame limiting is treated like batch limiting so maxed_batch_step mode, etc will apply here.
+29 -7
View File
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, NamedTuple
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
@@ -21,7 +22,6 @@ 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
@@ -30,16 +30,30 @@ LAST_LATENT_FORMAT = None
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(temporal_compression=6, tae_model="taem1.pth"),
"hunyuanvideo": VideoModelInfo(temporal_compression=4, tae_model="taehv.pth"),
"cosmos1cv8x8x8": VideoModelInfo(),
"wan21": VideoModelInfo(fps=16, temporal_compression=4, tae_model="taew2_1.pth"),
"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",
),
}
@@ -417,15 +431,23 @@ class BetterPreviewer(_ORIG_PREVIEWER):
if (self.oom_count and not self.oom_retry) or self.taesd is None:
return self.fallback_previewer(x0, quiet=True)
is_video = x0.ndim == 5
used_fallback = False
start_time = time()
try:
dargs = (
self._decode_latent_taevid(x0)
if is_video
else self._decode_latent_taesd(x0)
)
return self.decoded_to_image(*dargs, is_video=is_video)
result = self.decoded_to_image(*dargs, is_video=is_video)
except torch.OutOfMemoryError:
return self.fallback_previewer(x0)
used_fallback = True
result = self.fallback_previewer(x0)
if SETTINGS.btp_verbose:
tqdm.write(
f"BlehPreview: used fallback: {used_fallback}, decode time: {time() - start_time:0.2f}",
)
return result
def bleh_get_previewer(
+21 -9
View File
@@ -158,7 +158,7 @@ class TAEVidContext:
for xt_new in handler(i, xt, b):
self.work_queue.insert(0, TWorkItem(xt_new, i + 1))
def apply(self, x: torch.Tensor, *, show_progress_bar=False) -> torch.Tensor:
def apply(self, x: torch.Tensor, *, show_progress=False) -> torch.Tensor:
if x.ndim != 5:
raise ValueError("Expected 5 dimensional tensor")
self.reset(x)
@@ -167,7 +167,7 @@ class TAEVidContext:
model = self.model
model_len = len(model)
with tqdm(range(self.T), disable=not show_progress_bar) as pbar:
with tqdm(range(self.T), disable=not show_progress) as pbar:
while work_queue:
xt, i = work_queue.pop(0)
if i == model_len:
@@ -269,13 +269,13 @@ class TAEVid(nn.Module):
x: torch.Tensor,
model: nn.Module,
*,
show_progress_bar=False,
show_progress=False,
) -> torch.Tensor:
padding = (0, 0, 0, 0, 0, 0, 1, 0)
n, t, c, h, w = x.shape
x = x.reshape(n * t, c, h, w)
# parallel over input timesteps, iterate over blocks
for b in tqdm(model, disable=not show_progress_bar):
for b in tqdm(model, disable=not show_progress):
if not isinstance(b, MemBlock):
x = b(x)
continue
@@ -290,12 +290,24 @@ class TAEVid(nn.Module):
t = nt // n
return x.view(n, t, c, h, w)
def decode(self, x: torch.Tensor, *, parallel=True) -> torch.Tensor:
def apply(
self,
x: torch.Tensor,
*,
decode=True,
parallel=True,
show_progress=False,
) -> torch.Tensor:
model = self.decoder if decode else self.encoder
if parallel:
result = self.apply_parallel(x, self.decoder)
else:
result = TAEVidContext(self.decoder).apply(x)
return result[:, self.frames_to_trim :]
return self.apply_parallel(x, model, show_progress=show_progress)
return TAEVidContext(model).apply(x, show_progress=show_progress)
def decode(self, *args: list, **kwargs: dict) -> torch.Tensor:
return self.apply(*args, decode=True, **kwargs)[:, self.frames_to_trim :]
def encode(self, *args: list, **kwargs: dict) -> torch.Tensor:
return self.apply(*args, decode=False, **kwargs)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.c(x)
+128
View File
@@ -0,0 +1,128 @@
import torch # noqa: I001
import folder_paths
from comfy import model_management
from ..better_previews.previewer import VIDEO_FORMATS # noqa: TID252
from ..better_previews.tae_vid import TAEVid # noqa: TID252
class TAEVideoNodeBase:
FUNCTION = "go"
CATEGORY = "latent"
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"latent_type": (("wan21", "hunyuanvideo", "mochi"),),
"parallel_mode": (
"BOOLEAN",
{
"default": False,
"tooltip": "Parallel mode is faster but requires more memory.",
},
),
},
}
@classmethod
def get_taevid_model(
cls,
latent_type: str,
) -> tuple[TAEVid, torch.device, torch.dtype]:
vmi = VIDEO_FORMATS.get(latent_type)
if vmi is None or vmi.tae_model is None:
raise ValueError("Bad latent type")
tae_model_path = folder_paths.get_full_path("vae_approx", vmi.tae_model)
if tae_model_path is None:
if latent_type == "wan21":
model_src = "taew2_1.pth from https://github.com/madebyollin/taehv"
elif latent_type == "hunyuanvideo":
model_src = "taehv.pth from https://github.com/madebyollin/taehv"
else:
model_src = "taem1.pth from https://github.com/madebyollin/taem1"
err_string = f"Missing TAE video model. Download {model_src} and place it in the models/vae_approx directory"
raise RuntimeError(err_string)
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),
device,
dtype,
)
@classmethod
def go(cls, *, latent, latent_type: str, parallel_mode: bool) -> tuple:
pass
class TAEVideoDecode(TAEVideoNodeBase):
RETURN_TYPES = ("IMAGE",)
CATEGORY = "latent"
DESCRIPTION = "Fast decoding of Wan, Hunyuan and Mochi video latents with the video equivalent of TAESD."
@classmethod
def INPUT_TYPES(cls) -> dict:
result = super().INPUT_TYPES()
result["required"] |= {
"latent": ("LATENT",),
}
return result
@classmethod
def go(cls, *, latent: dict, latent_type: str, parallel_mode: bool) -> tuple:
model, device, dtype = cls.get_taevid_model(latent_type)
samples = latent["samples"].detach().to(device=device, dtype=dtype, copy=True)
img = (
model.decode(
samples.transpose(1, 2),
parallel=parallel_mode,
show_progress=True,
)
.movedim(2, -1)
.to(
dtype=torch.float,
device="cpu",
)
)
img = img.reshape(-1, *img.shape[-3:])
return (img,)
class TAEVideoEncode(TAEVideoNodeBase):
RETURN_TYPES = ("LATENT",)
CATEGORY = "latent"
DESCRIPTION = "Fast encoding of Wan, Hunyuan and Mochi video latents with the video equivalent of TAESD."
@classmethod
def INPUT_TYPES(cls) -> dict:
result = super().INPUT_TYPES()
result["required"] |= {
"image": ("IMAGE",),
}
return result
@classmethod
def go(cls, *, image: torch.Tensor, latent_type: str, parallel_mode: bool) -> tuple:
model, device, dtype = cls.get_taevid_model(latent_type)
image = image.detach().to(device=device, dtype=dtype, copy=True)
if image.ndim == 4:
image = image.unsqueeze(0)
latent = (
model.encode(
image.movedim(-1, 2),
parallel=parallel_mode,
show_progress=True,
)
.transpose(1, 2)
.to(
dtype=torch.float,
device="cpu",
)
)
return ({"samples": latent},)
+3 -2
View File
@@ -24,13 +24,14 @@ class Settings:
self.btp_oom_retry = btp.get("oom_retry", True)
self.btp_whitelist = frozenset(btp.get("whitelist_formats", frozenset()))
self.btp_blacklist = frozenset(btp.get("blacklist_formats", frozenset()))
self.btp_video_parallel = btp.get("video_parallel", True)
self.btp_video_parallel = btp.get("video_parallel", False)
self.btp_video_max_frames = btp.get("video_max_frames", -1)
self.btp_video_temporal_upscale_level = btp.get(
"video_temporal_upscale_level",
2,
0,
)
self.btp_animate_preview = btp.get("animate_preview", "none")
self.btp_verbose = btp.get("verbose", False)
@staticmethod
def get_cfg_path(filename) -> Path: