Support TAE video models, stage 1

This commit is contained in:
blepping
2025-03-18 07:56:08 -06:00
parent 926ceb6416
commit 0916e65794
8 changed files with 503 additions and 58 deletions
+2
View File
@@ -303,3 +303,5 @@ Also may be an item from [Filters](#filters).
## Credits
Latent blending and scaling and filter functions based on implementation from https://github.com/WASasquatch/FreeU_Advanced - thanks!
TAE video model support based on code https://github.com/madebyollin/taehv/.
+1 -3
View File
@@ -4,9 +4,6 @@ BLEH_VERSION = 1
settings.load_settings()
if settings.SETTINGS.btp_enabled:
from .py import betterTaesdPreview # noqa: F401
from .py.nodes import (
blockCFG,
deepShrink,
@@ -41,6 +38,7 @@ NODE_CLASS_MAPPINGS = {
"BlehSetSamplerPreset": samplers.BlehSetSamplerPreset,
"BlehCast": misc.BlehCast,
"BlehSetSigmas": misc.BlehSetSigmas,
"BlehEnsurePreviwer": misc.BlehEnsurePreviewer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+12
View File
@@ -48,3 +48,15 @@ betterTaesdPreviews:
# and only use the fallback if the normal previewer fails.
# When disabled, we use the fallback starting from the first OOM.
oom_retry: true
# List of lowercase latent format names from https://github.com/comfyanonymous/ComfyUI/blob/master/comfy/latent_formats.py
# If the list is empty, this disables the whitelist. Otherwise, Bleh will
# only handle previewing for formats in the list.
whitelist_formats: []
# List of lowercase latent format names (see above).
# Bleh will delegate to the normal previewer for any latent formats in the blacklist.
blacklist_formats: []
# Controls whether video previewing uses parallel mode (faster, requires more memory).
video_parallel: true
+3
View File
@@ -0,0 +1,3 @@
from .previewer import ensure_previewer
__all__ = ("ensure_previewer",)
@@ -1,15 +1,26 @@
from __future__ import annotations
import math
from time import time
from typing import NamedTuple
from typing import TYPE_CHECKING, NamedTuple
import folder_paths
import latent_preview
import torch
from comfy.latent_formats import LatentFormat
from comfy.cli_args import LatentPreviewMethod
from comfy.cli_args import args as comfy_args
from comfy.model_management import device_supports_non_blocking
from comfy.taesd.taesd import TAESD
from PIL import Image
from tqdm import tqdm
from .settings import SETTINGS
from ..settings import SETTINGS # noqa: TID252
from .tae_vid import TAEVid
if TYPE_CHECKING:
from pathlib import Path
from comfy import latent_formats
_ORIG_PREVIEWER = latent_preview.TAESDPreviewerImpl
_ORIG_GET_PREVIEWER = latent_preview.get_previewer
@@ -17,11 +28,25 @@ _ORIG_GET_PREVIEWER = latent_preview.get_previewer
LAST_LATENT_FORMAT = None
class VideoModelInfo(NamedTuple):
fps: int = 24
temporal_compression: int = 8
tae_model: str | Path | None = None
VIDEO_FORMATS = {
"mochi": VideoModelInfo(temporal_compression=6),
"hunyuanvideo": VideoModelInfo(temporal_compression=4),
"cosmos1cv8x8x8": VideoModelInfo(),
"wan21": VideoModelInfo(fps=16, temporal_compression=4, tae_model="taew2_1.pth"),
}
class FallbackPreviewerModel(torch.nn.Module):
@torch.no_grad()
def __init__(
self,
latent_format: LatentFormat,
latent_format: latent_formats.LatentFormat,
*,
dtype: torch.dtype,
device: torch.device,
@@ -58,18 +83,29 @@ class FallbackPreviewerModel(torch.nn.Module):
return x.add_(1.0).mul_(127.5).clamp_(0.0, 255.0)
class BetterTAESDPreviewer(_ORIG_PREVIEWER):
def __init__(self, taesd):
del taesd.taesd_encoder
self.latent_format = LAST_LATENT_FORMAT
class BetterPreviewer(_ORIG_PREVIEWER):
def __init__(
self,
*,
taesd: torch.nn.Module | None = None,
latent_format: latent_formats.LatentFormat,
vid_info: VideoModelInfo | None = None,
):
self.latent_format = latent_format
self.vid_info = vid_info
self.fallback_previewer_model = None
self.device = (
None
if SETTINGS.btp_preview_device is None
else torch.device(SETTINGS.btp_preview_device)
)
if self.device and self.device != next(taesd.parameters()).device:
taesd = taesd.to(self.device)
if taesd is not None:
if hasattr(taesd, "taesd_encoder"):
del taesd.taesd_encoder
if hasattr(taesd, "encoder"):
del taesd.encoder
if self.device and self.device != next(taesd.parameters()).device:
taesd = taesd.to(self.device)
self.taesd = taesd
self.stamp = None
self.cached = None
@@ -97,7 +133,7 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
# Popping upscale layers trick from https://github.com/madebyollin/
def maybe_pop_upscale_layers(self, *, width=None, height=None) -> None:
skip = self.skip_upscale_layers
if skip == 0:
if skip == 0 or not isinstance(self.taesd, TAESD):
return
upscale_layers = tuple(
idx
@@ -153,19 +189,28 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
self.stamp = now
return False
def prepare_decode_latent(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]:
def calculate_indexes(self, batch_size: int) -> range:
max_batch = self.max_batch_preview
batch = x0.shape[0]
if not self.maxed_batch_step_mode:
indexes = range(min(max_batch, batch))
else:
indexes = range(
0,
batch,
math.ceil(batch / max_batch),
)[:max_batch]
x0 = x0[indexes, :]
batch, (height, width) = x0.shape[0], x0.shape[-2:]
return range(min(max_batch, batch_size))
return range(
0,
batch_size,
math.ceil(batch_size / max_batch),
)[:max_batch]
def prepare_decode_latent(
self,
x0: torch.Tensor,
*,
frames_to_batch=True,
) -> tuple[torch.Tensor, int, int]:
if frames_to_batch and x0.ndim == 5:
x0 = x0.transpose(2, 1).reshape(-1, x0.shape[1], *x0.shape[-2:])
batch = x0.shape[0]
x0 = x0[self.calculate_indexes(batch), :]
batch = x0.shape[0]
height, width = x0.shape[-2:]
if self.device and x0.device != self.device:
x0 = x0.to(
device=self.device,
@@ -178,8 +223,34 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
)
return x0, cols, rows
def _decode_latent(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]:
x0, cols, rows = self.prepare_decode_latent(x0)
def _decode_latent_taevid(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]:
height, width = x0.shape[-2:]
if self.device and x0.device != self.device:
x0 = x0.to(
device=self.device,
non_blocking=device_supports_non_blocking(x0.device),
)
decoded = self.taesd.decode(x0.transpose(1, 2)).movedim(2, -1)
del x0
decoded = decoded.reshape(-1, *decoded.shape[2:])
batch = decoded.shape[0]
decoded = decoded[self.calculate_indexes(batch), :]
cols, rows = self.calc_cols_rows(
min(batch, self.max_batch_preview),
width,
height,
)
return (
decoded.mul_(255.0).round_().clamp_(min=0, max=255.0).detach(),
cols,
rows,
)
def _decode_latent_taesd(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]:
x0, cols, rows = self.prepare_decode_latent(
x0,
frames_to_batch=not isinstance(self.taesd, TAEVid),
)
height, width = x0.shape[-2:]
if self.skip_upscale_layers < 0:
self.maybe_pop_upscale_layers(
@@ -259,14 +330,14 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
if not quiet:
fallback_mode = "using fallback" if self.oom_fallback else "skipping"
tqdm.write(
f"*** BlehBetterTAESDPreviews: Got out of memory error while decoding preview - {fallback_mode}.",
f"*** BlehBetterPreviews: Got out of memory error while decoding preview - {fallback_mode}.",
)
if not self.oom_fallback:
return self.blank
if not self.init_fallback_previewer(x0.device, x0.dtype):
self.oom_fallback = False
tqdm.write(
"*** BlehBetterTAESDPreviews: Couldn't initialize fallback previewer, giving up on previews.",
"*** BlehBetterPreviews: Couldn't initialize fallback previewer, giving up on previews.",
)
return self.blank
x0, cols, rows = self.prepare_decode_latent(x0)
@@ -284,29 +355,81 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
return self.cached
if x0.shape[0] == 0:
return self.blank # Shouldn't actually be possible.
if self.oom_count and not self.oom_retry:
if (self.oom_count and not self.oom_retry) or self.taesd is None:
return self.fallback_previewer(x0, quiet=True)
try:
return self.decoded_to_image(*self._decode_latent(x0))
if isinstance(self.taesd, TAEVid):
return self.decoded_to_image(*self._decode_latent_taevid(x0))
return self.decoded_to_image(*self._decode_latent_taesd(x0))
except torch.OutOfMemoryError:
return self.fallback_previewer(x0)
def bleh_get_previewer_wrapper(
def bleh_get_previewer(
device,
latent_format: LatentFormat,
latent_format: latent_formats.LatentFormat,
*args: list,
**kwargs: dict,
):
global LAST_LATENT_FORMAT # noqa: PLW0603
LAST_LATENT_FORMAT = latent_format
) -> object | None:
preview_method = comfy_args.preview_method
format_name = latent_format.__class__.__name__.lower()
if (
not SETTINGS.btp_enabled
or format_name in SETTINGS.btp_blacklist
or (SETTINGS.btp_whitelist and format_name not in SETTINGS.btp_whitelist)
):
return _ORIG_GET_PREVIEWER(device, latent_format, *args, **kwargs)
tae_model = None
if preview_method in {LatentPreviewMethod.TAESD, LatentPreviewMethod.Auto}:
vid_info = VIDEO_FORMATS.get(format_name)
if vid_info is not None and vid_info.tae_model is not None:
tae_model_path = folder_paths.get_full_path(
"vae_approx",
vid_info.tae_model,
)
tae_model = (
TAEVid(
checkpoint_path=tae_model_path,
latent_channels=latent_format.latent_channels,
device=device,
decoder_time_upscale=(False, False, True),
).to(device)
if tae_model_path is not None
else None
)
if tae_model is None and latent_format.taesd_decoder_name is not None:
taesd_path = folder_paths.get_full_path(
"vae_approx",
f"{latent_format.taesd_decoder_name}.pth",
)
tae_model = (
TAESD(
None,
taesd_path,
latent_channels=latent_format.latent_channels,
).to(device)
if taesd_path is not None
else None
)
return BetterPreviewer(
taesd=tae_model,
latent_format=latent_format,
vid_info=vid_info,
)
if (
preview_method == LatentPreviewMethod.NoPreviews
or latent_format.latent_rgb_factors is None
):
return None
if preview_method == LatentPreviewMethod.Latent2RGB:
return BetterPreviewer(latent_format=latent_format)
return _ORIG_GET_PREVIEWER(device, latent_format, *args, **kwargs)
if not isinstance(latent_preview.TAESDPreviewerImpl, BetterTAESDPreviewer):
latent_preview.BLEH_ORIG_TAESDPreviewerImpl = _ORIG_PREVIEWER
latent_preview.TAESDPreviewerImpl = BetterTAESDPreviewer
def ensure_previewer():
if latent_preview.get_previewer != bleh_get_previewer:
latent_preview.BLEH_ORIG_get_previewer = _ORIG_GET_PREVIEWER
latent_preview.get_previewer = bleh_get_previewer
if latent_preview.get_previewer != bleh_get_previewer_wrapper:
latent_preview.BLEH_ORIG_get_previewer = _ORIG_GET_PREVIEWER
latent_preview.get_previewer = bleh_get_previewer_wrapper
ensure_previewer()
+274
View File
@@ -0,0 +1,274 @@
# Modified from https://github.com/madebyollin/taehv/blob/main/taehv.py
# ruff: noqa: N806
from __future__ import annotations
from typing import TYPE_CHECKING, NamedTuple
import torch
from torch import nn
from tqdm.auto import tqdm
if TYPE_CHECKING:
from pathlib import Path
F = torch.nn.functional
class TWorkItem(NamedTuple):
input_tensor: torch.Tensor
block_index: int
def conv(n_in, n_out, **kwargs: dict) -> nn.Conv2d:
return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module):
@classmethod
def forward(cls, x):
return torch.tanh(x / 3) * 3
class MemBlock(nn.Module):
def __init__(self, n_in, n_out):
super().__init__()
self.conv = nn.Sequential(
conv(n_in * 2, n_out),
nn.ReLU(inplace=True),
conv(n_out, n_out),
nn.ReLU(inplace=True),
conv(n_out, n_out),
)
self.skip = (
nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity()
)
self.act = nn.ReLU(inplace=True)
def forward(self, x, past):
return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x))
class TPool(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f * stride, n_f, 1, bias=False)
def forward(self, x):
return self.conv(x.reshape(-1, self.stride * x.shape[1], *x.shape[-2]))
class TGrow(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f, n_f * stride, 1, bias=False)
def forward(self, x):
x = self.conv(x)
return x.reshape(-1, *x.shape[1:])
def apply_model_with_memblocks(model, x, *, show_progress_bar=False):
if x.ndim != 5:
raise ValueError("Expected 5 dimensional tensor")
N, T, C, H, W = x.shape
out = []
# iterate over input timesteps and also iterate over blocks.
# because of the cursed TPool/TGrow blocks, this is not a nested loop,
# it's actually a ***graph traversal*** problem! so let's make a queue
work_queue = [
TWorkItem(xt, 0)
for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))
]
# in addition to manually managing our queue, we also need to manually manage our progressbar.
# we'll update it for every source node that we consume.
progress_bar = tqdm(range(T), disable=not show_progress_bar)
# we'll also need a separate addressable memory per node as well
mem = [None] * len(model)
while work_queue:
xt, i = work_queue.pop(0)
if i == 0:
# new source node consumed
progress_bar.update(1)
if i == len(model):
# reached end of the graph, append result to output list
out.append(xt)
continue
# fetch the block to process
b = model[i]
if isinstance(b, MemBlock):
# mem blocks are simple since we're visiting the graph in causal order
if mem[i] is None:
xt_new = b(xt, xt * 0)
mem[i] = xt
else:
xt_new = b(xt, mem[i])
mem[i].copy_(
xt,
) # inplace might reduce mysterious pytorch memory allocations? doesn't help though
# add successor to work queue
work_queue.insert(0, TWorkItem(xt_new, i + 1))
elif isinstance(b, TPool):
# pool blocks are miserable
if mem[i] is None:
mem[i] = [] # pool memory is itself a queue of inputs to pool
mem[i].append(xt)
if len(mem[i]) > b.stride:
# pool mem is in invalid state, we should have pooled before this
raise RuntimeError("Internal error: Invalid mem state")
if len(mem[i]) < b.stride:
# pool mem is not yet full, go back to processing the work queue
pass
else:
# pool mem is ready, run the pool block
N, C, H, W = xt.shape
xt = b(torch.cat(mem[i], 1).view(N * b.stride, C, H, W))
# reset the pool mem
mem[i] = []
# add successor to work queue
work_queue.insert(0, TWorkItem(xt, i + 1))
elif isinstance(b, TGrow):
xt = b(xt)
C, H, W = xt.shape[1:]
# each tgrow has multiple successor nodes
for xt_next in reversed(
xt.view(N, b.stride * C, H, W).chunk(b.stride, 1),
):
# add successor to work queue
work_queue.insert(0, TWorkItem(xt_next, i + 1))
else:
# normal block with no funny business
xt = b(xt)
# add successor to work queue
work_queue.insert(0, TWorkItem(xt, i + 1))
progress_bar.close()
return torch.stack(out, 1)
class TAEVid(nn.Module):
def __init__(
self,
*,
checkpoint_path: str | Path,
latent_channels: int,
image_channels: int = 3,
device="cpu",
decoder_time_upscale=(True, True),
decoder_space_upscale=(True, True, True),
):
super().__init__()
self.latent_channels = latent_channels
self.image_channels = image_channels
self.encoder = nn.Sequential(
conv(image_channels, 64),
nn.ReLU(inplace=True),
TPool(64, 2),
conv(64, 64, stride=2, bias=False),
MemBlock(64, 64),
MemBlock(64, 64),
MemBlock(64, 64),
TPool(64, 2),
conv(64, 64, stride=2, bias=False),
MemBlock(64, 64),
MemBlock(64, 64),
MemBlock(64, 64),
TPool(64, 1),
conv(64, 64, stride=2, bias=False),
MemBlock(64, 64),
MemBlock(64, 64),
MemBlock(64, 64),
conv(64, 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]),
nn.ReLU(inplace=True),
MemBlock(n_f[0], n_f[0]),
MemBlock(n_f[0], n_f[0]),
MemBlock(n_f[0], n_f[0]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1),
TGrow(n_f[0], 1),
conv(n_f[0], n_f[1], bias=False),
MemBlock(n_f[1], n_f[1]),
MemBlock(n_f[1], n_f[1]),
MemBlock(n_f[1], n_f[1]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1),
TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1),
conv(n_f[1], n_f[2], bias=False),
MemBlock(n_f[2], n_f[2]),
MemBlock(n_f[2], n_f[2]),
MemBlock(n_f[2], n_f[2]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[2] else 1),
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),
)
if checkpoint_path is None:
return
self.load_state_dict(
self.patch_tgrow_layers(
torch.load(checkpoint_path, map_location=device, weights_only=True),
),
)
def patch_tgrow_layers(self, sd: dict) -> dict:
new_sd = self.state_dict()
for i, layer in enumerate(self.decoder):
if isinstance(layer, TGrow):
key = f"decoder.{i}.conv.weight"
if sd[key].shape[0] > new_sd[key].shape[0]:
# take the last-timestep output channels
sd[key] = sd[key][-new_sd[key].shape[0] :]
return sd
@classmethod
def apply_parallel(
cls,
x: torch.Tensor,
model: nn.Module,
*,
show_progress_bar=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):
# tqdm.write(f"BLOCK({b.__class__.__name__}): {x.shape}")
if not isinstance(b, MemBlock):
x = b(x)
continue
nt, c, h, w = x.shape
t = nt // n
mem = F.pad(x.reshape(n, t, c, h, w), padding, value=0)[:, :t].reshape(
x.shape,
)
x = b(x, mem)
del mem
nt, c, h, w = x.shape
t = nt // n
return x.view(n, t, c, h, w)
def decode(self, x: torch.Tensor, *, parallel=True) -> torch.Tensor:
if parallel:
result = self.apply_parallel(x, self.decoder)
else:
result = apply_model_with_memblocks(self.decoder, x)
return result[:, self.frames_to_trim :]
# def encode_video(self, x, parallel=True, show_progress_bar=True):
# return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar)
# def decode_video(self, x, parallel=True, show_progress_bar=True):
# x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar)
# return x[:, self.frames_to_trim :]
def forward(self, x):
return self.c(x)
+31
View File
@@ -7,6 +7,8 @@ from decimal import Decimal
import torch
from comfy import model_management
from ..better_previews import ensure_previewer # noqa: TID252
class DiscardPenultimateSigma:
@classmethod
@@ -287,3 +289,32 @@ class BlehSetSigmas:
argb = sigmas_b
sigmas_out[start_index : start_index + newlen] = opfun(arga, argb)
return (sigmas_out.to(torch.float),)
class BlehEnsurePreviewer:
DESCRIPTION = "This node ensures Bleh is used for previews. Can be used if other custom nodes overwrite the Bleh previewer. It will pass through any value unchanged."
FUNCTION = "go"
OUTPUT_NODE = False
CATEGORY = "hacks"
WILDCARD = Wildcard("*")
RETURN_TYPES = (WILDCARD,)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"any_input": (
cls.WILDCARD,
{
"forceInput": True,
"description": "You can connect any type of input here, but take to ensure that you connect the output from this node to an input that is compatible.",
},
),
},
}
@classmethod
def go(cls, *, any_input):
ensure_previewer()
return (any_input,)
+18 -16
View File
@@ -7,22 +7,24 @@ class Settings:
def update(self, obj):
btp = obj.get("betterTaesdPreviews", None)
if btp is None:
self.btp_enabled = False
else:
self.btp_enabled = True
max_size = max(8, btp.get("max_size", 768))
self.btp_max_width = max(8, btp.get("max_width", max_size))
self.btp_max_height = max(8, btp.get("max_height", max_size))
self.btp_max_batch = max(1, btp.get("max_batch", 4))
self.btp_max_batch_cols = max(1, btp.get("max_batch_cols", 2))
self.btp_throttle_secs = btp.get("throttle_secs", 1)
self.btp_skip_upscale_layers = btp.get("skip_upscale_layers", 0)
self.btp_preview_device = btp.get("preview_device")
self.btp_maxed_batch_step_mode = btp.get("maxed_batch_step_mode", False)
self.btp_compile_previewer = btp.get("compile_previewer", False)
self.btp_oom_fallback = btp.get("oom_fallback", "latent2rgb")
self.btp_oom_retry = btp.get("oom_retry", True)
self.btp_enabled = btp is not None and btp.get("enabled", True) is True
if not self.btp_enabled:
return
max_size = max(8, btp.get("max_size", 768))
self.btp_max_width = max(8, btp.get("max_width", max_size))
self.btp_max_height = max(8, btp.get("max_height", max_size))
self.btp_max_batch = max(1, btp.get("max_batch", 4))
self.btp_max_batch_cols = max(1, btp.get("max_batch_cols", 2))
self.btp_throttle_secs = btp.get("throttle_secs", 1)
self.btp_skip_upscale_layers = btp.get("skip_upscale_layers", 0)
self.btp_preview_device = btp.get("preview_device")
self.btp_maxed_batch_step_mode = btp.get("maxed_batch_step_mode", False)
self.btp_compile_previewer = btp.get("compile_previewer", False)
self.btp_oom_fallback = btp.get("oom_fallback", "latent2rgb")
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)
@staticmethod
def get_cfg_path(filename) -> Path: