Add BlehBlockCFG node

Add tooltips and descriptions for most nodes

Improvements to TAESD previews

Various cleanups and lint squashing

Add more scaling and blend types

Hopefully improve latent normalization (may change seeds)
This commit is contained in:
blepping
2024-08-30 11:21:26 -06:00
parent 810cfd13f0
commit c98990a3cb
18 changed files with 876 additions and 191 deletions
+19 -2
View File
@@ -13,6 +13,7 @@ A ComfyUI nodes collection... eventually.
7. Ensure a seed is set even when `add_noise` is turned off in a sampler. Yes, that's right: if you don't have `add_noise` enabled _no_ seed gets set for samplers like `euler_a` and it's not possible to reproduce generations. (look for the [BlehForceSeedSampler](#blehforceseedsampler) node). For `SamplerCustomAdvanced` you can use `BlehDisableNoise` to accomplish the same thing.
8. Allows swapping to a refiner model at a predefined time (look for the [BlehRefinerAfter](#blehrefinerafter) node).
9. Allow defining arbitrary model patches (look for the [BlehBlockOps](#blehblockops) node).
10. Experimental blockwise CFG type effect (look for the [BlehBlockCFG](#blehblockcfg) node).
## Configuration
@@ -32,12 +33,14 @@ Current defaults:
|-|-|-|
|`enabled`|`true`|Toggles whether enhanced TAESD previews are enabled|
|`max_size`|`768`|Max width or height for previews. Note this does not affect TAESD decoding, just the preview image|
|`max_width`|`max_size`|Same as `max_size` except allows setting the width independently. Previews may not work well with non-square max dimensions.|
|`max_height`|`max_size`|Same as `max_size` except allows setting the height independently. Previews may not work well with non-square max dimensions.|
|`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`|`2`|Max frequency to decode the latents for previewing. `0.25` would be every quarter second, `2` would be once every two seconds|
|`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.|
|`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.|
|`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.|
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.
@@ -125,6 +128,20 @@ Allows switching to a refiner model at a predefined time. There are three time m
you likely can only use this to swap between models that are closely related. For example, switching from SD 1.5 to
SDXL is not going to work at all.
### BlehBlockCFG
Experimental model patch that attempts to guide either `cond` (positive prompt) or `uncond` (negative prompt) away from its opposite.
In other words, when applied to `cond` it will try to push it further away from what `uncond` is doing and vice versa. Stronger effect when
applied to `cond` or output blocks. The defaults are reasonable for SD 1.5 (or as reasonable as weird stuff like this can be).
Enter comma separated blocks numbers starting with one of **I**input, **O**utput or **M**iddle like `i4,m0,o4`.
The patch can be applied to the same model multiple times.
Is it good, or even doing what I think? Who knows!
_Note_: Probably only works with SD 1.x and SDXL.
### BlehBlockOps
Very experimental advanced node that allows defining model patches using YAML. This node is still under development and may be changed.
+5 -1
View File
@@ -37,4 +37,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"BlehDeepShrink": "Kohya Deep Shrink (bleh)",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "BLEH_VERSION"]
__all__ = ("BLEH_VERSION", "NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS")
from .py.nodes import blockCFG
NODE_CLASS_MAPPINGS["BlehBlockCFG"] = blockCFG.BlockCFGBleh
+33
View File
@@ -0,0 +1,33 @@
# Copy this file to blehconfig.yaml
betterTaesdPreviews:
# If disabled, will use the old ComfyUI previewer.
enabled: true
# Maximum preview size (applies to both height and width).
max_size: 768
# Maximum preview width. If set, will override max_size.
max_width: 768
# Maximum preview height. If set, will override max_size.
max_height: 768
# Maximum batch items to preview.
max_batch: 4
# Maximum columns to use when previewing batches.
max_batch_cols: 2
# Minimum time between updating previews. The default will update the preview at most once per second.
throttle_secs: 1
# When enabled and previewing batches, you will see previews spread across the batch. Otherwise it will be the first max_batch items.
maxed_batch_step_mode: false
# Allows overriding the preview device, for example you could set it to "cpu". Note: Generally should be left
# alone unless you know you need to change it. Previewing on CPU will likely be quite slow.
preview_device: null
# 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
-9
View File
@@ -1,9 +0,0 @@
betterTaesdPreviews:
enabled: true
max_size: 768
max_batch: 4
max_batch_cols: 2
throttle_secs: 1
maxed_batch_step_mode: false
preview_device: null
skip_upscale_layers: 0
+7
View File
@@ -2,6 +2,13 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20240830
* Added the `BlehBlockCFG` node (see README for usage and details).
* More scaling/blending types. Some of them don't work well with scaling and will be filtered, you can set the environment variable `COMFYUI_BLEH_OVERRIDE_NO_SCALE` if you want the full list to be available (but you might just get garbage if you try to use them for scaling).
* Possibly better normalization function (may change seeds). Set the environment variable `COMFYUI_BLEH_ORIG_NORMALIZE` to disable.
* TAESD previews should be faster. Also now can dynamically set the number of upscale layers to skip based on the preview size limits. Additionally it's possible to set the max preview width/height seperately - see the YAML example config.
## 20240506
* Add many new scaling types.
+104 -35
View File
@@ -2,8 +2,8 @@ import math
from time import time
import latent_preview
import numpy as np
import torch
from comfy.model_management import device_supports_non_blocking
from PIL import Image
from .settings import SETTINGS
@@ -14,14 +14,6 @@ _ORIG_PREVIEWER = latent_preview.TAESDPreviewerImpl
class BetterTAESDPreviewer(_ORIG_PREVIEWER):
def __init__(self, taesd):
del taesd.taesd_encoder
if SETTINGS.btp_skip_upscale_layers > 0:
upscale_layers = tuple(
idx
for idx, layer in enumerate(taesd.taesd_decoder)
if isinstance(layer, torch.nn.Upsample)
)
for idx in range(1, min(SETTINGS.btp_skip_upscale_layers, 3) + 1):
taesd.taesd_decoder.pop(upscale_layers[-idx])
self.device = (
None
if SETTINGS.btp_preview_device is None
@@ -33,46 +25,113 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
self.stamp = None
self.cached = None
self.blank = Image.new("RGB", size=(1, 1))
self.stream = None
self.prev_work = None
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
self.throttle_secs = SETTINGS.btp_throttle_secs
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.maybe_pop_upscale_layers()
def maybe_pop_upscale_layers(self, *, width=None, height=None):
skip = self.skip_upscale_layers
if skip == 0:
return
upscale_layers = tuple(
idx
for idx, layer in enumerate(self.taesd.taesd_decoder)
if isinstance(layer, torch.nn.Upsample)
)
num_upscale_layers = len(upscale_layers)
if skip < 0:
if width is None or height is None:
return
aggressive = skip == -2
skip = 0
max_width, max_height = self.preview_max_width, self.preview_max_height
while skip < num_upscale_layers and (
width > max_width or height > max_height
):
width //= 2
height //= 2
if not aggressive and width < max_width and height < max_height:
# Popping another would overshoot the size requirement.
break
skip += 1
if not aggressive and (width <= max_width or height <= max_height):
# At least one dimension is within the size requirement.
break
if skip > 0:
skip = min(skip, num_upscale_layers)
for idx in range(1, skip + 1):
self.taesd.taesd_decoder.pop(upscale_layers[-idx])
self.skip_upscale_layers = 0
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),
min(
max(*preview_image.size),
max(self.preview_max_width, self.preview_max_height),
),
)
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:
) and now - self.stamp < self.throttle_secs:
return True
self.stamp = now
return False
def _decode_latent(self, x0):
max_batch = SETTINGS.btp_max_batch
batch_size = x0.shape[0]
if not SETTINGS.btp_maxed_batch_step_mode:
indexes = range(min(max_batch, batch_size))
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_size,
math.ceil(batch_size / max_batch),
batch,
math.ceil(batch / max_batch),
)[:max_batch]
x0 = x0[indexes, :]
batch, _channels, height, width = x0.shape
if self.device and x0.device != self.device:
x0 = x0.to(self.device)
samples = (self.taesd.decode(x0) + 1.0) / 2.0
samples = torch.clamp(samples, min=0.0, max=1.0) * 255.0
return samples.to(dtype=torch.uint8).detach()
x0 = x0.to(
device=self.device,
non_blocking=device_supports_non_blocking(x0.device),
)
cols, rows = self.calc_cols_rows(
min(batch, self.max_batch_preview),
width,
height,
)
if self.skip_upscale_layers < 0:
self.maybe_pop_upscale_layers(
width=width * 8 * cols,
height=height * 8 * rows,
)
return (
(
self.taesd.decode(x0)
.movedim(1, -1)
.add_(1)
.mul_(0.5)
.clamp_(min=0, max=1)
.mul_(255)
.detach()
),
cols,
rows,
)
def calc_cols_rows(self, batch_size, width, height):
max_cols = SETTINGS.btp_max_batch_cols
max_cols = self.max_batch_cols
ratio = height / width
if ratio >= 1.45:
# Very tall images - prioritize horizontal layout.
@@ -85,28 +144,38 @@ class BetterTAESDPreviewer(_ORIG_PREVIEWER):
rows = math.ceil(batch_size / cols)
return cols, rows
def decoded_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:
def decoded_to_image(self, samples, cols, rows):
batch, width, height = samples.shape[:-1]
samples = samples.to(
device="cpu",
dtype=torch.uint8,
non_blocking=device_supports_non_blocking(samples.device),
).numpy()
if batch == 1:
self.cached = Image.fromarray(samples[0])
return self.cached
cols, rows = self.calc_cols_rows(batch_size, width, height)
cols, rows = self.calc_cols_rows(batch, width, height)
self.cached = result = Image.new("RGB", size=(width * cols, height * rows))
for idx in range(batch_size):
img_size = (width * cols, height * rows)
if self.cached is not None and self.cached.size == img_size:
result = self.cached
else:
self.cached = result = Image.new("RGB", size=(width * cols, height * rows))
for idx in range(batch):
result.paste(
Image.fromarray(samples[idx]),
box=((idx % cols) * width, ((idx // cols) % rows) * height),
)
self.cached = result
return result
def decode_latent_to_preview(self, x0):
if self.check_use_cached():
return self.cached
return self.decoded_to_image(self._decode_latent(x0).cpu())
if x0.shape[0] == 0:
return self.blank # Shouldn't actually be possible.
return self.decoded_to_image(*self._decode_latent(x0))
latent_preview.TAESDPreviewerImpl = BetterTAESDPreviewer
if not isinstance(latent_preview.TAESDPreviewerImpl, BetterTAESDPreviewer):
latent_preview.BLEH_ORIG_TAESDPreviewerImpl = _ORIG_PREVIEWER
latent_preview.TAESDPreviewerImpl = BetterTAESDPreviewer
+266 -41
View File
@@ -1,15 +1,20 @@
# Credits:
# Blending, slice and filtering functions based on https://github.com/WASasquatch/FreeU_Advanced
# Blending, slice and filtering functions based on https://github.com/WASasquatch/FreeU_Advanced
from __future__ import annotations
import math
import os
import kornia.filters as kf
import numpy as np
import torch
from torch import fft
from torch import FloatTensor, LongTensor, fft
OVERRIDE_NO_SCALE = "COMFYUI_BLEH_OVERRIDE_NO_SCALE" in os.environ
USE_ORIG_NORMALIZE = "COMFYUI_BLEH_ORIG_NORMALIZE" in os.environ
def normalize(latent, target_min=None, target_max=None):
def normalize_orig(latent, target_min=None, target_max=None, **_unused_kwargs: dict):
min_val = latent.min()
max_val = latent.max()
@@ -22,6 +27,26 @@ def normalize(latent, target_min=None, target_max=None):
return normalized * (target_max - target_min) + target_min
def normalize(latent, *, reference_latent=None, dim=(-3, -2, -1)):
if reference_latent is None:
return latent
min_val, max_val = (
latent.amin(dim=dim, keepdim=True),
latent.amax(dim=dim, keepdim=True),
)
target_min, target_max = (
reference_latent.amin(dim=dim, keepdim=True),
reference_latent.amax(dim=dim, keepdim=True),
)
normalized = (latent - min_val) / (max_val - min_val)
return normalized * (target_max - target_min) + target_min
if USE_ORIG_NORMALIZE:
normalize = normalize_orig
def hslerp(a, b, t):
if a.shape != b.shape:
raise ValueError("Input tensors a and b must have the same shape.")
@@ -36,15 +61,10 @@ def hslerp(a, b, t):
device=a.device,
dtype=a.dtype,
)
interpolation_tensor[0, 0, 0, 0] = 1.0
interpolation_tensor[0, 0, 0, 0] = 1.0 if t < 0.5 else -1.0
result = (1 - t) * a + t * b
norm = (torch.norm(b - a, dim=1, keepdim=True) / 6) * interpolation_tensor
if t < 0.5:
result += norm
else:
result -= norm
result += (torch.norm(b - a, dim=1, keepdim=True) / 6) * interpolation_tensor
return result
@@ -66,6 +86,18 @@ def hslerp_alt(a, b, t):
return result.add_(norm)
# This should be more correct but the results are worse. :(
def hslerp_alt_(a, b, t):
if a.shape != b.shape:
raise ValueError("Input tensors a and b must have the same shape.")
t_expanded = t.broadcast_to(a.shape[-2:])
while t_expanded.ndim < a.ndim:
t_expanded = t_expanded.unsqueeze(0)
interp = torch.where(t_expanded < 0.5, 1.0, -1.0)
result = (1 - t) * a + t * b
return result.add_((torch.norm(b - a, dim=1, keepdim=True) / 6) * interp)
# Copied from ComfyUI
def slerp_orig(b1, b2, r):
c = b1.shape[-1]
@@ -101,36 +133,233 @@ def slerp_orig(b1, b2, r):
return res
# From https://gist.github.com/Birch-san/230ac46f99ec411ed5907b0a3d728efa
def altslerp( # noqa: PLR0914
v0: FloatTensor,
v1: FloatTensor,
t: float | FloatTensor,
dot_threshold=0.9995,
dim=-1,
):
# Normalize the vectors to get the directions and angles
v0_norm: FloatTensor = torch.linalg.norm(v0, dim=dim)
v1_norm: FloatTensor = torch.linalg.norm(v1, dim=dim)
v0_normed: FloatTensor = v0 / v0_norm.unsqueeze(dim)
v1_normed: FloatTensor = v1 / v1_norm.unsqueeze(dim)
# Dot product with the normalized vectors
dot: FloatTensor = (v0_normed * v1_normed).sum(dim)
dot_mag: FloatTensor = dot.abs()
# if dp is NaN, it's because the v0 or v1 row was filled with 0s
# If absolute value of dot product is almost 1, vectors are ~colinear, so use lerp
gotta_lerp: LongTensor = dot_mag.isnan() | (dot_mag > dot_threshold)
can_slerp: LongTensor = ~gotta_lerp
t_batch_dim_count: int = (
max(0, t.dim() - v0.dim()) if isinstance(t, torch.Tensor) else 0
)
t_batch_dims: torch.Size = (
t.shape[:t_batch_dim_count] if isinstance(t, torch.Tensor) else torch.Size([])
)
out: FloatTensor = torch.zeros_like(v0.expand(*t_batch_dims, *(dim,) * v0.dim()))
# if no elements are lerpable, our vectors become 0-dimensional, preventing broadcasting
if gotta_lerp.any():
lerped: FloatTensor = torch.lerp(v0, v1, t)
out: FloatTensor = lerped.where(gotta_lerp.unsqueeze(dim), out)
# if no elements are slerpable, our vectors become 0-dimensional, preventing broadcasting
if can_slerp.any():
# Calculate initial angle between v0 and v1
theta_0: FloatTensor = dot.arccos().unsqueeze(dim)
sin_theta_0: FloatTensor = theta_0.sin()
# Angle at timestep t
theta_t: FloatTensor = theta_0 * t
sin_theta_t: FloatTensor = theta_t.sin()
# Finish the slerp algorithm
s0: FloatTensor = (theta_0 - theta_t).sin() / sin_theta_0
s1: FloatTensor = sin_theta_t / sin_theta_0
slerped: FloatTensor = s0 * v0 + s1 * v1
out: FloatTensor = slerped.where(can_slerp.unsqueeze(dim), out)
return out
class BlendMode:
__slots__ = ("allow_scale", "f", "norm", "norm_dims", "rev")
class _Empty:
pass
def __init__(
self,
f,
norm=None,
norm_dims=(-3, -2, -1),
rev=False,
allow_scale=True,
):
self.f = f
self.norm = norm
self.norm_dims = norm_dims
self.rev = rev
self.allow_scale = allow_scale
def edited(
self,
*,
f=_Empty,
norm=_Empty,
norm_dims=_Empty,
rev=_Empty,
allow_scale=_Empty,
):
empty = self._Empty
return self.__class__(
f if f is not empty else self.f,
norm=norm if norm is not empty else self.norm,
norm_dims=norm_dims if norm_dims is not empty else self.norm_dims,
rev=rev if rev is not empty else self.rev,
allow_scale=allow_scale if allow_scale is not empty else self.allow_scale,
)
def __call__(self, a, b, t):
if self.rev:
a, b = b, a
if self.norm is None:
return self.f(a, b, t)
ref = (1 - t) * a + t * b
return self.norm(self.f(a, b, t), reference_latent=ref, dim=self.norm_dims)
BLENDING_MODES = {
# Args:
# - a (tensor): Latent input 1
# - b (tensor): Latent input 2
# - t (float): Blending factor
# Interpolates between tensors a and b using normalized linear interpolation.
"bislerp": lambda a, b, t: normalize((1 - t) * a + t * b),
"bislerp": BlendMode(lambda a, b, t: (1 - t) * a + t * b),
# "nbislerp": BlendMode(lambda a, b, t: (1 - t) * a + t * b, normalize),
"slerp": BlendMode(lambda a, b, t: altslerp(a, b, t, dim=0)),
# Transfer the color from `b` to `a` by t` factor
"colorize": lambda a, b, t: a + (b - a) * t,
"colorize": BlendMode(lambda a, b, t: a + (b - a) * t),
# Interpolates between tensors a and b using cosine interpolation.
"cosinterp": lambda a, b, t: (
a + b - (a - b) * torch.cos(t * torch.tensor(math.pi))
)
/ 2,
"cosinterp": BlendMode(
lambda a, b, t: (a + b - (a - b) * torch.cos(t * torch.tensor(math.pi))) / 2,
),
# Interpolates between tensors a and b using cubic interpolation.
"cuberp": lambda a, b, t: a + (b - a) * (3 * t**2 - 2 * t**3),
"cuberp": BlendMode(lambda a, b, t: a + (b - a) * (3 * t**2 - 2 * t**3)),
# Interpolates between tensors a and b using normalized linear interpolation,
# with a twist when t is greater than or equal to 0.5.
"hslerp": hslerp,
"hslerp": BlendMode(hslerp),
# Adds tensor b to tensor a, scaled by t.
"inject": lambda a, b, t: a + b * t,
"inject": BlendMode(lambda a, b, t: a + b * t),
# Interpolates between tensors a and b using linear interpolation.
"lerp": lambda a, b, t: (1 - t) * a + t * b,
"lerp": BlendMode(lambda a, b, t: (1 - t) * a + t * b),
# Simulates a brightening effect by adding tensor b to tensor a, scaled by t.
"lineardodge": lambda a, b, t: normalize(a + b * t),
"lineardodge": BlendMode(lambda a, b, t: a + b * t),
# "nlineardodge": BlendMode(lambda a, b, t: a + b * t, normalize),
# Simulates a brightening effect by dividing a by (1 - b) with a small epsilon to avoid division by zero.
"colordodge": BlendMode(lambda a, b, _t: a / (1 - b + 1e-6), allow_scale=False),
"difference": BlendMode(
lambda a, b, t: abs(a - b) * t,
normalize,
allow_scale=False,
),
"exclusion": BlendMode(
lambda a, b, t: (a + b - 2 * a * b) * t,
normalize,
allow_scale=False,
),
"glow": BlendMode(
lambda a, b, _t: torch.where(
a <= 1,
a**2 / (1 - b + 1e-6),
b * (a - 1) / (a + 1e-6),
),
allow_scale=False,
),
"hardlight": BlendMode(
lambda a, b, t: (
2 * a * b * (a < 0.5).float()
+ (1 - 2 * (1 - a) * (1 - b)) * (a >= 0.5).float()
)
* t,
allow_scale=False,
),
"linearlight": BlendMode(
lambda a, b, _t: torch.where(b <= 0.5, a + 2 * b - 1, a + 2 * (b - 0.5)),
),
"multiply": BlendMode(
lambda a, b, t: a * t * b * t,
normalize,
allow_scale=False,
),
"overlay": BlendMode(
lambda a, b, t: (2 * a * b + a**2 - 2 * a * b * a) * t
if torch.all(b < 0.5)
else (1 - 2 * (1 - a) * (1 - b)) * t,
allow_scale=False,
),
# Combines tensors a and b using the Pin Light formula.
"pinlight": BlendMode(
lambda a, b, _t: torch.where(
b <= 0.5,
torch.min(a, 2 * b),
torch.max(a, 2 * b - 1),
),
),
"reflect": BlendMode(
lambda a, b, _t: torch.where(
b <= 1,
b**2 / (1 - a + 1e-6),
a * (b - 1) / (b + 1e-6),
),
allow_scale=False,
),
"screen": BlendMode(
lambda a, b, t: 1 - (1 - a) * (1 - b) * (1 - t),
allow_scale=False,
),
"subtract": BlendMode(lambda a, b, t: a * t - b * t, allow_scale=False),
"vividlight": BlendMode(
lambda a, b, _t: torch.where(
b <= 0.5,
a / (1 - 2 * b + 1e-6),
(a + 2 * b - 1) / (2 * (1 - b) + 1e-6),
),
allow_scale=False,
),
}
BLENDING_MODES |= {
f"norm{k}": v.edited(norm=normalize)
for k, v in BLENDING_MODES.items()
if k != "hslerp" and v.norm is None
}
BLENDING_MODES |= {f"rev{k}": v.edited(rev=True) for k, v in BLENDING_MODES.items()}
BIDERP_MODES = {
k: v.edited(norm_dims=0)
for k, v in BLENDING_MODES.items()
if (v.allow_scale or OVERRIDE_NO_SCALE) and not k.endswith("slerp")
}
BIDERP_MODES |= {
"hslerp": hslerp_alt,
"bislerp": slerp_orig,
"altbislerp": altslerp,
"revaltbislerp": lambda a, b, t: altslerp(b, a, t),
"bibislerp": BLENDING_MODES["bislerp"].edited(norm_dims=0),
"revhslerp": lambda a, b, t: hslerp_alt(b, a, t),
"revbislerp": lambda a, b, t: slerp_orig(b, a, t),
"revbibislerp": BLENDING_MODES["revbislerp"].edited(norm_dims=0),
}
for k in tuple(BLENDING_MODES.keys()):
if k == "hslerp":
continue
BLENDING_MODES[f"rev{k}"] = lambda a, b, t, f=BLENDING_MODES[k]: f(b, a, t)
FILTER_PRESETS = {
"none": (),
@@ -183,27 +412,22 @@ FILTER_PRESETS = {
}
BIDERP_MODES = {k: v for k, v in BLENDING_MODES.items() if not k.endswith("slerp")}
BIDERP_MODES |= {
"hslerp": hslerp_alt,
"bislerp": slerp_orig,
"bibislerp": BLENDING_MODES["bislerp"],
"revhslerp": lambda a, b, t, f=hslerp_alt: f(b, a, t),
"revbislerp": lambda a, b, t, f=slerp_orig: f(b, a, t),
"revbibislerp": BLENDING_MODES["revbislerp"],
}
ENHANCE_METHODS = (
"lowpass",
"multilowpass",
"highpass",
"multihighpass",
"bandpass",
"randhilowpass",
"randmultihilowpass",
"randhibandpass",
"randlowbandpass",
"gaussianblur",
"multigaussianblur",
"edge",
"multiedge",
"sharpen",
"multisharpen",
"korniabilateralblur",
"korniagaussianblur",
"korniasharpen",
@@ -252,7 +476,7 @@ FILTER_SIZES = (
def make_filter(channels, dtype, size=3):
a = FILTER_SIZES[size - 1]
filt = torch.tensor(a[:, None] * a[None, :], dtype=dtype)
filt = filt / torch.sum(filt)
filt /= torch.sum(filt)
return filt[None, None, :, :].repeat((channels, 1, 1, 1))
@@ -262,7 +486,7 @@ def antialias_tensor(x, antialias_size):
return torch.nn.functional.conv2d(x, filt, groups=channels, padding="same")
def enhance_tensor(
def enhance_tensor( # noqa: PLR0911
x,
name,
scale=1.0,
@@ -350,6 +574,7 @@ def scale_samples(
samples,
width,
height,
*,
mode="bicubic",
mode_h=None,
antialias_size=0,
@@ -371,11 +596,11 @@ def scale_samples(
dtype=torch.uint8,
).tolist()
mode, mode_h = (
m if mode not in ("random", "randomaa") else RAND_UPSCALE_METHODS[ridx]
m if mode not in {"random", "randomaa"} else RAND_UPSCALE_METHODS[ridx]
for ridx, m in zip(ridxs, (mode, mode_h))
)
mode_h = mode
if mode in ("bicubic", "nearest-exact", "bilinear", "area"):
if mode in {"bicubic", "nearest-exact", "bilinear", "area"}:
result = torch.nn.functional.interpolate(
samples,
size=(height, width),
@@ -397,7 +622,7 @@ def scale_samples(
# Modified from ComfyUI
def biderp(samples, width, height, mode="bislerp", mode_h=None):
def biderp(samples, width, height, mode="bislerp", mode_h=None): # noqa: PLR0914
if mode_h is None:
mode_h = mode
@@ -494,7 +719,7 @@ def ffilter(x, threshold, scale, scales=None, strength=1.0):
crow - scale_threshold : crow + scale_threshold,
ccol - scale_threshold : ccol + scale_threshold,
] = scaled_scale_value
mask = mask + (scale_mask - mask) * strength
mask += (scale_mask - mask) * strength
x_freq *= mask
+192
View File
@@ -0,0 +1,192 @@
from functools import partial
class BlockCFGBleh:
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "bleh/model_patches"
DESCRIPTION = (
"Applies a CFG type effect to the model blocks themselves during evaluation."
)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (
"MODEL",
{
"tooltip": "Model to patch",
},
),
"commasep_block_numbers": (
"STRING",
{
"default": "i4,m0,o4",
"tooltip": "Comma separated list of block numbers, each should start with one of i(input), m(iddle), o(utput)",
},
),
"scale": (
"FLOAT",
{
"default": 0.25,
"min": -100.0,
"max": 100.0,
"step": 0.001,
"round": False,
"tooltip": "Effect strength",
},
),
"start_percent": (
"FLOAT",
{
"default": 0.2,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "Start time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"end_percent": (
"FLOAT",
{
"default": 0.8,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "End time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"skip_mode": (
"BOOLEAN",
{
"default": False,
"tooltip": "For output blocks, this causes the effect to apply to the skip connection. For input blocks it patches after the skip connection. No effect for middle blocks.",
},
),
"apply_to": (
("cond", "uncond"),
{
"default": "uncond",
"tooltip": "Guides the specified target away from its opposite. cond=positive prompt, uncond=negative prompt.",
},
),
},
}
@classmethod
def patch(
cls,
*,
model,
commasep_block_numbers,
scale,
start_percent,
end_percent,
skip_mode,
apply_to,
):
input_blocks = {}
middle_blocks = {}
output_blocks = {}
for item_ in commasep_block_numbers.split(","):
item = item_.strip().lower()
if not item:
continue
block_type = item[0]
if block_type not in "imo":
raise ValueError
block, *tidx = item[1:].split(".", 1)
block = int(block)
tidx = int(tidx) if tidx else -1
if block_type == "i":
bd = input_blocks
else:
bd = output_blocks if block_type == "o" else middle_blocks
bd[block] = tidx
if (
scale == 0
or end_percent <= 0
or start_percent >= 1
or not (input_blocks or middle_blocks or output_blocks)
):
return (model,)
ms = model.get_model_object("model_sampling")
sigma_start = ms.percent_to_sigma(start_percent)
sigma_end = ms.percent_to_sigma(end_percent)
reverse = apply_to != "cond"
def check_applies(block_list, transformer_options):
block_num = transformer_options["block"][1]
sigma_tensor = transformer_options["sigmas"].max()
sigma = sigma_tensor.detach().cpu().item()
block_def = block_list.get(block_num)
return (
sigma_end <= sigma <= sigma_start
and block_def is not None
and (
block_def == -1
or block_def == transformer_options.get("transformer_index")
)
)
def apply_cfg_fun(tensor, primary_offset):
secondary_offset = 0 if primary_offset == 1 else 1
if reverse:
primary_offset, secondary_offset = secondary_offset, primary_offset
result = tensor.clone()
batch = tensor.shape[0] // 2
primary_idxs, secondary_idxs = (
tuple(range(batch * offs, batch + batch * offs))
for offs in (primary_offset, secondary_offset)
)
# print(f"\nIDXS: cond={primary_idxs}, uncond={secondary_idxs}")
result[primary_idxs, ...] = (
tensor[primary_idxs, ...]
+ (tensor[primary_idxs, ...] - tensor[secondary_idxs, ...]) * scale
)
return result
def non_output_block_patch(h, transformer_options, *, block_list):
cond_or_uncond = transformer_options["cond_or_uncond"]
if len(cond_or_uncond) != 2 or not check_applies(
block_list,
transformer_options,
):
return h
return apply_cfg_fun(h, cond_or_uncond[0])
def output_block_patch(h, hsp, transformer_options, *, block_list):
cond_or_uncond = transformer_options["cond_or_uncond"]
if len(cond_or_uncond) != 2 or not check_applies(
block_list,
transformer_options,
):
return h, hsp
return (
(apply_cfg_fun(h, cond_or_uncond[0]), hsp)
if not skip_mode
else (h, apply_cfg_fun(hsp, cond_or_uncond[0]))
)
m = model.clone()
if input_blocks:
(
m.set_model_input_block_patch
if skip_mode
else m.set_model_input_block_patch_after_skip
)(
partial(non_output_block_patch, block_list=input_blocks),
)
if middle_blocks:
m.set_model_patch(
partial(non_output_block_patch, block_list=middle_blocks),
"middle_block_patch",
)
if output_blocks:
m.set_model_output_block_patch(
partial(output_block_patch, block_list=output_blocks),
)
return (m,)
+92 -36
View File
@@ -1,7 +1,5 @@
# Adapted from the ComfyUI built-in node
import bisect
from .. import latent_utils # noqa: TID252
@@ -9,6 +7,7 @@ class DeepShrinkBleh:
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "bleh/model_patches"
DESCRIPTION = "Model patch that enables generating at higher resolution than the model was trained for by downscaling the image near the start of generation."
upscale_methods = (
"bicubic",
@@ -22,39 +21,101 @@ class DeepShrinkBleh:
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"model": (
"MODEL",
{
"tooltip": "Model to patch",
},
),
"commasep_block_numbers": (
"STRING",
{
"default": "3",
"tooltip": "A comma separated list of input block numbers, the default should work for SD 1.5 and SDXL.",
},
),
"downscale_factor": (
"FLOAT",
{"default": 2.0, "min": 1.0, "max": 32.0, "step": 0.1},
{
"default": 2.0,
"min": 1.0,
"max": 32.0,
"step": 0.1,
"tooltip": "Controls how much the block will get downscaled while the effect is active.",
},
),
"start_percent": (
"FLOAT",
{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001},
{
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "Start time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"start_fadeout_percent": (
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001},
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "When enabled, the downscale_factor will fade out such that at end_percent it will be around 1.0 (no downscaling). May reduce artifacts... or cause them!",
},
),
"end_percent": (
"FLOAT",
{"default": 0.35, "min": 0.0, "max": 1.0, "step": 0.001},
{
"default": 0.35,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "End time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"downscale_after_skip": (
"BOOLEAN",
{
"default": True,
"tooltip": "Controls whether the downscale effect occurs after the skip conection. Generally should be left enabled.",
},
),
"downscale_method": (
latent_utils.UPSCALE_METHODS,
{
"default": "bicubic",
"tooltip": "Mode used for downscaling. Bicubic is generally a safe choice.",
},
),
"upscale_method": (
latent_utils.UPSCALE_METHODS,
{
"default": "bicubic",
"tooltip": "Mode used for upscaling. Bicubic is generally a safe choice.",
},
),
"antialias_downscale": (
"BOOLEAN",
{
"default": False,
"tooltip": "Experimental option to anti-alias (smooth) the latent after downscaling.",
},
),
"antialias_upscale": (
"BOOLEAN",
{
"default": False,
"tooltip": "Experimental option to anti-alias (smooth) the latent after upscaling.",
},
),
"downscale_after_skip": ("BOOLEAN", {"default": True}),
"downscale_method": (latent_utils.UPSCALE_METHODS,),
"upscale_method": (latent_utils.UPSCALE_METHODS,),
"antialias_downscale": ("BOOLEAN", {"default": False}),
"antialias_upscale": ("BOOLEAN", {"default": False}),
},
}
@classmethod
def patch(
self,
cls,
*,
model,
commasep_block_numbers,
downscale_factor,
@@ -75,44 +136,35 @@ class DeepShrinkBleh:
raise ValueError(
"BlehDeepShrink: Bad value for block numbers: must be comma-separated list of numbers between 1-32",
)
antialias_downscale = antialias_downscale and downscale_method in (
antialias_downscale = antialias_downscale and downscale_method in {
"bicubic",
"bilinear",
)
antialias_upscale = antialias_upscale and upscale_method in (
}
antialias_upscale = antialias_upscale and upscale_method in {
"bicubic",
"bilinear",
)
}
if start_fadeout_percent < start_percent:
start_fadeout_percent = start_percent
elif start_fadeout_percent > end_percent:
# No fadeout.
start_fadeout_percent = 1000.0
sigma_start = model.model.model_sampling.percent_to_sigma(start_percent)
sigma_end = model.model.model_sampling.percent_to_sigma(end_percent)
# Arbitrary number that should have good enough precision
pct_steps = 400
pct_incr = 1.0 / pct_steps
sig2pct = tuple(
model.model.model_sampling.percent_to_sigma(x / pct_steps)
for x in range(pct_steps, -1, -1)
)
ms = model.get_model_object("model_sampling")
sigma_start = ms.percent_to_sigma(start_percent)
sigma_end = ms.percent_to_sigma(end_percent)
def input_block_patch(h, transformer_options):
sigma = transformer_options["sigmas"][0].cpu().item()
block_num = transformer_options["block"][1]
sigma_tensor = transformer_options["sigmas"].max()
sigma = sigma_tensor.detach().cpu().item()
if (
sigma > sigma_start
or sigma < sigma_end
or transformer_options["block"][1] not in block_numbers
or block_num not in block_numbers
):
return h
# This is obviously terrible but I couldn't find a better way to get the percentage from the current sigma.
idx = bisect.bisect_right(sig2pct, sigma)
if idx >= len(sig2pct):
# Sigma out of range somehow?
return h
pct = pct_incr * (pct_steps - idx)
pct = 1.0 - (ms.timestep(sigma_tensor).detach().cpu() / 999)
if (
pct < start_fadeout_percent
or start_fadeout_percent > end_percent
@@ -144,9 +196,13 @@ class DeepShrinkBleh:
)
def output_block_patch(h, hsp, transformer_options):
if h.shape[-2:] == hsp.shape[-2:]:
return h, hsp
sigma = transformer_options["sigmas"][0].cpu().item()
if (
h.shape[-2:] == hsp.shape[-2:]
or sigma > sigma_start
or sigma < sigma_end
):
return h, hsp
return latent_utils.scale_samples(
h,
hsp.shape[-1],
+9 -6
View File
@@ -11,7 +11,7 @@ from einops import rearrange
class HyperTile:
def __init__(
def __init__( # noqa: PLR0917
self,
model,
seed,
@@ -144,6 +144,11 @@ class HyperTile:
class HyperTileBleh:
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "bleh/model_patches"
DESCRIPTION = "Model patch that speeds up generation at some cost of quality."
@classmethod
def INPUT_TYPES(cls):
return {
@@ -177,12 +182,10 @@ class HyperTileBleh:
},
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "bleh/model_patches"
@classmethod
def patch(
self,
cls,
*,
model,
seed,
tile_size,
+18 -12
View File
@@ -16,8 +16,10 @@ class DiscardPenultimateSigma:
FUNCTION = "go"
RETURN_TYPES = ("SIGMAS",)
CATEGORY = "sampling/custom_sampling/sigmas"
DESCRIPTION = "Discards the next to last sigma in the list."
def go(self, enabled, sigmas):
@classmethod
def go(cls, enabled, sigmas):
if not enabled or len(sigmas) < 2:
return (sigmas,)
return (torch.cat((sigmas[:-2], sigmas[-1:])),)
@@ -40,6 +42,11 @@ class SeededDisableNoise:
class BlehDisableNoise:
DESCRIPTION = "Allows setting a seed even when disabling noise. Used for SamplerCustomAdvanced or other nodes that take a NOISE input."
RETURN_TYPES = ("NOISE",)
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/noise"
@classmethod
def INPUT_TYPES(cls):
return {
@@ -51,13 +58,10 @@ class BlehDisableNoise:
},
}
def go(self, noise_seed):
@classmethod
def go(cls, noise_seed):
return (SeededDisableNoise(noise_seed),)
RETURN_TYPES = ("NOISE",)
FUNCTION = "go"
CATEGORY = "sampling/custom_sampling/noise"
class Wildcard(str):
__slots__ = ()
@@ -67,16 +71,18 @@ class Wildcard(str):
class BlehPlug:
DESCRIPTION = "This node can be used to plug up an input but act like the input was not actually connected. Can be used to prevent something like Use Everywhere nodes from supplying an input without having to set up blacklists or other configuration."
FUNCTION = "go"
OUTPUT_NODE = False
CATEGORY = "hacks"
WILDCARD = Wildcard("*")
RETURN_TYPES = (WILDCARD,)
@classmethod
def INPUT_TYPES(cls):
return {}
def go(self):
@classmethod
def go(cls):
return (None,)
RETURN_TYPES = (WILDCARD,)
FUNCTION = "go"
OUTPUT_NODE = False
CATEGORY = "hacks"
+66 -19
View File
@@ -80,12 +80,11 @@ class PatchTypeTransformerReplace(PatchTypeTransformer):
to["patches_replace"] = patches
patches[self.name] = val
@torch.no_grad()
def __call__(self, key, model_options, *args: list[Any]):
return self._call(key, self.get_patches(model_options), *args)
@torch.no_grad()
def _call(self, key, patches, *args: list[Any]):
@classmethod
def _call(cls, key, patches, *args: list[Any]):
p = patches.get(key)
if p:
return p(*args)
@@ -101,8 +100,8 @@ class PatchTypeModel(PatchTypeTransformer):
class PatchTypeModelWrapper(PatchTypeModel):
@torch.no_grad()
def _call(self, patches, apply_model, opts):
@classmethod
def _call(cls, patches, apply_model, opts):
if not patches:
return apply_model(opts["input"], opts["timestep"], **opts["c"])
return patches[0](apply_model, opts)
@@ -115,17 +114,25 @@ class PatchTypeSamplerPostCfgFunction(PatchTypeModel):
def set_patches(self, model_options, val):
model_options[self.name] = val
@torch.no_grad()
def _call(self, patches, opts):
result = opts["denoised"]
_call_result_key = "denoised"
@classmethod
def _call(cls, patches, opts):
curr_opts = opts.copy()
key = cls._call_result_key
for p in patches:
result = p(opts | {"denoised": result})
result = p(curr_opts)
curr_opts[key] = result
return result
class PatchTypeSamplerPreCfgFunction(PatchTypeSamplerPostCfgFunction):
_call_result_key = "conds_out"
class PatchTypeSamplerCfgFunction(PatchTypeModel):
@torch.no_grad()
def _call(self, patches, opts):
@classmethod
def _call(cls, patches, opts):
if not patches:
cond_pred, uncond_pred = opts["cond_denoised"], opts["uncond_denoised"]
return uncond_pred + (cond_pred - uncond_pred) * opts["cond_scale"]
@@ -150,6 +157,9 @@ PATCH_TYPES = {
"sampler_post_cfg_function": PatchTypeSamplerPostCfgFunction(
"sampler_post_cfg_function",
),
"sampler_pre_cfg_function": PatchTypeSamplerPostCfgFunction(
"sampler_pre_cfg_function",
),
}
@@ -168,7 +178,7 @@ class ModelConditionalState:
class ModelPatchConditional:
def __init__(
def __init__( # noqa: PLR0917
self,
model_default,
model_matched,
@@ -265,28 +275,65 @@ class ModelPatchConditionalNode:
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "bleh/model_patches"
DESCRIPTION = "Experimental model patch that lets you control when other model patches are active."
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_default": ("MODEL",),
"model_matched": ("MODEL",),
"model_default": (
"MODEL",
{
"tooltip": "Fallback model patches, used when start/end/interval do not match.",
},
),
"model_matched": (
"MODEL",
{"tooltip": "Model patches used when start/end/interval match."},
),
"start_percent": (
"FLOAT",
{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001},
{
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "Start time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"end_percent": (
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001},
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "End time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"interval": (
"INT",
{
"default": 1,
"min": -999,
"max": 999,
"tooltip": "Step interval to use model_matched. If positive 3 would mean activate every third step, if negative -3 would mean skip every third step.",
},
),
"base_on_default": (
"BOOLEAN",
{
"default": True,
"tooltip": "When true, the active set of patches will be applied to model_default, otherwise they will be applied to model_matched.",
},
),
"interval": ("INT", {"default": 1, "min": -999, "max": 999}),
"base_on_default": ("BOOLEAN", {"default": True}),
},
}
@classmethod
def patch(
self,
cls,
*,
model_default,
model_matched=None,
start_percent: float = 0.0,
+18 -12
View File
@@ -6,6 +6,7 @@ import importlib
import operator as pyop
from collections import OrderedDict
from enum import Enum, auto
from itertools import starmap
import torch
import yaml
@@ -122,7 +123,7 @@ class OpType(Enum):
# count, [ops]
REPEAT = auto()
#
# scale, type
APPLY_ENHANCEMENT = auto()
@@ -207,7 +208,7 @@ class Compare:
def __init__(self, typ: str, value):
self.typ = getattr(CompareType, typ.upper().strip())
if self.typ in (CompareType.OR, CompareType.AND, CompareType.NOT):
if self.typ in {CompareType.OR, CompareType.AND, CompareType.NOT}:
self.value = tuple(ConditionGroup(v) for v in value)
self.field = None
return
@@ -282,7 +283,7 @@ class ConditionGroup:
conds = tuple(conds.items())
if isinstance(conds[0], str):
conds = (conds,)
self.conds = tuple(Condition(ct, cv) for ct, cv in conds)
self.conds = tuple(starmap(Condition, conds))
def test(self, state: dict) -> bool:
return all(c.test(state) for c in self.conds)
@@ -319,7 +320,7 @@ class Operation:
if extra:
errstr = f"Unexpected argument keys for operation {typ}: {extra}"
raise ValueError(errstr)
self.args = tuple(args.get(k, v) for k, v in defaults.items())
self.args = tuple(starmap(args.get, defaults.items()))
else:
if len(args) > len(defaults):
raise ValueError("Too many arguments for operation")
@@ -458,7 +459,7 @@ class OpFlip(Operation):
def op(self, t, _state):
return torch.flip(
t,
dims=(2 if self.args[0] in ("v", "vertical") else 3,),
dims=(2 if self.args[0] in {"v", "vertical"} else 3,),
)
@@ -605,7 +606,8 @@ class OpNoise(Operation):
class OpDebug(Operation):
def op(self, t, state):
@classmethod
def op(cls, t, state):
stcopy = {
k: v
if not isinstance(v, torch.Tensor)
@@ -779,11 +781,12 @@ class BlehBlockOps:
},
}
@classmethod
def patch(
self,
cls,
model,
rules: str,
sigmas_opt=None,
sigmas_opt: None | torch.Tensor = None,
):
rules = rules.strip()
if len(rules) == 0:
@@ -919,7 +922,7 @@ class BlehBlockOps:
}
set_state_step(state, args["timestep"].max().item())
x = pre_model(state)
args = args | {"input": x}
args = args | {"input": x} # noqa: PLR6104
if orig_model_function_wrapper is not None:
result = orig_model_function_wrapper(apply_model, args)
else:
@@ -956,9 +959,11 @@ class BlehLatentScaleBy:
CATEGORY = "latent"
@classmethod
def upscale(
self,
samples,
cls,
*,
samples: dict,
method_horizontal: str,
method_vertical: str,
scale_width: float,
@@ -997,8 +1002,9 @@ class BlehLatentOps:
CATEGORY = "latent"
@classmethod
def go(
self,
cls,
samples,
rules: str,
):
+31 -6
View File
@@ -4,6 +4,10 @@ import comfy.model_management as mm
class BlehRefinerAfter:
DESCRIPTION = "Allows switching to another model at a certain point in sampling. Only works with models that are closely related as the sampling type and conditioning must match. Can be used to switch to a refiner model near the end of sampling."
RETURN_TYPES = ("MODEL",)
CATEGORY = "bleh/model_patches"
@classmethod
def INPUT_TYPES(cls):
return {
@@ -14,16 +18,34 @@ class BlehRefinerAfter:
"percent",
"sigma",
),
{
"tooltip": "Controls how start_time is interpreted. Timestep will be 999 at the start of sampling and 0 at the end - it is basically the inverse of the sampling percentage with a multiplier. Percent is the percent of sampling (not steps) and will be 0.0 at the start of sampling and 1.0 at the end. Sigma is an advanced option - if you don't know what it is, you don't need to use it.",
},
),
"start_time": (
"FLOAT",
{
"default": 199.0,
"min": 0.0,
"max": 999.0,
"tooltip": "Time the refiner_model will become active. The type of value you enter here will depend on what time_mode is set to.",
},
),
"model": (
"MODEL",
{
"tooltip": "Model to patch. This will also be the active model until the start_time condition is met.",
},
),
"refiner_model": (
"MODEL",
{
"tooltip": "Model to switch to after the start_time condition is met.",
},
),
"start_time": ("FLOAT", {"default": 199.0, "min": 0.0, "max": 999.0}),
"model": ("MODEL",),
"refiner_model": ("MODEL",),
},
}
RETURN_TYPES = ("MODEL",)
CATEGORY = "bleh/model_patches"
FUNCTION = "patch"
@staticmethod
@@ -60,6 +82,7 @@ class BlehRefinerAfter:
def check_time(sigma):
return sigma.item() <= start_time
case "percent":
if start_time > 1.0 or start_time < 0.0:
raise ValueError(
@@ -72,6 +95,7 @@ class BlehRefinerAfter:
def check_time(sigma):
return sigma.item() <= ms.percent_to_sigma(start_time)
case "timestep":
if start_time <= 0.0:
return (model,)
@@ -80,6 +104,7 @@ class BlehRefinerAfter:
def check_time(sigma):
return ms.timestep(sigma) <= start_time
case _:
raise ValueError("BlehRefinerAfter: invalid time mode")
+8 -8
View File
@@ -15,6 +15,10 @@ class SamplerChain(NamedTuple):
class BlehInsaneChainSampler:
RETURN_TYPES = ("SAMPLER", "BLEH_SAMPLER_CHAIN")
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "build"
@classmethod
def INPUT_TYPES(cls):
return {
@@ -27,11 +31,6 @@ class BlehInsaneChainSampler:
},
}
RETURN_TYPES = ("SAMPLER", "BLEH_SAMPLER_CHAIN")
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "build"
def build(
self,
sampler: object | None = None,
@@ -98,15 +97,16 @@ class BlehInsaneChainSampler:
class BlehForceSeedSampler:
DESCRIPTION = "ComfyUI has a bug where it will not set any seed if you have add_noise disabled in the sampler. This node is a workaround for that which ensures a seed alway gets set."
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"sampler": ("SAMPLER",)},
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "go"
def go(
+7 -4
View File
@@ -11,7 +11,9 @@ class Settings:
self.btp_enabled = False
else:
self.btp_enabled = True
self.btp_max_size = max(8, btp.get("max_size", 768))
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)
@@ -19,12 +21,13 @@ class Settings:
self.btp_preview_device = btp.get("preview_device")
self.btp_maxed_batch_step_mode = btp.get("maxed_batch_step_mode", False)
def get_cfg_path(self, filename):
@staticmethod
def get_cfg_path(filename) -> Path:
my_path = Path.resolve(Path(__file__).parent)
return my_path.parent / filename
def try_update_from_json(self, filename):
import json
import json # noqa: PLC0415
try:
with Path.open(self.get_cfg_path(filename)) as fp:
@@ -35,7 +38,7 @@ class Settings:
def try_update_from_yaml(self, filename):
try:
import yaml
import yaml # noqa: PLC0415
with Path.open(self.get_cfg_path(filename)) as fp:
self.update(yaml.safe_load(fp))
+1
View File
@@ -8,6 +8,7 @@ ignore = [
"ANN204",
"ANN206",
"C901",
"CPY001",
"D100",
"D101",
"D102",