diff --git a/README.md b/README.md index 2e7e45d..6a9a8b7 100644 --- a/README.md +++ b/README.md @@ -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,21 @@ 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`. You may also use `*` rather than a block +number to select all blocks in the category, for example `i*, o*` matches all input and all output blocks. + +The patch can be applied to the same model multiple times. + +Is it good, or even doing what I think? Who knows! Both positive and negative scales seem to have positive effect on the generation. Low negative scales applied to `cond` seem to make the generation bright and colorful. + +_Note_: Probably only works with SD 1.x and SDXL. Middle block patching will probably only work if you have [FreeU_Advanced](https://github.com/WASasquatch/FreeU_Advanced) installed. + ### BlehBlockOps Very experimental advanced node that allows defining model patches using YAML. This node is still under development and may be changed. diff --git a/__init__.py b/__init__.py index e418eac..072135a 100644 --- a/__init__.py +++ b/__init__.py @@ -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 diff --git a/blehconfig.json.example b/blehconfig.example.json similarity index 100% rename from blehconfig.json.example rename to blehconfig.example.json diff --git a/blehconfig.example.yaml b/blehconfig.example.yaml new file mode 100644 index 0000000..22887d9 --- /dev/null +++ b/blehconfig.example.yaml @@ -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 diff --git a/blehconfig.yaml.example b/blehconfig.yaml.example deleted file mode 100644 index c3f7786..0000000 --- a/blehconfig.yaml.example +++ /dev/null @@ -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 diff --git a/changelog.md b/changelog.md index 8231eb2..0d5ee87 100644 --- a/changelog.md +++ b/changelog.md @@ -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. diff --git a/py/betterTaesdPreview.py b/py/betterTaesdPreview.py index bc8d875..5c051b4 100644 --- a/py/betterTaesdPreview.py +++ b/py/betterTaesdPreview.py @@ -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 diff --git a/py/latent_utils.py b/py/latent_utils.py index aa852e7..5a71c35 100644 --- a/py/latent_utils.py +++ b/py/latent_utils.py @@ -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 diff --git a/py/nodes/blockCFG.py b/py/nodes/blockCFG.py new file mode 100644 index 0000000..e45f974 --- /dev/null +++ b/py/nodes/blockCFG.py @@ -0,0 +1,193 @@ +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). You may also use * instead of a block number to select all blocks in the category.", + }, + ), + "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 idx, item_ in enumerate(commasep_block_numbers.split(",")): + item = item_.strip().lower() + if not item: + continue + block_type = item[0] + if block_type not in "imo" or len(item) < 2: + errstr = f"Bad block definition at item {idx}" + raise ValueError(errstr) + if item[1] == "*": + block = tidx = -1 + else: + 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) + ok_time = sigma_end <= sigma <= sigma_start + if not ok_time: + return False + if block_def is None: + return -1 in block_list + return block_def in {-1, 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[secondary_idxs, ...] + ).mul_(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,) diff --git a/py/nodes/deepShrink.py b/py/nodes/deepShrink.py index be12694..e19fed8 100644 --- a/py/nodes/deepShrink.py +++ b/py/nodes/deepShrink.py @@ -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], diff --git a/py/nodes/hyperTile.py b/py/nodes/hyperTile.py index 52e1ecb..ebd1eba 100644 --- a/py/nodes/hyperTile.py +++ b/py/nodes/hyperTile.py @@ -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, diff --git a/py/nodes/misc.py b/py/nodes/misc.py index b069bca..c78e7a3 100644 --- a/py/nodes/misc.py +++ b/py/nodes/misc.py @@ -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" diff --git a/py/nodes/modelPatchConditional.py b/py/nodes/modelPatchConditional.py index 234dc01..c86df1d 100644 --- a/py/nodes/modelPatchConditional.py +++ b/py/nodes/modelPatchConditional.py @@ -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, diff --git a/py/nodes/ops.py b/py/nodes/ops.py index 84c3b28..0a609cd 100644 --- a/py/nodes/ops.py +++ b/py/nodes/ops.py @@ -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, ): diff --git a/py/nodes/refinerAfter.py b/py/nodes/refinerAfter.py index 07cec0c..8673ba4 100644 --- a/py/nodes/refinerAfter.py +++ b/py/nodes/refinerAfter.py @@ -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") diff --git a/py/nodes/samplers.py b/py/nodes/samplers.py index ec3520d..bccd664 100644 --- a/py/nodes/samplers.py +++ b/py/nodes/samplers.py @@ -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( diff --git a/py/settings.py b/py/settings.py index 528b08e..0a80e04 100644 --- a/py/settings.py +++ b/py/settings.py @@ -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)) diff --git a/ruff.toml b/ruff.toml index 4668c97..3fafb71 100644 --- a/ruff.toml +++ b/ruff.toml @@ -8,6 +8,7 @@ ignore = [ "ANN204", "ANN206", "C901", + "CPY001", "D100", "D101", "D102",