From 926ceb64166b455afab8181eefe0a9e0f5009acf Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 13 Mar 2025 07:34:16 -0600 Subject: [PATCH] OOM fallback feature in previewer, other previewer updates. --- README.md | 7 ++ blehconfig.example.json | 5 +- blehconfig.example.yaml | 17 +++++ changelog.md | 5 ++ py/betterTaesdPreview.py | 160 ++++++++++++++++++++++++++++++++++----- py/settings.py | 3 + 6 files changed, 178 insertions(+), 19 deletions(-) diff --git a/README.md b/README.md index 0e4c068..55a339a 100644 --- a/README.md +++ b/README.md @@ -2,6 +2,8 @@ A ComfyUI nodes collection of utility and model patching functions. Also includes improved previewer that allows previewing batches during generation. +For recent user-visible changes, please see the [ChangeLog](changelog.md). + ## Features 1. Better TAESD previews (see below). @@ -42,6 +44,9 @@ Current defaults: |`maxed_batch_step_mode`|`false`|When `false`, you will see the first `max_batch` previews, when `true` you will see previews spread across the batch| |`preview_device`|`null`|`null` (use the default device) or a string with a PyTorch device name like `"cpu"`, `"cuda:0"`, etc. Can be used to run TAESD previews on CPU or other available devices. Not recommended to change this unless you really need to, using the CPU device may prevent out of memory errors but will likely significantly slow down generation.| |`skip_upscale_layers`|`0`|The TAESD model has three upscale layers, each doubles the size of the result. Skipping some of them will significantly speed up TAESD previews at the cost of smaller preview image results. You can set this to `-1` to automatically pop layers until at least one dimension is within the max width/height or `-2` to aggressively pop until _both_ dimensions are within the limit.| +|`compile_previewer`|`false`|Controls whether the previewer gets compiled with `torch.compile`. May be a boolean or an object in which case the object will be used as argument to `torch.compile`. Note: May cause a delay/memory spike on the first preview.| +|`oom_fallback`|`latent2rgb`|May be set to `none` or `latent2rgb`. Controls what happens if trying to decode the preview runs out of memory.| +|`oom_retry`|`true`|If set to `false`, we will give up and use the `oom_fallback` behavior after hitting the first OOM. Otherwise, we'll attempt to decode with the normal previewer each time a preview is requested, even if that previously ran out of memory.| These defaults are conservative. I would recommend setting `throttle_secs` to something relatively high (like 5-10) especially if you are generating batches at high resolution. @@ -49,6 +54,8 @@ Slightly more detailed explanation for `maxed_batch_step_mode`: If max previews More detailed explanation for skipping upscale layers: Latents (the thing you're running the TAESD preview on) are 8 times smaller than the image you get decoding by normal VAE or TAESD. The TAESD decoder has three upscale layers, each doubling the size: `1 * 2 * 2 * 2 = 8`. So for example if normal decoding would get you a `1280x1280` image, skipping one TAESD upscale layer will get you a `640x640` result, skipping two will get you `320x320` and so on. I did some testing running TAESD decode on CPU for a `1280x1280` image: the base speed is about `1.95` sec base, `1.15` sec with one upscale layer skipped, `0.44` sec with two upscale layers skipped and `0.16` sec with all three upscale layers popped (of course you only get a `160x160` preview at that point). The upshot is if you are using TAESD to preview large images or batches or you want to run TAESD on CPU (normally pretty slow) you would probably benefit from setting `skip_upscale_layers` to `1` or `2`. Also if your max preview size is `768` and you are decoding a `1280x1280` image, it's just going to get scaled down to `768x768` anyway. +**Note**: Other node packs that patch ComfyUI's previewer behavior may interfere with this feature. One I am aware of is [ComfyUI-VideoHelperSuite](https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite) - if you have displaying animated previews turned on, it will overwrite Bleh's patched previewer. Or possibly, depending on the load order, Bleh will prevent it from working correctly. + ### BlehModelPatchConditional **Note**: Very experimental. diff --git a/blehconfig.example.json b/blehconfig.example.json index 12ded32..33334a4 100644 --- a/blehconfig.example.json +++ b/blehconfig.example.json @@ -7,6 +7,9 @@ "throttle_secs": 1, "maxed_batch_step_mode": false, "preview_device": null, - "skip_upscale_layers": 0 + "skip_upscale_layers": 0, + "compile_previewer": false, + "oom_fallback": "latent2rgb", + "oom_retry": true } } diff --git a/blehconfig.example.yaml b/blehconfig.example.yaml index 22887d9..161a048 100644 --- a/blehconfig.example.yaml +++ b/blehconfig.example.yaml @@ -31,3 +31,20 @@ betterTaesdPreviews: # Allows skipping upscale layers in the TAESD model, may increase performance when previewing large images or batches. # May be set to -1 (conservative) or -2 (aggressive) to automatically calculate how many to skip. See README.md for details. skip_upscale_layers: 0 + + # Controls whether the previewer model is compiled (using torch.compile). Only works if your + # Torch version and GPU support compiling. This also may cause a delay/memory spike on decoding the first preview. + # This may be a boolean or object with arguments to pass to torch.compile. For example: + # compile_previewer: + # mode: max-autotune + # backend: inductor + compile_previewer: false + + # Controls behavior if we run out of memory trying to decode the preview. + # Possible values: none, latent2rgb + oom_fallback: "latent2rgb" + + # When enabled, we will try to use the normal previewer on each call + # and only use the fallback if the normal previewer fails. + # When disabled, we use the fallback starting from the first OOM. + oom_retry: true diff --git a/changelog.md b/changelog.md index d1818c8..b366dbe 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,11 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20250313 + +* Added OOM fallback to the previewer. +* Added ability to compile the previewer (and a few other related options). + ## 20250119 * Fixed SageAttention 1.x support and setting `tensor_layout` should work properly for SageAttention 2.x now. Please create an issue if you experience problems. diff --git a/py/betterTaesdPreview.py b/py/betterTaesdPreview.py index 44246bc..7d7aab8 100644 --- a/py/betterTaesdPreview.py +++ b/py/betterTaesdPreview.py @@ -1,20 +1,68 @@ -import logging import math from time import time +from typing import NamedTuple import latent_preview import torch +from comfy.latent_formats import LatentFormat from comfy.model_management import device_supports_non_blocking from PIL import Image +from tqdm import tqdm from .settings import SETTINGS _ORIG_PREVIEWER = latent_preview.TAESDPreviewerImpl +_ORIG_GET_PREVIEWER = latent_preview.get_previewer + +LAST_LATENT_FORMAT = None + + +class FallbackPreviewerModel(torch.nn.Module): + @torch.no_grad() + def __init__( + self, + latent_format: LatentFormat, + *, + dtype: torch.dtype, + device: torch.device, + scale_factor: float = 8.0, + upscale_mode: str = "bilinear", + ): + super().__init__() + + raw_factors = latent_format.latent_rgb_factors + raw_bias = latent_format.latent_rgb_factors_bias + factors = torch.tensor(raw_factors, device=device, dtype=dtype).transpose(0, 1) + bias = ( + torch.tensor(raw_bias, device=device, dtype=dtype) + if raw_bias is not None + else None + ) + self.lin = torch.nn.Linear( + factors.shape[1], + factors.shape[0], + device=device, + dtype=dtype, + bias=bias is not None, + ) + self.upsample = torch.nn.Upsample(scale_factor=scale_factor, mode=upscale_mode) + self.requires_grad_(False) # noqa: FBT003 + self.lin.weight.copy_(factors) + if bias is not None: + self.lin.bias.copy_(bias) + + @torch.no_grad() + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.lin(x.movedim(1, -1)).movedim(-1, 1) + x = self.upsample(x).movedim(1, -1) + 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 + self.fallback_previewer_model = None self.device = ( None if SETTINGS.btp_preview_device is None @@ -26,6 +74,9 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): self.stamp = None self.cached = None self.blank = Image.new("RGB", size=(1, 1)) + self.oom_fallback = SETTINGS.btp_oom_fallback == "latent2rgb" + self.oom_retry = SETTINGS.btp_oom_retry + self.oom_count = 0 self.skip_upscale_layers = SETTINGS.btp_skip_upscale_layers self.preview_max_width = SETTINGS.btp_max_width self.preview_max_height = SETTINGS.btp_max_height @@ -33,9 +84,18 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): self.max_batch_preview = SETTINGS.btp_max_batch self.maxed_batch_step_mode = SETTINGS.btp_maxed_batch_step_mode self.max_batch_cols = SETTINGS.btp_max_batch_cols + self.compile_previewer = SETTINGS.btp_compile_previewer self.maybe_pop_upscale_layers() + if self.compile_previewer: + compile_kwargs = ( + {} + if not isinstance(self.compile_previewer, dict) + else self.compile_previewer + ) + self.taesd = torch.compile(self.taesd, **compile_kwargs) - def maybe_pop_upscale_layers(self, *, width=None, height=None): + # 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: return @@ -69,7 +129,11 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): self.taesd.taesd_decoder.pop(upscale_layers[-idx]) self.skip_upscale_layers = 0 - def decode_latent_to_preview_image(self, preview_format, x0): + def decode_latent_to_preview_image( + self, + preview_format: str, + x0: torch.Tensor, + ) -> tuple[str, Image, int]: preview_image = self.decode_latent_to_preview(x0) return ( preview_format, @@ -80,7 +144,7 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): ), ) - def check_use_cached(self): + def check_use_cached(self) -> bool: now = time() if ( self.cached is not None and self.stamp is not None @@ -89,7 +153,7 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): self.stamp = now return False - def _decode_latent(self, x0): + def prepare_decode_latent(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]: max_batch = self.max_batch_preview batch = x0.shape[0] if not self.maxed_batch_step_mode: @@ -101,7 +165,7 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): math.ceil(batch / max_batch), )[:max_batch] x0 = x0[indexes, :] - batch, _channels, height, width = x0.shape + batch, (height, width) = x0.shape[0], x0.shape[-2:] if self.device and x0.device != self.device: x0 = x0.to( device=self.device, @@ -112,6 +176,11 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): width, height, ) + return x0, cols, rows + + def _decode_latent(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]: + x0, cols, rows = self.prepare_decode_latent(x0) + height, width = x0.shape[-2:] if self.skip_upscale_layers < 0: self.maybe_pop_upscale_layers( width=width * 8 * cols, @@ -121,17 +190,21 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): ( self.taesd.decode(x0) .movedim(1, -1) - .add_(1) - .mul_(0.5) - .clamp_(min=0, max=1) - .mul_(255) + .add_(1.0) + .mul_(127.5) + .clamp_(min=0, max=255.0) .detach() ), cols, rows, ) - def calc_cols_rows(self, batch_size, width, height): + def calc_cols_rows( + self, + batch_size: int, + width: int, + height: int, + ) -> tuple[int, int]: max_cols = self.max_batch_cols ratio = height / width if ratio >= 1.45: @@ -145,8 +218,8 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): rows = math.ceil(batch_size / cols) return cols, rows - def decoded_to_image(self, samples, cols, rows): - batch, height, width = samples.shape[:-1] + def decoded_to_image(self, samples: torch.Tensor, cols: int, rows: int) -> Image: + batch, (height, width) = samples.shape[0], samples.shape[-3:-1] samples = samples.to( device="cpu", dtype=torch.uint8, @@ -169,20 +242,71 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER): ) return result - def decode_latent_to_preview(self, x0): + @torch.no_grad() + def init_fallback_previewer(self, device: torch.device, dtype: torch.dtype) -> bool: + if self.fallback_previewer_model is not None: + return True + if self.latent_format is None: + return False + self.fallback_previewer_model = FallbackPreviewerModel( + self.latent_format, + device=device, + dtype=dtype, + ) + return True + + def fallback_previewer(self, x0: torch.Tensor, *, quiet=False) -> Image: + 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}.", + ) + 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.", + ) + return self.blank + x0, cols, rows = self.prepare_decode_latent(x0) + try: + return self.decoded_to_image( + self.fallback_previewer_model(x0), + cols, + rows, + ) + except torch.OutOfMemoryError: + return self.blank + + def decode_latent_to_preview(self, x0: torch.Tensor) -> Image: if self.check_use_cached(): 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: + return self.fallback_previewer(x0, quiet=True) try: return self.decoded_to_image(*self._decode_latent(x0)) except torch.OutOfMemoryError: - logging.warning( - "*** BlehBetterTAESDPreviews: Got out of memory error while decoding preview - skipping.", - ) - return self.blank + return self.fallback_previewer(x0) + + +def bleh_get_previewer_wrapper( + device, + latent_format: LatentFormat, + *args: list, + **kwargs: dict, +): + global LAST_LATENT_FORMAT # noqa: PLW0603 + LAST_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 + +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 diff --git a/py/settings.py b/py/settings.py index 0a80e04..30e32d2 100644 --- a/py/settings.py +++ b/py/settings.py @@ -20,6 +20,9 @@ class Settings: 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) @staticmethod def get_cfg_path(filename) -> Path: