Refactoring and cleanups

This commit is contained in:
blepping
2024-01-28 06:17:05 -07:00
parent 0bcd6c5f85
commit bef8ff691c
6 changed files with 353 additions and 185 deletions
+3 -3
View File
@@ -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
+4 -3
View File
@@ -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"]
+91 -65
View File
@@ -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
+177 -90
View File
@@ -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(),
)
+44 -24
View File
@@ -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
+34
View File
@@ -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"]