Support TAE video models, stage 1
This commit is contained in:
@@ -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
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user