diff --git a/blehconfig.example.yaml b/blehconfig.example.yaml index f0a862b..3a772f2 100644 --- a/blehconfig.example.yaml +++ b/blehconfig.example.yaml @@ -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 diff --git a/py/better_previews/previewer.py b/py/better_previews/previewer.py index f18c93f..7f8c1ad 100644 --- a/py/better_previews/previewer.py +++ b/py/better_previews/previewer.py @@ -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 diff --git a/py/better_previews/tae_vid.py b/py/better_previews/tae_vid.py index a315780..82678ad 100644 --- a/py/better_previews/tae_vid.py +++ b/py/better_previews/tae_vid.py @@ -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) diff --git a/py/settings.py b/py/settings.py index 94bc065..262b761 100644 --- a/py/settings.py +++ b/py/settings.py @@ -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: