Unbork borked ported TAEvid code

Refactor stuff
Initial internal support for animated previews
Untested support for Hunyuan and Mochi TAEvid decoding
Add separate "batch" limit for video latent frames
This commit is contained in:
blepping
2025-03-20 06:14:21 -06:00
parent 0916e65794
commit c55189cc5a
4 changed files with 202 additions and 107 deletions
+11
View File
@@ -60,3 +60,14 @@ betterTaesdPreviews:
# Controls whether video previewing uses parallel mode (faster, requires more memory).
video_parallel: true
# 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.
video_max_frames: -1
# One of: video, batch, both, none
# Probably will only work if you set the preview type to webp.
# When active, rows/columns are ignored and batch items or video frames will be
# produced as an animated WEBP.
# NOTE: Does not work correctly yet.
animate_preview: none
+72 -14
View File
@@ -20,6 +20,7 @@ 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
@@ -35,13 +36,39 @@ class VideoModelInfo(NamedTuple):
VIDEO_FORMATS = {
"mochi": VideoModelInfo(temporal_compression=6),
"hunyuanvideo": VideoModelInfo(temporal_compression=4),
"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"),
}
class ImageWrapper:
def __init__(self, frames: tuple, frame_duration: int):
self._frames = frames
self._frame_duration = frame_duration
def save(self, fp, format: str | None, **kwargs: dict): # noqa: A002
if len(self._frames) == 1:
return self._frames[0].save(fp, format, **kwargs)
kwargs |= {
"loop": 0,
"save_all": True,
"append_images": self._frames[1:],
"duration": self._frame_duration,
}
return self._frames[0].save(fp, "webp", **kwargs)
def resize(self, *args: list, **kwargs: dict) -> ImageWrapper:
return ImageWrapper(
tuple(frame.resize(*args, **kwargs) for frame in self._frames),
frame_duration=self._frame_duration,
)
def __getattr__(self, key):
return getattr(self._frames[0], key)
class FallbackPreviewerModel(torch.nn.Module):
@torch.no_grad()
def __init__(
@@ -172,7 +199,7 @@ class BetterPreviewer(_ORIG_PREVIEWER):
) -> tuple[str, Image, int]:
preview_image = self.decode_latent_to_preview(x0)
return (
preview_format,
preview_format if not isinstance(preview_image, ImageWrapper) else "WEBP",
preview_image,
min(
max(*preview_image.size),
@@ -189,8 +216,12 @@ class BetterPreviewer(_ORIG_PREVIEWER):
self.stamp = now
return False
def calculate_indexes(self, batch_size: int) -> range:
max_batch = self.max_batch_preview
def calculate_indexes(self, batch_size: int, *, is_video=False) -> range:
max_batch = (
SETTINGS.btp_video_max_frames if is_video else self.max_batch_preview
)
if max_batch < 0:
return range(batch_size)
if not self.maxed_batch_step_mode:
return range(min(max_batch, batch_size))
return range(
@@ -205,10 +236,11 @@ class BetterPreviewer(_ORIG_PREVIEWER):
*,
frames_to_batch=True,
) -> tuple[torch.Tensor, int, int]:
if frames_to_batch and x0.ndim == 5:
is_video = x0.ndim == 5
if frames_to_batch and is_video:
x0 = x0.transpose(2, 1).reshape(-1, x0.shape[1], *x0.shape[-2:])
batch = x0.shape[0]
x0 = x0[self.calculate_indexes(batch), :]
x0 = x0[self.calculate_indexes(batch, is_video=is_video), :]
batch = x0.shape[0]
height, width = x0.shape[-2:]
if self.device and x0.device != self.device:
@@ -230,7 +262,10 @@ class BetterPreviewer(_ORIG_PREVIEWER):
device=self.device,
non_blocking=device_supports_non_blocking(x0.device),
)
decoded = self.taesd.decode(x0.transpose(1, 2)).movedim(2, -1)
decoded = self.taesd.decode(
x0.transpose(1, 2),
parallel=SETTINGS.btp_video_parallel,
).movedim(2, -1)
del x0
decoded = decoded.reshape(-1, *decoded.shape[2:])
batch = decoded.shape[0]
@@ -289,7 +324,22 @@ class BetterPreviewer(_ORIG_PREVIEWER):
rows = math.ceil(batch_size / cols)
return cols, rows
def decoded_to_image(self, samples: torch.Tensor, cols: int, rows: int) -> Image:
@classmethod
def decoded_to_animation(cls, samples: np.ndarray) -> ImageWrapper:
batch = samples.shape[0]
return ImageWrapper(
tuple(Image.fromarray(samples[idx]) for idx in range(batch)),
frame_duration=250,
)
def decoded_to_image(
self,
samples: torch.Tensor,
cols: int,
rows: int,
*,
is_video=False,
) -> Image | ImageWrapper:
batch, (height, width) = samples.shape[0], samples.shape[-3:-1]
samples = samples.to(
device="cpu",
@@ -299,8 +349,12 @@ class BetterPreviewer(_ORIG_PREVIEWER):
if batch == 1:
self.cached = Image.fromarray(samples[0])
return self.cached
if SETTINGS.btp_animate_preview == "both" or (
is_video,
SETTINGS.btp_animate_preview,
) in {(True, "video"), (False, "batch")}:
return self.decoded_to_animation(samples)
cols, rows = self.calc_cols_rows(batch, width, height)
img_size = (width * cols, height * rows)
if self.cached is not None and self.cached.size == img_size:
result = self.cached
@@ -357,10 +411,14 @@ class BetterPreviewer(_ORIG_PREVIEWER):
return self.blank # Shouldn't actually be possible.
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
try:
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))
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)
except torch.OutOfMemoryError:
return self.fallback_previewer(x0)
@@ -392,7 +450,7 @@ def bleh_get_previewer(
checkpoint_path=tae_model_path,
latent_channels=latent_format.latent_channels,
device=device,
decoder_time_upscale=(False, False, True),
decoder_time_upscale=(False, False),
).to(device)
if tae_model_path is not None
else None
+117 -93
View File
@@ -11,6 +11,7 @@ from torch import nn
from tqdm.auto import tqdm
if TYPE_CHECKING:
from collections.abc import Iterable
from pathlib import Path
F = torch.nn.functional
@@ -21,13 +22,13 @@ class TWorkItem(NamedTuple):
block_index: int
def conv(n_in, n_out, **kwargs: dict) -> nn.Conv2d:
def conv(n_in: int, n_out: int, **kwargs: dict) -> nn.Conv2d:
return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module):
@classmethod
def forward(cls, x):
def forward(cls, x: torch.Tensor) -> torch.Tensor:
return torch.tanh(x / 3) * 3
@@ -46,8 +47,8 @@ class MemBlock(nn.Module):
)
self.act = nn.ReLU(inplace=True)
def forward(self, x, past):
return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x))
def forward(self, x: torch.Tensor, past: torch.Tensor) -> torch.Tensor:
return self.act(self.conv(torch.cat((x, past), 1)) + self.skip(x))
class TPool(nn.Module):
@@ -56,8 +57,9 @@ class TPool(nn.Module):
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]))
def forward(self, x: torch.Tensor) -> torch.Tensor:
c, h, w = x.shape[-3:]
return self.conv(x.reshape(-1, self.stride * c, h, w))
class TGrow(nn.Module):
@@ -66,87 +68,117 @@ class TGrow(nn.Module):
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 forward(self, x: torch.Tensor) -> torch.Tensor:
orig_shape = x.shape
return self.conv(x).reshape(-1, *orig_shape[-3:])
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
class TAEVidContext:
def __init__(self, model):
self.model = model
self.HANDLERS = {
MemBlock: self.handle_memblock,
TPool: self.handle_tpool,
TGrow: self.handle_tgrow,
}
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))
def reset(self, x: torch.Tensor) -> None:
N, T, C, H, W = x.shape
self.N, self.T = N, T
self.work_queue = [
TWorkItem(xt, 0)
for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))
]
self.mem = [None] * len(self.model)
def handle_memblock(
self,
i: int,
xt: torch.Tensor,
b: nn.Module,
) -> Iterable[torch.Tensor]:
mem = self.mem
# mem blocks are simple since we're visiting the graph in causal order
if mem[i] is None:
xt_new = b(xt, torch.zeros_like(xt))
mem[i] = xt
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)
xt_new = b(xt, mem[i])
# inplace might reduce mysterious pytorch memory allocations? doesn't help though
mem[i].copy_(xt)
return (xt_new,)
def handle_tpool(
self,
i: int,
xt: torch.Tensor,
b: nn.Module,
) -> Iterable[torch.Tensor]:
mem = self.mem
# 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
return ()
# 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] = []
return (xt,)
def handle_tgrow(
self,
_i: int,
xt: torch.Tensor,
b: nn.Module,
) -> Iterable[torch.Tensor]:
xt = b(xt)
C, H, W = xt.shape[1:]
return reversed(
xt.view(self.N, b.stride * C, H, W).chunk(b.stride, 1),
)
@classmethod
def handle_default(
cls,
_i: int,
xt: torch.Tensor,
b: nn.Module,
) -> Iterable[torch.Tensor]:
return (b(xt),)
def handle_block(self, i: int, xt: torch.Tensor, b: nn.Module) -> None:
handler = self.HANDLERS.get(b.__class__, self.handle_default)
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:
if x.ndim != 5:
raise ValueError("Expected 5 dimensional tensor")
self.reset(x)
out = []
work_queue = self.work_queue
model = self.model
model_len = len(model)
with tqdm(range(self.T), disable=not show_progress_bar) as pbar:
while work_queue:
xt, i = work_queue.pop(0)
if i == model_len:
# reached end of the graph, append result to output list
out.append(xt)
continue
if i == 0:
# new source node consumed
pbar.update(1)
self.handle_block(i, xt, model[i])
return torch.stack(out, 1)
class TAEVid(nn.Module):
@@ -241,7 +273,6 @@ class TAEVid(nn.Module):
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
@@ -260,15 +291,8 @@ class TAEVid(nn.Module):
if parallel:
result = self.apply_parallel(x, self.decoder)
else:
result = apply_model_with_memblocks(self.decoder, x)
result = TAEVidContext(self.decoder).apply(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):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.c(x)
+2
View File
@@ -25,6 +25,8 @@ class Settings:
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_max_frames = btp.get("video_max_frames", -1)
self.btp_animate_preview = btp.get("animate_preview", "none")
@staticmethod
def get_cfg_path(filename) -> Path: