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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user