Add basic TAE video encode/decode nodes
Tweak some other TAE video stuff
This commit is contained in:
@@ -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 = {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user