diff --git a/README.md b/README.md index 6ad1388..e774ebf 100644 --- a/README.md +++ b/README.md @@ -18,7 +18,7 @@ Restart ComfyUI to apply new changes. * Supports setting max preview size (ComfyUI default is hardcoded to 512 max). * Supports showing previews for more than the first latent in the batch. * Supports throttling previews. Do you really need your expensive TAESD preview to get updated 3 times a second? -* Supports using CUDA streams to avoid waiting for a synchronize. This might be faster. +* Supports using CUDA streams to avoid waiting for a synchronize. Increases speed slightly at the cost of higher VRAM usage. For comparison, a batch of 8 768x768 images with throttle at `0.5` sec is `1.29s/it` with it on and `1.42s/it` with it off for me. Current defaults from `blehconfig.json` @@ -29,9 +29,9 @@ Current defaults from `blehconfig.json` |`max_batch`|`4`|Max number of latents in a batch to preview| |`max_batch_cols`|`2`|Max number of columns to use when previewing batches| |`throttle_secs`|`1`|Max frequency to decode the latents for previewing. `0.25` would be every 1/4 sec, `2` would be only once every two seconds| -|`use_cuda`|`true`|Use special logic for CUDA (and maybe pretend-CUDA like ROCM) to reduce the performance impact of preview generation| +|`use_cuda`|`false`|Use special logic for CUDA (and maybe pretend-CUDA like ROCM) to reduce the performance impact of preview generation| -I would recommend setting `throttle_secs` to something relatively high like 5-10 sec especially if you are generating batches at high resolution. +I would recommend setting `throttle_secs` to something relatively high like 5-10 sec especially if you are generating batches at high resolution. `use_cuda` now defaults to `false` as it can substantially increase VRAM requirements. ### BlehHyperTile diff --git a/__init__.py b/__init__.py index 14d80a2..3b589d5 100644 --- a/__init__.py +++ b/__init__.py @@ -1,17 +1,18 @@ from .py import settings + settings.load_settings() if settings.SETTINGS.btp_enabled: - from .py import betterTaesdPreview + from .py import betterTaesdPreview from .py import hypertile NODE_CLASS_MAPPINGS = { - "BlehHyperTile": hypertile.HyperTileBleh, + "BlehHyperTile": hypertile.HyperTileBleh, } NODE_DISPLAY_NAME_MAPPINGS = { - "HyperTile (bleh)": "HyperTile (bleh)", + "HyperTile (bleh)": "HyperTile (bleh)", } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/py/betterTaesdPreview.py b/py/betterTaesdPreview.py index e99e2af..95a45f7 100644 --- a/py/betterTaesdPreview.py +++ b/py/betterTaesdPreview.py @@ -1,85 +1,111 @@ import math from time import time +import latent_preview import numpy as np import torch from PIL import Image -import latent_preview - from .settings import SETTINGS _ORIG_PREVIEWER = latent_preview.TAESDPreviewerImpl + class BetterTAESDPreviewer(_ORIG_PREVIEWER): - def __init__(self, taesd): - self.taesd = taesd - self.stamp = None - self.cached = None - self.blank = Image.new("RGB", size=(1,1)) - self.stream = None - self.prev_work = None - self.use_cuda = SETTINGS.btp_use_cuda and hasattr(torch, "cuda") and torch.cuda.is_available() + def __init__(self, taesd): + self.taesd = taesd + self.stamp = None + self.cached = None + self.blank = Image.new("RGB", size=(1, 1)) + self.stream = None + self.prev_work = None + self.cpudev = torch.device("cpu") + self.use_cuda = ( + SETTINGS.btp_use_cuda + and hasattr(torch, "cuda") + and torch.cuda.is_available() + ) - def decode_latent_to_preview_image(self, preview_format, x0): - preview_image = self.decode_latent_to_preview(x0) - return (preview_format, preview_image, min(max(*preview_image.size), SETTINGS.btp_max_size)) + def decode_latent_to_preview_image(self, preview_format, x0): + preview_image = self.decode_latent_to_preview(x0) + return ( + preview_format, + preview_image, + min(max(*preview_image.size), SETTINGS.btp_max_size), + ) - def check_use_cached(self): - now = time() - if self.cached is not None and self.stamp is not None: - if now - self.stamp < SETTINGS.btp_throttle_secs: - return True - self.stamp = now - return False + def check_use_cached(self): + now = time() + if ( + self.cached is not None and self.stamp is not None + ) and now - self.stamp < SETTINGS.btp_throttle_secs: + return True + self.stamp = now + return False - def _decode_latent(self, x0): - samples = (self.taesd.decode(x0[:SETTINGS.btp_max_batch]) + 1.0) / 2.0 - samples = torch.clamp(samples, min = 0.0, max = 1.0) * 255.0 - return samples.to(dtype = torch.uint8).detach() + def _decode_latent(self, x0): + samples = (self.taesd.decode(x0[: SETTINGS.btp_max_batch]) + 1.0) / 2.0 + samples = torch.clamp(samples, min=0.0, max=1.0) * 255.0 + return samples.to(dtype=torch.uint8).detach() - def decode_latent_to_preview(self, x0, cpudev = torch.device("cpu")): - if self.check_use_cached(): - return self.cached - if x0.device == cpudev or not self.use_cuda: - return self._decode_latent_to_preview_nocuda(x0) - if self.stream is None: - self.stream = torch.cuda.Stream() - elif not self.stream.query(): - return self.cached or self.blank - work = None - if self.prev_work is not None: - # We will only arrive here if the stream is ready. Sync just to be safe, should be instant. - self.stream.synchronize() - work = self.prev_work - del self.prev_work - # The default stream may be still processing the current step. - self.stream.wait_stream(torch.cuda.default_stream()) - try: - torch.cuda.set_stream(self.stream) - self.prev_work = self._decode_latent(x0).to(device=cpudev, non_blocking=True) - finally: - torch.cuda.set_stream(torch.cuda.default_stream()) - return self.work_to_image(work) if work is not None else self.blank + def decode_latent_to_preview(self, x0): + use_cached = self.check_use_cached() + if x0.device == self.cpudev or not self.use_cuda: + return ( + self.cached if use_cached else self._decode_latent_to_preview_nocuda(x0) + ) + if self.stream is None: + self.stream = torch.cuda.Stream() + elif not self.stream.query(): + return self.cached or self.blank + work = None + if self.prev_work is not None: + # We will only arrive here if the stream is ready. Sync just to be safe, should be instant. + self.stream.synchronize() + work = self.prev_work + del self.prev_work + result = self.work_to_image(work) if work is not None else self.blank + if use_cached: + return result + # The original stream may be still processing the current step. + orig_stream = torch.cuda.current_stream() + self.stream.wait_stream(orig_stream) + try: + torch.cuda.set_stream(self.stream) + self.prev_work = self._decode_latent(x0).to( + device=self.cpudev, + non_blocking=True, + ) + finally: + torch.cuda.set_stream(orig_stream) + return result - def work_to_image(self, samples): - samples = tuple(np.moveaxis(x, 0, 2) for x in samples.numpy()) - batch_size = len(samples) - height, width, _ = samples[0].shape - if batch_size < 2: - self.cached = Image.fromarray(samples[0]) - return self.cached - cols = min(math.ceil(batch_size / 2), SETTINGS.btp_max_batch_cols) - rows = math.ceil(batch_size / cols) - self.cached = result = Image.new("RGB", size = (width * cols, height * rows)) - for idx in range(batch_size): - result.paste(Image.fromarray(samples[idx]), - box = ((idx % cols) * width, ((idx // cols) % rows) * height), - ) - self.cached = result - return result + def calc_cols_rows(self, batch_size, width, height): + ratio = width / height + cols = min(math.ceil(batch_size / 2), SETTINGS.btp_max_batch_cols) + rows = math.ceil(batch_size / cols) + return cols, rows + + def work_to_image(self, samples): + samples = tuple(np.moveaxis(x, 0, 2) for x in samples.numpy()) + batch_size = len(samples) + height, width, _ = samples[0].shape + if batch_size < 2: + self.cached = Image.fromarray(samples[0]) + return self.cached + cols, rows = self.calc_cols_rows(batch_size, width, height) + + self.cached = result = Image.new("RGB", size=(width * cols, height * rows)) + for idx in range(batch_size): + result.paste( + Image.fromarray(samples[idx]), + box=((idx % cols) * width, ((idx // cols) % rows) * height), + ) + self.cached = result + return result + + def _decode_latent_to_preview_nocuda(self, x0): + return self.work_to_image(self._decode_latent(x0).cpu()) - def _decode_latent_to_preview_nocuda(self, x0): - return self.work_to_image(self._decode_latent(x0).cpu()) latent_preview.TAESDPreviewerImpl = BetterTAESDPreviewer diff --git a/py/hypertile.py b/py/hypertile.py index b7a7ec3..246152c 100644 --- a/py/hypertile.py +++ b/py/hypertile.py @@ -2,118 +2,205 @@ # Originally taken from: https://github.com/tfernd/HyperTile/ # Modified version of ComfyUI main code # https://github.com/comfyanonymous/ComfyUI/blob/master/comfy_extras/nodes_hypertile.py +from __future__ import annotations import math import random -import torch - from einops import rearrange -from comfy.ldm.modules.diffusionmodules.openaimodel import forward_timestep_embed, timestep_embedding, th, apply_control - -def get_closest_divisors(hw: int, aspect_ratio: float) -> tuple[int, int]: - pairs = [(i, hw // i) for i in range(int(math.sqrt(hw)), 1, -1) if hw % i == 0] - pair = min(((i, hw // i) for i in range(2, hw + 1) if hw % i == 0), - key=lambda x: abs(x[1] / x[0] - aspect_ratio)) - pairs.append(pair) - res = min(pairs, key=lambda x: max(x) / min(x)) - return res -def calc_optimal_hw(hw: int, aspect_ratio: float) -> tuple[int, int]: - hcand = round(math.sqrt(hw * aspect_ratio)) - wcand = hw // hcand +class HyperTile: + def __init__( + self, + model, + seed, + tile_size, + swap_size, + max_depth, + scale_depth, + start_step, + end_step, + ): + self.rng = random.Random() + self.rng.seed(seed) + self.model = model + self.latent_tile_size = max(32, tile_size) // 8 + self.swap_size = swap_size + self.max_depth = max_depth + self.scale_depth = scale_depth + self.start_step = start_step + self.end_step = end_step + # Temporary storage for rearranged tensors in the output part + self.temp = None + + def patch(self): + model = self.model + model.set_model_attn1_patch(self.attn1_in) + model.set_model_attn1_output_patch(self.attn1_out) + return model + + @staticmethod + def get_closest_divisors(hw: int, aspect_ratio: float) -> tuple[int, int]: + pairs = tuple( + (i, hw // i) for i in range(int(math.sqrt(hw)), 1, -1) if hw % i == 0 + ) + pair = min( + ((i, hw // i) for i in range(2, hw + 1) if hw % i == 0), + key=lambda x: abs(x[1] / x[0] - aspect_ratio), + ) + pairs.append(pair) + return min(pairs, key=lambda x: max(x) / min(x)) + + @classmethod + def calc_optimal_hw(cls, hw: int, aspect_ratio: float) -> tuple[int, int]: + hcand = round(math.sqrt(hw * aspect_ratio)) + wcand = hw // hcand + + if hcand * wcand == hw: + return hcand, wcand - if hcand * wcand != hw: wcand = round(math.sqrt(hw / aspect_ratio)) hcand = hw // wcand if hcand * wcand != hw: - return get_closest_divisors(hw, aspect_ratio) + return cls.get_closest_divisors(hw, aspect_ratio) - return hcand, wcand + return hcand, wcand -def random_divisor(value: int, min_value: int, /, max_options: int = 1, rand_obj=random.Random()) -> int: - min_value = min(min_value, value) + def random_divisor( + self, + value: int, + min_value: int, + /, + max_options: int = 1, + ) -> int: + min_value = min(min_value, value) + # All big divisors of value (inclusive) + divisors = tuple(i for i in range(min_value, value + 1) if value % i == 0) + ns = tuple(value // i for i in divisors[:max_options]) # has at least 1 element + idx = self.rng.randint(0, len(ns) - 1) if len(ns) - 1 > 0 else 0 + return ns[idx] - # All big divisors of value (inclusive) - divisors = [i for i in range(min_value, value + 1) if value % i == 0] + def check_timestep(self, extra_options): + current_timestep = self.model.model.model_sampling.timestep( + extra_options["sigmas"][0], + ).item() + return current_timestep <= self.start_step and current_timestep >= self.end_step - ns = [value // i for i in divisors[:max_options]] # has at least 1 element + def attn1_in(self, q, k, v, extra_options): + if not self.check_timestep(extra_options): + self.temp = None + return q, k, v - if len(ns) - 1 > 0: - idx = rand_obj.randint(0, len(ns) - 1) - else: - idx = 0 + model_chans = q.shape[-2] + orig_shape = extra_options["original_shape"] + + apply_to = tuple( + (orig_shape[-2] / (2**i)) * (orig_shape[-1] / (2**i)) + for i in range(self.max_depth + 1) + ) + if model_chans not in apply_to: + return q, k, v + + aspect_ratio = orig_shape[-1] / orig_shape[-2] + + hw = q.size(1) + h, w = ( + round(math.sqrt(hw * aspect_ratio)), + round(math.sqrt(hw / aspect_ratio)), + ) + + factor = (2 ** apply_to.index(model_chans)) if self.scale_depth else 1 + + nh = self.random_divisor(h, self.latent_tile_size * factor, self.swap_size) + nw = self.random_divisor(w, self.latent_tile_size * factor, self.swap_size) + + if nh * nw <= 1: + return q, k, v + + q = rearrange( + q, + "b (nh h nw w) c -> (b nh nw) (h w) c", + h=h // nh, + w=w // nw, + nh=nh, + nw=nw, + ) + self.temp = (nh, nw, h, w) + return q, k, v + + def attn1_out(self, out, _extra_options): + if self.temp is None: + return out + nh, nw, h, w = self.temp + self.temp = None + out = rearrange(out, "(b nh nw) hw c -> b nh nw hw c", nh=nh, nw=nw) + return rearrange( + out, + "b nh nw (h w) c -> b (nh h nw w) c", + h=h // nh, + w=w // nw, + ) - return ns[idx] class HyperTileBleh: @classmethod - def INPUT_TYPES(s): - return {"required": { - "model": ("MODEL",), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "tile_size": ("INT", {"default": 256, "min": 1, "max": 2048}), - "swap_size": ("INT", {"default": 2, "min": 1, "max": 128}), - "max_depth": ("INT", {"default": 0, "min": 0, "max": 10}), - "scale_depth": ("BOOLEAN", {"default": False}), - "start_step": ("INT", { "default": 0, "min": 0, "max": 1000, "step": 1, "display": "number" }), - "end_step": ("INT", { "default": 1000, "min": 0, "max": 1000, "step": 0.1, }), - }} + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}), + "tile_size": ("INT", {"default": 256, "min": 1, "max": 2048}), + "swap_size": ("INT", {"default": 2, "min": 1, "max": 128}), + "max_depth": ("INT", {"default": 0, "min": 0, "max": 10}), + "scale_depth": ("BOOLEAN", {"default": False}), + "start_step": ( + "INT", + { + "default": 0, + "min": 0, + "max": 1000, + "step": 1, + "display": "number", + }, + ), + "end_step": ( + "INT", + { + "default": 1000, + "min": 0, + "max": 1000, + "step": 0.1, + }, + ), + }, + } + RETURN_TYPES = ("MODEL",) FUNCTION = "patch" - CATEGORY = "bleh/model_patches" - def patch(self, model, seed, tile_size, swap_size, max_depth, scale_depth, start_step, end_step): - latent_tile_size = max(32, tile_size) // 8 - temp = None - - rand_obj = random.Random() - rand_obj.seed(seed) - - def hypertile_in(q, k, v, extra_options): - nonlocal temp - - current_timestep = model.model.model_sampling.timestep(extra_options['sigmas'][0]).item() - if current_timestep > start_step or current_timestep < end_step: - temp = None - return q, k, v - - model_chans = q.shape[-2] - orig_shape = extra_options['original_shape'] - - apply_to = [] - for i in range(max_depth + 1): - apply_to.append((orig_shape[-2] / (2 ** i)) * (orig_shape[-1] / (2 ** i))) - if model_chans not in apply_to: - return q, k, v - - aspect_ratio = orig_shape[-1] / orig_shape[-2] - - hw = q.size(1) - h, w = round(math.sqrt(hw * aspect_ratio)), round(math.sqrt(hw / aspect_ratio)) - - factor = (2 ** apply_to.index(model_chans)) if scale_depth else 1 - nh = random_divisor(h, latent_tile_size * factor, swap_size, rand_obj) - nw = random_divisor(w, latent_tile_size * factor, swap_size, rand_obj) - - if nh * nw > 1: - q = rearrange(q, "b (nh h nw w) c -> (b nh nw) (h w) c", h=h // nh, w=w // nw, nh=nh, nw=nw) - temp = (nh, nw, h, w) - return q, k, v - - def hypertile_out(out, extra_options): - nonlocal temp - if temp is not None: - nh, nw, h, w = temp - temp = None - out = rearrange(out, "(b nh nw) hw c -> b nh nw hw c", nh=nh, nw=nw) - out = rearrange(out, "b nh nw (h w) c -> b (nh h nw w) c", h=h // nh, w=w // nw) - return out - - m = model.clone() - m.set_model_attn1_patch(hypertile_in) - m.set_model_attn1_output_patch(hypertile_out) - return (m, ) + def patch( + self, + model, + seed, + tile_size, + swap_size, + max_depth, + scale_depth, + start_step, + end_step, + ): + return ( + HyperTile( + model.clone(), + seed, + tile_size, + swap_size, + max_depth, + scale_depth, + start_step, + end_step, + ).patch(), + ) diff --git a/py/settings.py b/py/settings.py index 5cbee0d..10793d0 100644 --- a/py/settings.py +++ b/py/settings.py @@ -1,29 +1,49 @@ -import json -import os from pathlib import Path -class Settings: - def __init__(self, obj = {}): - btp = obj.get("betterTaesdPreviews", None) - if btp is None: - self.btp_enabled = False - else: - self.btp_enabled = True - self.btp_max_size = btp.get("max_size", 768) - self.btp_max_batch = btp.get("max_batch", 4) - self.btp_max_batch_cols = btp.get("max_batch_cols", 2) - self.btp_throttle_secs = btp.get("throttle_secs", 1) - self.btp_use_cuda = btp.get("use_cuda", True) -SETTINGS = None +class Settings: + def __init__(self): + self.btp_enabled = False + + def update(self, obj): + btp = obj.get("betterTaesdPreviews", None) + if btp is None: + self.btp_enabled = False + else: + self.btp_enabled = True + self.btp_max_size = btp.get("max_size", 768) + self.btp_max_batch = btp.get("max_batch", 4) + self.btp_max_batch_cols = btp.get("max_batch_cols", 2) + self.btp_throttle_secs = btp.get("throttle_secs", 1) + self.btp_use_cuda = btp.get("use_cuda", True) + + def get_cfg_path(self, filename): + my_path = Path.resolve(Path(__file__).parent) + return my_path.parent / filename + + def try_update_from_json(self, filename): + import json + + try: + with Path.open(self.get_cfg_path(filename)) as fp: + self.update(json.load(fp)) + except OSError: + return False + + def try_update_from_yaml(self, filename): + try: + import yaml + + with Path.open(self.get_cfg_path(filename)) as fp: + self.update(yaml.safe_load(fp)) + except (OSError, ImportError): + return False + + +SETTINGS = Settings() + def load_settings(): - global SETTINGS - my_path = Path(os.path.abspath(os.path.dirname(__file__))) - cfg_path = my_path.parent / "blehconfig.json" - try: - with open(cfg_path, "r") as fp: - SETTINGS = Settings(json.load(fp)) - except OSError: - SETTINGS = Settings() - return SETTINGS + if not SETTINGS.try_update_from_yaml("blehconfig.yaml"): + SETTINGS.try_update_from_json("blehconfig.json") + return SETTINGS diff --git a/ruff.toml b/ruff.toml new file mode 100644 index 0000000..bcbcdda --- /dev/null +++ b/ruff.toml @@ -0,0 +1,34 @@ +[lint] +ignore = [ + "ANN001", + "ANN101", + "ANN102", + "ANN201", + "ANN202", + "ANN204", + "ANN206", + "C901", + "D100", + "D101", + "D102", + "D103", + "D104", + "D107", + "D211", + "D213", + "E402", + "E501", + "EM101", + "ERA001", + "F403", + "F405", + "PLR0912", + "PLR0913", + "PLR0915", + "PLR2004", + "T201", + "TRY003", + "N802", + "N999", +] +select = ["ALL"]