From 0916e6579427cdf768ce10109433d5de96ec1fb4 Mon Sep 17 00:00:00 2001 From: blepping Date: Tue, 18 Mar 2025 07:56:08 -0600 Subject: [PATCH] Support TAE video models, stage 1 --- README.md | 2 + __init__.py | 4 +- blehconfig.example.yaml | 12 + py/better_previews/__init__.py | 3 + .../previewer.py} | 201 ++++++++++--- py/better_previews/tae_vid.py | 274 ++++++++++++++++++ py/nodes/misc.py | 31 ++ py/settings.py | 34 ++- 8 files changed, 503 insertions(+), 58 deletions(-) create mode 100644 py/better_previews/__init__.py rename py/{betterTaesdPreview.py => better_previews/previewer.py} (60%) create mode 100644 py/better_previews/tae_vid.py diff --git a/README.md b/README.md index 55a339a..8b297b0 100644 --- a/README.md +++ b/README.md @@ -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/. diff --git a/__init__.py b/__init__.py index 8918e05..625c8f9 100644 --- a/__init__.py +++ b/__init__.py @@ -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 = { diff --git a/blehconfig.example.yaml b/blehconfig.example.yaml index 161a048..f0a862b 100644 --- a/blehconfig.example.yaml +++ b/blehconfig.example.yaml @@ -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 diff --git a/py/better_previews/__init__.py b/py/better_previews/__init__.py new file mode 100644 index 0000000..ab48f12 --- /dev/null +++ b/py/better_previews/__init__.py @@ -0,0 +1,3 @@ +from .previewer import ensure_previewer + +__all__ = ("ensure_previewer",) diff --git a/py/betterTaesdPreview.py b/py/better_previews/previewer.py similarity index 60% rename from py/betterTaesdPreview.py rename to py/better_previews/previewer.py index 7d7aab8..f18c93f 100644 --- a/py/betterTaesdPreview.py +++ b/py/better_previews/previewer.py @@ -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() diff --git a/py/better_previews/tae_vid.py b/py/better_previews/tae_vid.py new file mode 100644 index 0000000..a315780 --- /dev/null +++ b/py/better_previews/tae_vid.py @@ -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) diff --git a/py/nodes/misc.py b/py/nodes/misc.py index f0c7a02..1f1c25f 100644 --- a/py/nodes/misc.py +++ b/py/nodes/misc.py @@ -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,) diff --git a/py/settings.py b/py/settings.py index 30e32d2..94bc065 100644 --- a/py/settings.py +++ b/py/settings.py @@ -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: