From 78520b756816e1323241a51ccb3fd48cd5f9f2ec Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 20 Mar 2025 09:14:25 -0600 Subject: [PATCH] Add basic TAE video encode/decode nodes Tweak some other TAE video stuff --- __init__.py | 3 + blehconfig.example.yaml | 2 +- py/better_previews/previewer.py | 36 +++++++-- py/better_previews/tae_vid.py | 30 +++++--- py/nodes/taevid.py | 128 ++++++++++++++++++++++++++++++++ py/settings.py | 5 +- 6 files changed, 185 insertions(+), 19 deletions(-) create mode 100644 py/nodes/taevid.py diff --git a/__init__.py b/__init__.py index 625c8f9..eb40a4c 100644 --- a/__init__.py +++ b/__init__.py @@ -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 = { diff --git a/blehconfig.example.yaml b/blehconfig.example.yaml index 025abb2..02a169c 100644 --- a/blehconfig.example.yaml +++ b/blehconfig.example.yaml @@ -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. diff --git a/py/better_previews/previewer.py b/py/better_previews/previewer.py index 6b0c071..1f9e824 100644 --- a/py/better_previews/previewer.py +++ b/py/better_previews/previewer.py @@ -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( diff --git a/py/better_previews/tae_vid.py b/py/better_previews/tae_vid.py index 266ebfe..9ae451b 100644 --- a/py/better_previews/tae_vid.py +++ b/py/better_previews/tae_vid.py @@ -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) diff --git a/py/nodes/taevid.py b/py/nodes/taevid.py new file mode 100644 index 0000000..d40b2ad --- /dev/null +++ b/py/nodes/taevid.py @@ -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},) diff --git a/py/settings.py b/py/settings.py index a33a535..c378c36 100644 --- a/py/settings.py +++ b/py/settings.py @@ -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: