OOM fallback feature in previewer, other previewer updates.

This commit is contained in:
blepping
2025-03-13 07:34:16 -06:00
parent 0ab8900dd0
commit 926ceb6416
6 changed files with 178 additions and 19 deletions
+7
View File
@@ -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.
+4 -1
View File
@@ -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
}
}
+17
View File
@@ -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
+5
View File
@@ -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.
+142 -18
View File
@@ -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
+3
View File
@@ -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: