August '25 updates (#26)

* Previewer refactors
Support visual previews for ACE-Step
Allow setting dtype for previewer model
Allow scaling q/k/v input and output for Sage

* Start updating docs and example configuration files

* Add slice blend modes

* Rewrite blend mode system.
More blend modes.
Added BlehCFGInitSampler.
Basic Sparge Attention support.
Make samplers a bit more compatible with wrapping wrappers.

* Add probinject and probsubtract_b blend modes
Add BlehLatentAsImage and BlehImageAsLatent nodes

* Add BlehModelPatchFastTerminate node
Update docs/changelog
Add cosinesimilarity blend modes (basically SLERP)

* Remove some dead code
This commit is contained in:
blepping
2025-08-09 05:27:34 -06:00
committed by GitHub
parent 9606ca236c
commit 347e7eee6c
12 changed files with 1398 additions and 209 deletions
+18 -11
View File
@@ -6,17 +6,18 @@ For recent user-visible changes, please see the [ChangeLog](changelog.md).
## Features
1. Better TAESD previews (see below).
2. Allow setting seed, timestep range and step interval for HyperTile (look for the [`BlehHyperTile`](#blehhypertile) node).
3. Allow applying Kohya Deep Shrink to multiple blocks, also allow gradually fading out the downscale factor (look for the [`BlehDeepShrink`](#blehdeepshrink) node).
4. Allow discarding penultimate sigma (look for the `BlehDiscardPenultimateSigma` node). This can be useful if you find certain samplers are ruining your image by spewing a bunch of noise into it at the very end (usually only an issue with `dpm2 a` or SDE samplers).
5. Allow more conveniently switching between samplers during sampling (look for the [BlehInsaneChainSampler](#blehinsanechainsampler) node).
6. Apply arbitrary model patches at an interval and/or for a percentage of sampling (look for the [BlehModelPatchConditional](#blehmodelpatchconditional) node).
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).
11. [SageAttention](https://github.com/thu-ml/SageAttention/) support either globally or as a sampler wrapper. Look for the [BlehSageAttentionSampler](#blehsageattentionsampler) and `BlehGlobalSageAttention` nodes.
* Better TAESD previews (see below).
* Visual previews for some audio models (currently only ACE-Steps).
* Allow setting seed, timestep range and step interval for HyperTile (look for the [`BlehHyperTile`](#blehhypertile) node).
* Allow applying Kohya Deep Shrink to multiple blocks, also allow gradually fading out the downscale factor (look for the [`BlehDeepShrink`](#blehdeepshrink) node).
* Allow discarding penultimate sigma (look for the `BlehDiscardPenultimateSigma` node). This can be useful if you find certain samplers are ruining your image by spewing a bunch of noise into it at the very end (usually only an issue with `dpm2 a` or SDE samplers).
* Allow more conveniently switching between samplers during sampling (look for the [BlehInsaneChainSampler](#blehinsanechainsampler) node).
* Apply arbitrary model patches at an interval and/or for a percentage of sampling (look for the [BlehModelPatchConditional](#blehmodelpatchconditional) node).
* 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.
* Allows swapping to a refiner model at a predefined time (look for the [BlehRefinerAfter](#blehrefinerafter) node).
* Allow defining arbitrary model patches (look for the [BlehBlockOps](#blehblockops) node).
* Experimental blockwise CFG type effect (look for the [BlehBlockCFG](#blehblockcfg) node).
* [SageAttention](https://github.com/thu-ml/SageAttention/) support either globally or as a sampler wrapper. Look for the [BlehSageAttentionSampler](#blehsageattentionsampler) and `BlehGlobalSageAttention` nodes.
## Configuration
@@ -30,6 +31,9 @@ Restart ComfyUI to apply any new changes.
* Supports showing previews for more than the first latent in the batch.
* Supports throttling previews. Do you really need your expensive high quality preview to get updated 3 times a second?
The previewer can now show visual previews for ACE-Steps latents. If you want to disable that feature, you can add `aceaudio` to the
`blacklist_formats` list. For example if you are using a YAML configuration file you could do: `blacklist_formats: ["aceaudio"]`
**General settings defaults:**
|Key|Default|Description|
@@ -193,6 +197,7 @@ If you run into custom nodes that don't seem to be honoring SageAttention (you c
**Note:** Requires manually installing SageAttention into your Python environment. Should work with SageAttention 1.0 and 2.0.x (2.0.x currently requires CUDA 8+). Link: https://github.com/thu-ml/SageAttention
This also supports SpargeAttention (only in simple usage mode) if you have it installed, although I personally haven't seen it outperform Sage2. You can use this by setting `sageattn_function` to `sparge` or `sparge1` (for the SageAttention1-based version) in the YAML options. It is also possible to pass the `cdfthreshd` and `simthreshd1` parameters this way. See: https://github.com/thu-ml/SpargeAttn
### BlehGlobalSageAttention
@@ -339,3 +344,5 @@ Also may be an item from [Filters](#filters).
Many latent blending and scaling and filter functions based on implementation from https://github.com/WASasquatch/FreeU_Advanced - thanks!
TAE video model support based on code from https://github.com/madebyollin/taehv/.
AFS (analytical first step) formula from https://arxiv.org/abs/2210.05475
+2 -1
View File
@@ -15,10 +15,11 @@ from .py.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
def blep_init():
bi = getattr(nodes, "_blepping_integrations", {})
bi = sys.modules.get("_blepping_integrations", {})
if "bleh" in bi:
return
bi["bleh"] = sys.modules[__name__]
sys.modules["_blepping_integrations"] = bi
nodes._blepping_integrations = bi # noqa: SLF001
samplers.add_sampler_presets()
+6 -1
View File
@@ -10,6 +10,11 @@
"skip_upscale_layers": 0,
"compile_previewer": false,
"oom_fallback": "latent2rgb",
"oom_retry": true
"oom_retry": true,
"whitelist_formats": [],
"blacklist_formats": [],
"video_parallel": false,
"video_max_frames": -1,
"video_temporal_upscale_level": 0
}
}
+10
View File
@@ -21,6 +21,10 @@ betterTaesdPreviews:
# Minimum time between updating previews. The default will update the preview at most once per second.
throttle_secs: 1
# Can be set to use a different throttle time when using the fallback previewer
# (anything other than TAESD or TAEVID). If unset or null it will use throttle_secs.
throttle_secs_fallback: null
# 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
@@ -28,6 +32,12 @@ betterTaesdPreviews:
# alone unless you know you need to change it. Previewing on CPU will likely be quite slow.
preview_device: null
# Can be set to override the previewer dtype (which probably defaults to float32).
# You may set it to a specific dtype: float32, float16, bfloat16
# Setting it to "keep" or null just leaves the dtype alone (which is probably float32).
# Setting it to "vae" will use whatever dtype ComfyUI is set to use for VAE.
preview_dtype: 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
+14
View File
@@ -2,6 +2,20 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20250809
This set of changes involves refactoring parts of the previewer. Please create an issue if you experience problems.
* It's now possible to set the previewer dtype, see the `preview_dtype` setting. Note: Previews probably have been using float32 which is likely slower/more memory intensive than necessary. I'd recommend setting it to `vae`, `bfloat16` or `float16`.
* The previewer can show a visual representation of ACE-Steps latents (audio model). If you don't like it then you can add `aceaudio` to the `blacklist_formats` list in your configuration.
* You can set `throttle_secs_fallback` to use a different throttle setting for the fallback previewer (which includes stuff like the ACE-Steps previewer).
* Basic support for SpargeAttention, see README and https://github.com/thu-ml/SpargeAttn
* The blend mode system has been rewritten.
* Added more blend modes. Partial list as I don't remember exactly what I did: probinject, probsubtract_b, inject_difference, inject_copysign_a, inject_copysign_b, inject_avoidsign_a, inject_avoidsign_b, slice_X (where x may be dimension, flat, etc), loplerp_X, cosinesimilarity (and variants - basically reinvented SLERP here), hybrid_lerp_cosinesimilarity, subtract_b, subtract_b_scaleup_a
* Added a `BlehCFGInitSampler` sampler wrapper than can be used to skip steps like CFGZeroStar zero init, but without the cost of calling the model.
* Added `BlehImageAsLatent` and `BlehLatentAsImage` nodes that let you convert IMAGE to LATENT vice versa. Note this is just to allow running image operations on latents or latent operations on images, it doesn't really convert anything.
* Added a `BlehModelPatchFastTerminate` node that speeds up catching attempts to interrupt a generation. Mostly useful for video models where steps can take a very long time.
## 20250504
This is a fairly large set of changes. Please create an issue if you experience problems.
+238 -58
View File
@@ -10,7 +10,7 @@ import torch
from comfy import latent_formats
from comfy.cli_args import LatentPreviewMethod
from comfy.cli_args import args as comfy_args
from comfy.model_management import device_supports_non_blocking
from comfy.model_management import device_supports_non_blocking, vae_dtype
from comfy.taesd.taesd import TAESD
from PIL import Image
from tqdm import tqdm
@@ -29,6 +29,59 @@ _ORIG_GET_PREVIEWER = latent_preview.get_previewer
LAST_LATENT_FORMAT = None
# Referenced from https://github.com/learnables/learn2learn/blob/752200384c3ca8caeb8487b5dd1afd6568e8ec01/learn2learn/utils/__init__.py#L51
def clone_module(module, *, memo: dict | None = None) -> torch.nn.Module:
if not isinstance(module, torch.nn.Module):
raise TypeError("Expected torch.nn.Module")
if memo is None:
memo = {}
clone = module.__new__(type(module))
for k in ("__dict__", "_parameters", "_buffers", "_modules"):
if not hasattr(clone, k):
continue
setattr(clone, k, getattr(module, k).copy())
# We don't care about the has_grad case here.
for k in getattr(clone, "_parameters", {}):
v = module._parameters[k] # noqa: SLF001
if v is None:
continue
ptr = v.data_ptr
new_v = memo.get(ptr)
if new_v is None:
new_v = v.clone()
memo[ptr] = new_v
clone._parameters[k] = new_v # noqa: SLF001
for k in getattr(clone, "_modules", {}):
# print("RECURSE", k)
clone._modules[k] = clone_module(module._modules[k], memo=memo) # noqa: SLF001
if hasattr(clone, "flatten_parameters"):
clone = clone._apply(lambda x: x) # noqa: SLF001
return clone
# Simple heuristic.
def get_module_device_dtype(
module: torch.nn.Module,
) -> tuple[torch.device, torch.dtype] | tuple[None, None]:
p = next(module.parameters(), None)
if p is None:
raise RuntimeError("Couldn't get module device/dtype!")
return p.device, p.dtype
def normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)):
min_val, max_val = (
latent.amin(dim=dim, keepdim=True),
latent.amax(dim=dim, keepdim=True),
)
normalized = (latent - min_val).div_(max_val - min_val)
return (
normalized.mul_(target_max - target_min)
.add_(target_min)
.clamp_(target_min, target_max)
)
class VideoModelInfo(NamedTuple):
latent_format: latent_formats.LatentFormat
fps: int = 24
@@ -95,7 +148,6 @@ class FallbackPreviewerModel(torch.nn.Module):
upscale_mode: str = "bilinear",
):
super().__init__()
raw_factors = latent_format.latent_rgb_factors
raw_bias = latent_format.latent_rgb_factors_bias
factors = torch.tensor(raw_factors, device=device, dtype=dtype).transpose(0, 1)
@@ -121,7 +173,29 @@ class FallbackPreviewerModel(torch.nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.lin(x.movedim(1, -1)).movedim(-1, 1)
x = self.upsample(x).movedim(1, -1)
return x.add_(1.0).mul_(127.5).clamp_(0.0, 255.0)
return x.add_(1.0).clamp_(0.0, 2.0).mul_(127.5).round_()
class ACEStepsPreviewerModel(torch.nn.Module):
@torch.no_grad()
def __init__(
self,
*,
dtype: torch.dtype,
device: torch.device,
normalize_dims: tuple = (-1,),
):
super().__init__()
self.dtype = dtype
self.device = device
self.normalize_dims = normalize_dims
@torch.no_grad()
def forward(self, x: torch.Tensor) -> torch.Tensor:
batch, temporal = x.shape[0], x.shape[-1]
x = normalize_to_scale(x, 0.0, 1.0, dim=self.normalize_dims) * 255.0
x = x.reshape(batch, -1, temporal)
return x[..., None].expand(*x.shape, 3)
class BetterPreviewer(_ORIG_PREVIEWER):
@@ -133,6 +207,12 @@ class BetterPreviewer(_ORIG_PREVIEWER):
vid_info: VideoModelInfo | None = None,
):
self.latent_format = latent_format
self.latent_format_name = (
"unknown"
if latent_format is None
else latent_format.__class__.__name__.lower()
)
self.spatial_compression = 8
self.vid_info = vid_info
self.fallback_previewer_model = None
self.device = (
@@ -140,14 +220,27 @@ class BetterPreviewer(_ORIG_PREVIEWER):
if SETTINGS.btp_preview_device is None
else torch.device(SETTINGS.btp_preview_device)
)
dtype = (
SETTINGS.btp_preview_dtype.lower()
if SETTINGS.btp_preview_dtype is not None
else None
)
self.dtype: str | torch.dtype | None = None
if dtype in {"vae", "keep"}:
self.dtype = dtype
elif dtype in {"float32", "float16", "bfloat16"}:
self.dtype = getattr(torch, dtype)
self.orig_previewer_model = (
None
if taesd is None
else clone_module(taesd).to(device="cpu", dtype=torch.float32)
)
if taesd is not None:
if hasattr(taesd, "taesd_encoder"):
del taesd.taesd_encoder
if hasattr(taesd, "encoder"):
del taesd.encoder
if self.device and self.device != next(taesd.parameters()).device:
taesd = taesd.to(self.device)
self.taesd = taesd
self.previewer_model = taesd
self.stamp = None
self.cached = None
self.blank = Image.new("RGB", size=(1, 1))
@@ -155,30 +248,75 @@ class BetterPreviewer(_ORIG_PREVIEWER):
self.oom_retry = SETTINGS.btp_oom_retry
self.oom_count = 0
self.skip_upscale_layers = SETTINGS.btp_skip_upscale_layers
self.skip_upscale_layers_state: tuple[int, int] | tuple[None, None] | None = (
None
)
self.preview_max_width = SETTINGS.btp_max_width
self.preview_max_height = SETTINGS.btp_max_height
self.throttle_secs = SETTINGS.btp_throttle_secs
self.throttle_secs_fallback = SETTINGS.btp_throttle_secs_fallback
self.max_batch_preview = SETTINGS.btp_max_batch
self.maxed_batch_step_mode = SETTINGS.btp_maxed_batch_step_mode
self.max_batch_cols = SETTINGS.btp_max_batch_cols
self.compile_previewer = SETTINGS.btp_compile_previewer
self.maybe_pop_upscale_layers()
if self.compile_previewer:
compile_kwargs = (
{}
if not isinstance(self.compile_previewer, dict)
else self.compile_previewer
def maybe_refresh_previewer(
self,
*,
device=None,
dtype=None,
width=None,
height=None,
) -> None:
if self.orig_previewer_model is None:
return
pdevice, pdtype = (
get_module_device_dtype(self.previewer_model)
if self.previewer_model is not None
else (None, None)
)
need_refresh = (
self.previewer_model is None
or (dtype is not None and pdtype != dtype)
or (device is not None and pdevice != device)
)
is_taesd = isinstance(self.orig_previewer_model, TAESD)
if is_taesd and not need_refresh:
need_refresh = (
self.skip_upscale_layers < 0
and self.skip_upscale_layers_state != (width, height)
)
self.taesd = torch.compile(self.taesd, **compile_kwargs)
if not need_refresh:
return
tqdm.write("Refreshing previewer")
self.previewer_model = clone_module(self.orig_previewer_model).to(
device=device,
dtype=dtype,
)
if is_taesd:
self.skip_upscale_layers_state = None
self.maybe_pop_upscale_layers(width=width, height=height)
if not self.compile_previewer:
return
tqdm.write("Compiling previewer")
compile_kwargs = (
{}
if not isinstance(self.compile_previewer, dict)
else self.compile_previewer
)
self.previewer_model = torch.compile(self.previewer_model, **compile_kwargs)
# Popping upscale layers trick from https://github.com/madebyollin/
def maybe_pop_upscale_layers(self, *, width=None, height=None) -> None:
if self.skip_upscale_layers_state:
return
self.skip_upscale_layers_state = (width, height)
skip = self.skip_upscale_layers
if skip == 0 or not isinstance(self.taesd, TAESD):
if skip == 0 or not isinstance(self.previewer_model, TAESD):
return
upscale_layers = tuple(
idx
for idx, layer in enumerate(self.taesd.taesd_decoder)
for idx, layer in enumerate(self.previewer_model.taesd_decoder)
if isinstance(layer, torch.nn.Upsample)
)
num_upscale_layers = len(upscale_layers)
@@ -203,8 +341,7 @@ class BetterPreviewer(_ORIG_PREVIEWER):
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
self.previewer_model.taesd_decoder.pop(upscale_layers[-idx])
def decode_latent_to_preview_image(
self,
@@ -223,9 +360,14 @@ class BetterPreviewer(_ORIG_PREVIEWER):
def check_use_cached(self) -> bool:
now = time()
throttle = (
self.throttle_secs
if self.previewer_model is not None
else self.throttle_secs_fallback
)
if (
self.cached is not None and self.stamp is not None
) and now - self.stamp < self.throttle_secs:
) and now - self.stamp < throttle:
return True
self.stamp = now
return False
@@ -253,15 +395,9 @@ class BetterPreviewer(_ORIG_PREVIEWER):
is_video = x0.ndim == 5
if frames_to_batch and is_video:
x0 = x0.transpose(2, 1).reshape(-1, x0.shape[1], *x0.shape[-2:])
batch = x0.shape[0]
x0 = x0[self.calculate_indexes(batch, is_video=is_video), :]
x0 = x0[self.calculate_indexes(x0.shape[0], is_video=is_video), :]
batch = x0.shape[0]
height, width = x0.shape[-2:]
if self.device and x0.device != self.device:
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,
@@ -269,14 +405,41 @@ class BetterPreviewer(_ORIG_PREVIEWER):
)
return x0, cols, rows
def prepare_previewer(
self,
x0: torch.Tensor,
*,
img_width: int | None = None,
img_height: int | None = None,
) -> torch.Tensor:
if self.dtype == "vae":
dtype = vae_dtype(x0)
elif self.dtype == "keep":
dtype = x0.dtype
else:
dtype = self.dtype
self.maybe_refresh_previewer(
dtype=dtype,
device=self.device or x0.device,
width=img_width,
height=img_height,
)
pdevice, pdtype = get_module_device_dtype(self.previewer_model)
# tqdm.write(
# f"\nPREVIEW: pdevice={pdevice}, pdtype={pdtype}, device={x0.device}, dtype={x0.dtype}",
# )
if x0.device == pdevice and x0.dtype == pdtype:
return x0
return x0.to(
device=pdevice,
dtype=pdtype,
non_blocking=device_supports_non_blocking(x0.device),
)
def _decode_latent_taevid(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]:
height, width = x0.shape[-2:]
if self.device and x0.device != self.device:
x0 = x0.to(
device=self.device,
non_blocking=device_supports_non_blocking(x0.device),
)
decoded = self.taesd.decode(
x0 = self.prepare_previewer(x0)
decoded = self.previewer_model.decode(
x0.transpose(1, 2),
parallel=SETTINGS.btp_video_parallel,
).movedim(2, -1)
@@ -295,7 +458,7 @@ class BetterPreviewer(_ORIG_PREVIEWER):
height,
)
return (
decoded.mul_(255.0).round_().clamp_(min=0, max=255.0).detach(),
decoded.clamp_(0.0, 1.0).mul_(255.0).round_().detach(),
cols,
rows,
)
@@ -303,21 +466,22 @@ class BetterPreviewer(_ORIG_PREVIEWER):
def _decode_latent_taesd(self, x0: torch.Tensor) -> tuple[torch.Tensor, int, int]:
x0, cols, rows = self.prepare_decode_latent(
x0,
frames_to_batch=not isinstance(self.taesd, TAEVid),
frames_to_batch=not isinstance(self.previewer_model, TAEVid),
)
height, width = x0.shape[-2:]
if self.skip_upscale_layers < 0:
self.maybe_pop_upscale_layers(
width=width * 8 * cols,
height=height * 8 * rows,
)
img_height, img_width = (
height * self.spatial_compression * rows,
width * self.spatial_compression * cols,
)
x0 = self.prepare_previewer(x0, img_width=img_width, img_height=img_height)
return (
(
self.taesd.decode(x0)
self.previewer_model.decode(x0)
.movedim(1, -1)
.add_(1.0)
.clamp_(0.0, 2.0)
.mul_(127.5)
.clamp_(min=0, max=255.0)
.round_()
.detach()
),
cols,
@@ -381,10 +545,20 @@ class BetterPreviewer(_ORIG_PREVIEWER):
@torch.no_grad()
def init_fallback_previewer(self, device: torch.device, dtype: torch.dtype) -> bool:
if self.fallback_previewer_model is not None:
return True
if self.latent_format is None:
return False
if (
self.fallback_previewer_model is not None
and self.fallback_previewer_model.dtype == dtype
and self.fallback_previewer_model.device == device
):
return True
if self.latent_format_name == "aceaudio":
self.fallback_previewer_model = ACEStepsPreviewerModel(
device=device,
dtype=dtype,
)
return True
self.fallback_previewer_model = FallbackPreviewerModel(
self.latent_format,
device=device,
@@ -421,7 +595,7 @@ class BetterPreviewer(_ORIG_PREVIEWER):
return self.cached
if x0.shape[0] == 0:
return self.blank # Shouldn't actually be possible.
if (self.oom_count and not self.oom_retry) or self.taesd is None:
if (self.oom_count and not self.oom_retry) or self.previewer_model is None:
return self.fallback_previewer(x0, quiet=True)
is_video = x0.ndim == 5
used_fallback = False
@@ -449,14 +623,21 @@ def bleh_get_previewer(
*args: list,
**kwargs: dict,
) -> object | None:
def orig_get_previewer():
return _ORIG_GET_PREVIEWER(device, latent_format, *args, **kwargs)
preview_method = comfy_args.preview_method
if preview_method == LatentPreviewMethod.NoPreviews:
return orig_get_previewer()
format_name = latent_format.__class__.__name__.lower()
if (
not SETTINGS.btp_enabled
or format_name in SETTINGS.btp_blacklist
or (SETTINGS.btp_whitelist and format_name not in SETTINGS.btp_whitelist)
):
return _ORIG_GET_PREVIEWER(device, latent_format, *args, **kwargs)
return orig_get_previewer()
tae_model = None
if preview_method in {LatentPreviewMethod.TAESD, LatentPreviewMethod.Auto}:
vid_info = VIDEO_FORMATS.get(format_name)
@@ -473,9 +654,9 @@ def bleh_get_previewer(
TAEVid(
checkpoint_path=tae_model_path,
latent_channels=latent_format.latent_channels,
device=device,
device=torch.device("cpu"),
decoder_time_upscale=decoder_time_upscale,
).to(device)
)
if tae_model_path is not None
else None
)
@@ -489,23 +670,22 @@ def bleh_get_previewer(
None,
taesd_path,
latent_channels=latent_format.latent_channels,
).to(device)
)
if taesd_path is not None
else None
)
return BetterPreviewer(
taesd=tae_model,
latent_format=latent_format,
vid_info=vid_info,
)
if (
preview_method == LatentPreviewMethod.NoPreviews
or latent_format.latent_rgb_factors is None
if tae_model is not None:
return BetterPreviewer(
taesd=tae_model,
latent_format=latent_format,
vid_info=vid_info,
)
if format_name == "aceaudio" or (
preview_method == LatentPreviewMethod.Latent2RGB
and latent_format.latent_rgb_factors is not None
):
return None
if preview_method == LatentPreviewMethod.Latent2RGB:
return BetterPreviewer(latent_format=latent_format)
return _ORIG_GET_PREVIEWER(device, latent_format, *args, **kwargs)
return orig_get_previewer()
def ensure_previewer():
+745 -87
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import math
import os
from functools import partial
from typing import ClassVar
import kornia.filters as kf
import numpy as np
@@ -49,6 +50,27 @@ if USE_ORIG_NORMALIZE:
normalize = normalize_orig
def normalize_to_scale(
latent: torch.Tensor,
target_min: float,
target_max: float,
*,
dim=(-3, -2, -1),
eps: float = 1e-07,
) -> torch.Tensor:
min_val, max_val = (
latent.amin(dim=dim, keepdim=True),
latent.amax(dim=dim, keepdim=True),
)
normalized = latent - min_val
normalized /= (max_val - min_val).add_(eps)
return (
normalized.mul_(target_max - target_min)
.add_(target_min)
.clamp_(target_min, target_max)
)
def hslerp(a, b, t):
if a.shape != b.shape:
raise ValueError("Input tensors a and b must have the same shape.")
@@ -145,6 +167,7 @@ def altslerp( # noqa: PLR0914
v0: FloatTensor,
v1: FloatTensor,
t: float | FloatTensor,
*,
dot_threshold=0.9995,
dim=-1,
):
@@ -196,27 +219,6 @@ def altslerp( # noqa: PLR0914
return out
# def prob_blend_(a, b, t, *, cpu=False):
# if not isinstance(t, torch.Tensor):
# t = torch.tensor((t,), dtype=a.dtype, device=a.device)
# tmin, tmax = t.aminmax()
# tmin, tmax = min(tmin, 0.0), max(tmax, 1.0)
# t = t - tmin
# tdiv = tmax - tmin
# if tdiv != 0:
# t /= tdiv
# t = t.broadcast_to(a.shape)
# probs = torch.rand(
# *a.shape,
# dtype=a.dtype,
# layout=a.layout,
# device="cpu" if cpu else a.device,
# )
# if probs.device != a.device:
# probs = probs.to(a.device)
# return torch.where(probs > t, a, b)
def stochasistic_blend(
a,
b,
@@ -249,22 +251,6 @@ def stochasistic_blend(
return blend(a, b, tadj)
def prob_blend(a, b, t, *, cpu=False):
t_device = torch.device("cpu") if cpu else a.device
if not isinstance(t, torch.Tensor):
t = torch.tensor((t,), dtype=a.dtype, device=t_device)
elif t.device != t_device:
t = t.detach().clone().to(t_device)
tmin, tmax = t.aminmax()
tmin, tmax = min(tmin, 0.0), max(tmax, 1.0)
t = t - tmin # noqa: PLR6104
tdiv = tmax - tmin
if tdiv != 0:
t /= tdiv
t = t.clamp_(0, 1).broadcast_to(a.shape)
return torch.where(torch.bernoulli(t).to(device=a.device, dtype=torch.bool), b, a)
def gaussian_smoothing(
t: torch.Tensor,
kernel_size,
@@ -308,31 +294,63 @@ def gaussian_smoothing(
return result
def prob_blend_smoothed(
a,
b,
t,
*,
cpu: bool = False,
blend=torch.lerp,
kernel_size: int | tuple | list = 3,
sigma: float | tuple | list = 1.0,
):
t_device = torch.device("cpu") if cpu else a.device
if not isinstance(t, torch.Tensor):
t = torch.tensor((t,), dtype=a.dtype, device=t_device)
elif t.device != t_device:
t = t.detach().clone().to(t_device)
tmin, tmax = t.aminmax()
tmin, tmax = min(tmin, 0.0), max(tmax, 1.0)
t = t - tmin # noqa: PLR6104
tdiv = tmax - tmin
if tdiv != 0:
t /= tdiv
t = t.clamp_(0, 1).broadcast_to(a.shape)
t = torch.bernoulli(t).to(device=a.device, dtype=a.dtype)
t = gaussian_smoothing(t, kernel_size, sigma)
return blend(a, b, t)
class ProbBlend:
@staticmethod
def output(a: torch.Tensor, b: torch.Tensor, b_t: torch.Tensor) -> torch.Tensor:
return torch.where(b_t.to(device=a.device, dtype=torch.bool), b, a)
def __call__(
self,
a,
b,
t,
*,
cpu=False,
collapse_dims=(),
**kwargs: dict,
):
t_device = torch.device("cpu") if cpu else a.device
if not isinstance(t, torch.Tensor):
t = torch.tensor((t,), dtype=a.dtype, device=t_device)
elif t.device != t_device:
t = t.detach().clone().to(t_device)
tmin, tmax = t.aminmax()
tmin, tmax = min(tmin, 0.0), max(tmax, 1.0)
t = t - tmin # noqa: PLR6104
tdiv = tmax - tmin
if tdiv != 0:
t /= tdiv
if collapse_dims:
dims = a.ndim
prob_shape = list(a.shape)
for didx in collapse_dims:
if didx >= dims:
continue
prob_shape[didx] = 1
else:
prob_shape = a.shape
t = torch.bernoulli(t.clamp_(0, 1).broadcast_to(prob_shape)).to(a)
return self.output(a, b, t, **kwargs)
class ProbBlendSmoothed(ProbBlend):
@staticmethod
def output(
a: torch.Tensor,
b: torch.Tensor,
b_t: torch.Tensor,
*,
output_blend=torch.lerp,
kernel_size: int | tuple | list = 3,
sigma: float | tuple | list = 1.0,
) -> torch.Tensor:
t = b_t.to(device=a.device, dtype=a.dtype)
t = gaussian_smoothing(t, kernel_size, sigma)
return output_blend(a, b, t)
prob_blend = ProbBlend()
prob_blend_smoothed = ProbBlendSmoothed()
# Originally referenced from https://github.com/54rt1n/ComfyUI-DareMerge
@@ -395,10 +413,397 @@ def gradient_blend(
return result
def slice_blend(
a: torch.Tensor,
b: torch.Tensor,
t: float | torch.Tensor,
*,
flatten=True,
dim=1,
flip_a=False,
flip_b=False,
flip_out=False,
) -> torch.Tensor:
if isinstance(t, torch.Tensor):
t = t.mean().clamp(0, 1)
else:
t = a.new_full((1,), t).clamp(0, 1)
if t == 0:
return a
if t == 1:
return b
orig_shape = a.shape
if a.ndim > 2 and flatten:
a = a.flatten(start_dim=dim)
b = b.flatten(start_dim=dim)
elsb = int(a.shape[dim] * t)
elsa = a.shape[dim] - elsb
astart, aend = (None, elsa) if not flip_a else (a.shape[dim] - elsa, None)
bstart, bend = (None, elsb) if flip_b else (a.shape[dim] - elsb, None)
aslice = tuple(
slice(None) if i != dim else slice(astart, aend) for i in range(a.ndim)
)
bslice = tuple(
slice(None) if i != dim else slice(bstart, bend) for i in range(a.ndim)
)
achunk, bchunk = a[aslice], b[bslice]
result = torch.cat((bchunk, achunk) if flip_out else (achunk, bchunk), dim=dim)
return result.reshape(orig_shape)
def slice_blend_smooth( # noqa: PLR0914
a: torch.Tensor,
b: torch.Tensor,
t: float | torch.Tensor,
*,
flatten: bool = True,
dim: int = 1,
fade_percent_l: float = 0.1,
fade_percent_r: float = 0.1,
always_fade: bool = False,
b_start_percent: float = 1.0,
b_blend_max: float = 1.0,
invert: bool = False, # Doesn't work propertly at the moment.
blend_function=torch.lerp,
) -> torch.Tensor:
if isinstance(t, torch.Tensor):
t = t.mean().clamp(0, 1)
else:
t = a.new_full((1,), t).clamp_(0, 1)
if invert:
t = 1 - t
b, a = a, b
if t == 0:
return a
b_start_percent = max(0.0, min(1.0, b_start_percent))
fade_percent_l = (
max(0.0, min(1.0, fade_percent_l))
if b_start_percent > 0 and not always_fade
else 0.0
)
fade_percent_r = (
max(0.0, min(1.0, fade_percent_r))
if b_start_percent < 1 and not always_fade
else 0.0
)
fade_mul = 1.0 / max(1.0, fade_percent_l + fade_percent_r)
orig_shape = a.shape
if flatten and dim < a.ndim - 1:
a = a.flatten(start_dim=dim)
b = b.flatten(start_dim=dim)
dim_els = a.shape[dim]
els_b = int(dim_els * t)
if invert:
els_b += int((dim_els - els_b) * (fade_percent_l + fade_percent_r) * fade_mul)
els_a = dim_els - els_b
b_start = int(els_a * b_start_percent)
b_end = b_start + els_b
elslfade, elsrfade = (
int(els_b * fade_percent_l * fade_mul),
int(els_b * fade_percent_r * fade_mul),
)
blend_mask = a.new_zeros(dim_els)
blend_mask[b_start:b_end] = b_blend_max
if elslfade > 0:
blend_mask[b_start : b_start + elslfade] = torch.linspace(
0.0,
b_blend_max,
steps=elslfade + 2,
device=blend_mask.device,
dtype=blend_mask.dtype,
)[1:-1]
if elsrfade > 0:
rfade_start = b_end - elsrfade
blend_mask[rfade_start : rfade_start + elsrfade] = torch.linspace(
b_blend_max,
0.0,
steps=elsrfade + 2,
device=blend_mask.device,
dtype=blend_mask.dtype,
)[1:-1]
blend_mask = blend_mask.view(
tuple(dim_els if d == dim else 1 for d in range(a.ndim)),
)
return blend_function(a, b, blend_mask).reshape(orig_shape)
def lop_lerp(
a: torch.Tensor,
b: torch.Tensor,
t: torch.tensor | float,
*,
a_ratio=1.0,
b_ratio=1.0,
):
if not isinstance(t, torch.Tensor):
t = a.new_full((1,), t)
return (a_ratio - t.clamp(max=a_ratio)).mul(a).add_(b * (t * b_ratio))
# # Thanks, ChatGPT though you did get the ratio reversed.
def cosine_similarity_blend_chatgpt_orig(
b: torch.Tensor,
a: torch.Tensor,
ratio: float | torch.Tensor,
*,
dim: int = -1,
eps: float = 1e-6,
) -> torch.Tensor:
a_n = a / (a.norm(dim=dim, keepdim=True).clamp_min(eps))
b_n = b / (b.norm(dim=dim, keepdim=True).clamp_min(eps))
c = torch.sum(a_n * b_n, dim=dim, keepdim=True)
s = 2 * ratio - 1
if not torch.is_tensor(s):
s = a.new_tensor(s)
a_ = 1 - c
alpha = a_ * (a_ - 2 * s**2)
beta = 2 * a_ * (c + s**2)
gamma = c**2 - s**2
disc = beta**2 - 4 * alpha * gamma
disc = disc.clamp_min(0.0)
sqrt_disc = torch.sqrt(disc)
lam1 = (-beta + sqrt_disc) / (2 * alpha).clamp_min(eps)
lam2 = (-beta - sqrt_disc) / (2 * alpha).clamp_min(eps)
lam = torch.where((lam1 >= 0) & (lam1 <= 1), lam1, lam2)
lam = torch.where((lam >= 0) & (lam <= 1), lam, ratio)
return torch.lerp(b, a, lam.expand_as(a))
def cosine_similarity_blend_chatgpt( # noqa: PLR0914
a: torch.Tensor,
b: torch.Tensor,
ratio: float,
*,
dim: int = -1,
eps: float = 1e-8,
small_angle: float = 1e-4,
opp_eps: float = 1e-6,
) -> torch.Tensor:
# --- normalize directions ---
mag_a = a.norm(dim=dim, keepdim=True).clamp_min(eps)
mag_b = b.norm(dim=dim, keepdim=True).clamp_min(eps)
a_n = a / mag_a
b_n = b / mag_b
# cosine & angle between a and b
cos_ab = (a_n * b_n).sum(dim=dim, keepdim=True).clamp(-1.0, 1.0)
theta = torch.acos(cos_ab)
# map blend ratio -> fraction along the arc
# we want angle from a -> out = theta * t
t = ratio if torch.is_tensor(ratio) else a.new_tensor(ratio)
# handle exact-opposite case: fallback to lerp then renorm
opp_mask = torch.abs(cos_ab + 1) < opp_eps
if opp_mask.any():
# simple normalized lerp + renormalize
lerp_dir = (1 - t) * a_n + t * b_n
lerp_dir /= lerp_dir.norm(dim=dim, keepdim=True).clamp_min(eps)
# magnitude later will apply
a_n = torch.where(opp_mask, lerp_dir, a_n)
b_n = torch.where(opp_mask, b_n, b_n) # no-op but keeps shapes aligned
theta = torch.where(
opp_mask,
torch.acos((a_n * b_n).sum(dim=dim, keepdim=True)),
theta,
)
# for very small angles, do lerp+renormalize
lerp_mask = theta < small_angle
if lerp_mask.any():
lerp_dir = (1 - t) * a_n + t * b_n
lerp_dir /= lerp_dir.norm(dim=dim, keepdim=True).clamp_min(eps)
# override only where theta is small
a_n = torch.where(lerp_mask, lerp_dir, a_n)
b_n = torch.where(lerp_mask, b_n, b_n)
cos_ab = (a_n * b_n).sum(dim=dim, keepdim=True).clamp(-1, 1)
theta = torch.acos(cos_ab)
# now true SLERP coefficients
sin_theta = torch.sin(theta).clamp_min(eps)
coef_a = torch.sin((1 - t) * theta) / sin_theta
coef_b = torch.sin(t * theta) / sin_theta
dir_out = coef_a * a_n + coef_b * b_n
# --- geometric magnitude interpolation ---
log_a = torch.log(mag_a)
log_b = torch.log(mag_b)
log_out = (1 - t) * log_a + t * log_b
mag_out = torch.exp(log_out)
return dir_out * mag_out
def cosine_similarity_blend_deepseek( # noqa: PLR0914
a: torch.Tensor,
b: torch.Tensor,
ratio: float,
*,
dim: int = -1,
eps=1e-08,
threshold=1e-06,
) -> torch.Tensor:
if not torch.is_tensor(ratio):
ratio = a.new_tensor(ratio)
# Compute magnitudes of a and b along the specified dimension
mag_a = torch.norm(a, p=2, dim=dim, keepdim=True).add_(eps)
mag_b = torch.norm(b, p=2, dim=dim, keepdim=True).add_(eps)
# Avoid division by zero during normalization
a_norm = a / mag_a
b_norm = b / mag_b
# Compute cosine similarity (dot product of normalized vectors)
d = (a_norm * b_norm).sum(dim=dim, keepdim=True).clamp_(-1.0, 1.0)
# Compute angle between a_norm and b_norm
theta = torch.acos(d)
# Calculate desired cosine similarity with b (s_b) from blend ratio
s_b = (2.0 * ratio - 1.0).clamp_(-1.0, 1.0)
# Compute angle from result to b based on s_b
angle_from_b = torch.acos(s_b)
# Calculate interpolation parameter t_val
t_val = (1.0 - angle_from_b / theta).clamp_(0.0, 1.0)
# Precompute sin_theta for slerp
sin_theta = torch.sin(theta)
# Linear interpolation fallback for small sin_theta
linear_part_norm = torch.lerp(a_norm, b_norm, t_val)
# linear_part_norm = (1.0 - t_val) * a_norm + t_val * b_norm
# Slerp computation
sin_t_theta = torch.sin(t_val * theta)
sin_comp_theta = torch.sin((1.0 - t_val) * theta)
slerp_denom = sin_theta + eps # Avoid division by zero
slerp_part_norm = (sin_comp_theta / slerp_denom) * a_norm + (
sin_t_theta / slerp_denom
) * b_norm
# Choose slerp unless sin_theta is too small (use linear then)
v_norm = torch.where(sin_theta < threshold, linear_part_norm, slerp_part_norm)
# Linearly interpolate magnitude
mag = torch.lerp(mag_a, mag_b, ratio)
# Scale normalized vector by interpolated magnitude
return v_norm * mag
DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND = "chatgpt"
COSINE_SIMILARITY_BLEND_BACKENDS = {
"altslerp": altslerp,
"deepseek": cosine_similarity_blend_deepseek,
"chatgpt": cosine_similarity_blend_chatgpt,
"chatgpt_orig": cosine_similarity_blend_chatgpt_orig,
}
def cosine_similarity_blend(
a: torch.Tensor,
b: torch.Tensor,
t: float | torch.Tensor,
*args: list,
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
**kwargs: dict,
) -> torch.Tensor:
fun = COSINE_SIMILARITY_BLEND_BACKENDS.get(backend)
if fun is None:
errstr = f"Bad cosine similarity blend backend {backend}, must be one of {tuple(COSINE_SIMILARITY_BLEND_BACKENDS)}"
raise ValueError(errstr)
return fun(a, b, t, *args, **kwargs)
def cosine_similarity_blend_avg(
a: torch.Tensor,
b: torch.Tensor,
t: float | torch.Tensor,
*,
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
dims=(-1, -2),
) -> torch.Tensor:
blend_fun = partial(cosine_similarity_blend, backend=backend)
multiplier = 1.0 / len(dims)
result = None
for dim in dims:
curr_result = blend_fun(a, b, t, dim=dim).mul_(multiplier)
result = curr_result if result is None else result.add_(curr_result)
return result
def cosine_similarity_blend_flat(
a: torch.Tensor,
b: torch.Tensor,
t: float | torch.Tensor,
*,
start_dim=0,
end_dim=1,
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
) -> torch.Tensor:
if a.shape != b.shape:
raise ValueError("Tensor shape mismatch, a and b must be the same shape")
if start_dim < 0:
start_dim = a.ndim + start_dim
if end_dim < 0:
end_dim = a.ndim + end_dim
if start_dim < 0 or end_dim < 0 or start_dim >= a.ndim or end_dim >= a.ndim:
raise ValueError("Bad start/end_dim parameters")
orig_shape = a.shape
a = a.flatten(start_dim=start_dim, end_dim=end_dim)
b = b.flatten(start_dim=start_dim, end_dim=end_dim)
if isinstance(t, torch.Tensor) and t.ndim == len(orig_shape):
t = t.flatten(start_dim=start_dim, end_dim=end_dim)
return cosine_similarity_blend(a, b, t, dim=start_dim, backend=backend).reshape(
orig_shape,
)
def blend_blend(
a: torch.Tensor,
b: torch.Tensor,
t: float | torch.Tensor,
*,
blend_mode_a: str = "lerp",
blend_mode_b="cosinesimilarity_flat_spatdims",
blend_blend: float | torch.Tensor = 0.5,
blend_mode_blend: str = "lerp",
blend_a_kwargs: dict | None = None,
blend_b_kwargs: dict | None = None,
blend_blend_kwargs: dict | None = None,
) -> torch.Tensor:
fun_a = BLENDING_MODES[blend_mode_a]
fun_b = BLENDING_MODES[blend_mode_b]
fun_blend = BLENDING_MODES[blend_mode_blend]
if not torch.is_tensor(blend_blend):
blend_blend = a.new_tensor(blend_blend)
return fun_blend(
fun_a(a, b, t, **({} if blend_a_kwargs is None else blend_a_kwargs)),
fun_b(a, b, t, **({} if blend_b_kwargs is None else blend_b_kwargs)),
blend_blend,
**({} if blend_blend_kwargs is None else blend_blend_kwargs),
)
class BlendMode:
__slots__ = (
"allow_scale",
"f",
"f_kwargs",
"f_raw",
"force_rescale",
"norm",
"norm_dims",
@@ -418,8 +823,15 @@ class BlendMode:
allow_scale=True,
rescale_dims=(-3, -2, -1),
force_rescale=False,
**kwargs: dict,
):
self.f = f
self.f_raw = f
self.f = f if not kwargs else partial(f, **kwargs)
self.f_kwargs = kwargs
if norm is True:
norm = normalize
elif norm is False:
norm = None
self.norm = norm
self.norm_dims = norm_dims
self.rev = rev
@@ -437,10 +849,13 @@ class BlendMode:
allow_scale=_Empty,
rescale_dims=_Empty,
force_rescale=_Empty,
):
preserve_kwargs=True,
**kwargs: dict,
) -> object:
empty = self._Empty
kwargs = (self.f_kwargs | kwargs) if preserve_kwargs else kwargs
return self.__class__(
f if f is not empty else self.f,
f if f is not empty else self.f_raw,
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,
@@ -451,6 +866,7 @@ class BlendMode:
force_rescale=force_rescale
if force_rescale is not empty
else self.force_rescale,
**kwargs,
)
def rescale(self, t, *, rescale_dims=_Empty):
@@ -465,7 +881,7 @@ class BlendMode:
tmax = torch.amax(t, keepdim=True, dim=rescale_dims)
return (t - tmin).div_(tmax - tmin).clamp_(0, 1), tmin, tmax
def __call__(self, a, b, t, *, norm_dims=_Empty):
def __call__(self, a, b, t, *, norm_dims=_Empty) -> torch.Tensor:
if not self.force_rescale:
return self.__call__internal(a, b, t, norm_dims=norm_dims)
a, amin, amax = self.rescale(a)
@@ -476,7 +892,7 @@ class BlendMode:
del amin, amax, bmin, bmax
return result.mul_(rmax.sub_(rmin)).add_(rmin)
def __call__internal(self, a, b, t, *, norm_dims=_Empty):
def __call__internal(self, a, b, t, *, norm_dims=_Empty) -> torch.Tensor:
if not isinstance(t, torch.Tensor) and isinstance(a, torch.Tensor):
t = a.new_full((1,), t)
if self.rev:
@@ -490,18 +906,160 @@ class BlendMode:
)
class BlendingModes:
def __init__(self, builtins=None):
self.builtins = {} if builtins is None else builtins
self.cache = {}
def get(self, k: str, default=None):
result = self.builtins.get(k)
if result is not None:
return result
result = self.cache.get(k)
if result is not None:
return result
return self.try_extended(k, default=default)
_simple_value_map: ClassVar = {
"true": True,
"false": False,
"()": (),
"none": None,
}
def parse_value(self, k: str, v: str):
k = k.strip().lower()
v = v.strip()
if not v:
raise ValueError("Empty value")
if len(v) > 1 and v[0] == "^":
literal_mode = True
v = v[1:]
else:
literal_mode = False
vl = v.lower()
result = self._simple_value_map.get(vl, vl)
if result is not vl:
return result
v0 = v[0]
if v0.isdigit() or v0 in "-+":
if "," in v:
result = tuple(
self.parse_value(k, subv)
for subv in (_subv for _subv in v.split(",") if _subv.strip())
)
if len(result) > 1 and not all(
subv.__class__ is result[0].__class__ for subv in result[1:]
):
errstr = f"Mismatched items in list for key {k}"
raise ValueError(errstr)
return result
return float(v) if "." in v else int(v)
if not literal_mode and k.startswith("blend"):
# It won't be a numeric value here.
result = self.builtins.get(v)
if result is None:
errstr = f"Unknown blend mode {v}"
raise ValueError(errstr)
return result
return v
def parse_arg(self, s: str, idx: int) -> tuple:
kv = s.split("=", 1)
if len(kv) != 2:
errstr = f"Failed to parse argument at position {idx}"
raise ValueError(errstr)
k, v = kv[0].strip(), kv[1].strip()
if not k:
errstr = f"Empty key at argument position {idx}"
raise ValueError(errstr)
try:
v_out = self.parse_value(k, v)
except ValueError as exc:
errstr = f"Parse failed at argument position {idx}: {exc}"
raise ValueError(errstr) from exc
return (k, v_out)
def try_extended(self, k: str, default=None) -> object:
if ":" not in k:
return default
name, *arglist = k.strip().split(":")
name = name.strip()
base_bm = self.builtins.get(name)
if base_bm is None:
errstr = f"Unknown mode {name} for extended blend specification"
raise ValueError(errstr)
bm_kwargs = dict(self.parse_arg(arg, idx) for idx, arg in enumerate(arglist))
bm = base_bm.edited(**bm_kwargs)
self.cache[k] = bm
return bm
def items(self):
return self.builtins.items()
def values(self):
return self.builtins.values()
def __contains__(self, k: str) -> bool:
return self.get(k) is not None
def __iter__(self):
return self.builtins.__iter__()
keys = __iter__
def __setitem__(self, k: str, v) -> str:
self.builtins[k] = v if isinstance(v, BlendMode) else BlendMode(v)
def __getitem__(self, k: str):
result = self.get(k)
if result is None:
raise KeyError(k)
return result
def __ior__(self, other: dict | object):
if isinstance(other, dict):
self.builtins |= {
k: v if isinstance(v, BlendMode) else BlendMode(v)
for k, v in other.items()
}
return self
self.builtins |= other.builtins
self.cache |= other.cache
return self
def __or__(self, other: dict | object) -> object:
clone = self.__class__()
clone.builtins = self.builtins.copy()
clone.cache = self.cache.copy()
if isinstance(other, dict):
clone.builtins |= {
k: v if isinstance(v, BlendMode) else BlendMode(v)
for k, v in other.items()
}
return clone
clone.builtins |= other.builtins
clone.cache |= other.cache
return clone
def copy(self):
return self | {}
BLENDING_MODES = {
# Args:
# - a (tensor): Latent input 1
# - b (tensor): Latent input 2
# - t (float): Blending factor
"a_only": BlendMode(lambda a, _b, t: a * t, allow_scale=False),
"b_only": BlendMode(lambda _a, b, t: b * t, allow_scale=False),
# Interpolates between tensors a and b using normalized linear interpolation.
"bislerp": BlendMode(
lambda a, b, t: ((1 - t) * a).add_(t * b),
normalize,
),
# "nbislerp": BlendMode(lambda a, b, t: (1 - t) * a + t * b, normalize),
"slerp": BlendMode(lambda a, b, t: altslerp(a, b, t, dim=-1)),
"slerp": BlendMode(altslerp),
# Transfer the color from `b` to `a` by t` factor
"colorize": BlendMode(lambda a, b, t: (b - a).mul_(t).add_(a)),
# Interpolates between tensors a and b using cosine interpolation.
@@ -516,45 +1074,138 @@ BLENDING_MODES = {
# with a twist when t is greater than or equal to 0.5.
"hslerp": BlendMode(hslerp),
"hslerpalt": BlendMode(hslerp_alt2),
"hslerpalt110x": BlendMode(partial(hslerp_alt2, sign_order=(1.1, -1.1))),
"hslerpalt125x": BlendMode(partial(hslerp_alt2, sign_order=(1.25, -1.25))),
"hslerpalt150x": BlendMode(partial(hslerp_alt2, sign_order=(1.5, -1.5))),
"hslerpalt300x": BlendMode(partial(hslerp_alt2, sign_order=(3.0, -3.0))),
"hslerpaltflipsign": BlendMode(partial(hslerp_alt2, sign_order=(-1.0, 1.0))),
"hslerpaltflipsign110x": BlendMode(partial(hslerp_alt2, sign_order=(-1.1, 1.1))),
"hslerpaltflipsign125x": BlendMode(partial(hslerp_alt2, sign_order=(-1.25, 1.25))),
"hslerpaltflipsign150x": BlendMode(partial(hslerp_alt2, sign_order=(-1.5, 1.5))),
"hslerpaltflipsign300x": BlendMode(partial(hslerp_alt2, sign_order=(-3.0, 3.0))),
"problerp0.25": BlendMode(partial(stochasistic_blend, fuzz=0.25)),
"problerp0.1": BlendMode(partial(stochasistic_blend, fuzz=0.1)),
"problerp0.025": BlendMode(partial(stochasistic_blend, fuzz=0.025)),
"hslerpalt110x": BlendMode(hslerp_alt2, sign_order=(1.1, -1.1)),
"hslerpalt125x": BlendMode(hslerp_alt2, sign_order=(1.25, -1.25)),
"hslerpalt150x": BlendMode(hslerp_alt2, sign_order=(1.5, -1.5)),
"hslerpalt300x": BlendMode(hslerp_alt2, sign_order=(3.0, -3.0)),
"hslerpaltflipsign": BlendMode(hslerp_alt2, sign_order=(-1.0, 1.0)),
"hslerpaltflipsign110x": BlendMode(hslerp_alt2, sign_order=(-1.1, 1.1)),
"hslerpaltflipsign125x": BlendMode(hslerp_alt2, sign_order=(-1.25, 1.25)),
"hslerpaltflipsign150x": BlendMode(hslerp_alt2, sign_order=(-1.5, 1.5)),
"hslerpaltflipsign300x": BlendMode(hslerp_alt2, sign_order=(-3.0, 3.0)),
"problerp0.25": BlendMode(stochasistic_blend, fuzz=0.25),
"problerp0.1": BlendMode(stochasistic_blend, fuzz=0.1),
"problerp0.025": BlendMode(stochasistic_blend, fuzz=0.025),
"probselect": BlendMode(prob_blend),
"probselect_channels": BlendMode(prob_blend, collapse_dims=(1,)),
"probselectsmoothed": BlendMode(prob_blend_smoothed),
"probselectsmoothed_ks5": BlendMode(partial(prob_blend_smoothed, kernel_size=5)),
"probselectsmoothed_ks9": BlendMode(partial(prob_blend_smoothed, kernel_size=9)),
"probselectsmoothed_channels": BlendMode(
prob_blend_smoothed,
collapse_dims=(1,),
),
"probselectsmoothed_ks5": BlendMode(prob_blend_smoothed, kernel_size=5),
"probselectsmoothed_ks9": BlendMode(prob_blend_smoothed, kernel_size=9),
"probselectsmoothed_ks9_sigma3": BlendMode(
partial(
prob_blend_smoothed,
kernel_size=9,
sigma=3.0,
prob_blend_smoothed,
kernel_size=9,
sigma=3.0,
),
"probinject": BlendMode(
lambda a, b, t, **kwargs: prob_blend(torch.zeros_like(b), b, t, **kwargs).add_(
a,
),
),
"probsubtract_b": BlendMode(
lambda a, b, t, **kwargs: a - prob_blend(torch.zeros_like(b), b, t, **kwargs),
),
"gradient": BlendMode(gradient_blend),
# Adds tensor b to tensor a, scaled by t.
"inject": BlendMode(lambda a, b, t: (b * t).add_(a)),
"injecthalf": BlendMode(lambda a, b, t: (b * (t * 0.5)).add_(a)),
"injectquarter": BlendMode(lambda a, b, t: (b * (t * 0.25)).add_(a)),
"inject_difference": BlendMode(lambda a, b, t: (a - b).mul_(t).add_(a)),
"inject_copysign_a": BlendMode(lambda a, b, t: (b * t).add_(a).copysign_(a)),
"inject_copysign_b": BlendMode(lambda a, b, t: (b * t).add_(a).copysign_(b)),
"inject_avoidsign_a": BlendMode(lambda a, b, t: (b * t).add_(a).copysign_(a.neg())),
"inject_avoidsign_b": BlendMode(lambda a, b, t: (b * t).add_(a).copysign_(b.neg())),
# Interpolates between tensors a and b using linear interpolation.
"lerp": BlendMode(lambda a, b, t: ((1 - t) * a).add_(t * b)),
# "lerp": BlendMode(lambda a, b, t: ((1.0 - t) * a).add_(t * b)),
"lerp": BlendMode(torch.lerp),
"lerp050x": BlendMode(lambda a, b, t: (((1 - t) * a).add_(t * b)).mul_(0.5)),
"lerp075x": BlendMode(lambda a, b, t: (((1 - t) * a).add_(t * b)).mul_(0.75)),
"lerp110x": BlendMode(lambda a, b, t: (((1 - t) * a).add_(t * b)).mul_(1.1)),
"lerp125x": BlendMode(lambda a, b, t: (((1 - t) * a).add_(t * b)).mul_(1.25)),
"lerp150x": BlendMode(lambda a, b, t: (((1 - t) * a).add_(t * b)).mul_(1.5)),
"lerp_copysign_a": BlendMode(
lambda a, b, t: ((1.0 - t) * a).add_(t * b).copysign_(a),
),
"lerp_copysign_b": BlendMode(
lambda a, b, t: ((1.0 - t) * a).add_(t * b).copysign_(b),
),
"lerp_avoidsign_a": BlendMode(
lambda a, b, t: ((1.0 - t) * a).add_(t * b).copysign_(a.neg()),
),
"lerp_avoidsign_b": BlendMode(
lambda a, b, t: ((1.0 - t) * a).add_(t * b).copysign_(b.neg()),
),
# Simulates a brightening effect by adding tensor b to tensor a, scaled by t.
"lineardodge": BlendMode(lambda a, b, t: (b * t).add_(a)),
"copysign": BlendMode(lambda a, b, _t: torch.copysign(a, b)),
"probcopysign": BlendMode(lambda a, b, t: torch.copysign(a, prob_blend(a, b, t))),
"slice_flat_d1": BlendMode(slice_blend, dim=1, flatten=True),
"slice_flat_d2": BlendMode(slice_blend, dim=2, flatten=True),
"slice_d1": BlendMode(slice_blend, dim=1, flatten=False),
"slice_d2": BlendMode(slice_blend, dim=2, flatten=False),
"slice_d3": BlendMode(slice_blend, dim=3, flatten=False),
"slice_d1_flip": BlendMode(
slice_blend,
dim=1,
flatten=False,
flip_a=True,
flip_b=True,
flip_out=True,
),
"slice_d2_flip": BlendMode(
slice_blend,
dim=2,
flatten=False,
flip_a=True,
flip_b=True,
flip_out=True,
),
"slice_d3_flip": BlendMode(
slice_blend,
dim=3,
flatten=False,
flip_a=True,
flip_b=True,
flip_out=True,
),
"slicesmooth_d1": BlendMode(slice_blend_smooth, dim=1, flatten=False),
"slicesmooth_d2": BlendMode(slice_blend_smooth, dim=2, flatten=False),
"slicesmooth_d3": BlendMode(slice_blend_smooth, dim=3, flatten=False),
"loplerp_a098": BlendMode(lop_lerp, a_ratio=0.98),
"loplerp_a101": BlendMode(lop_lerp, a_ratio=1.01),
"loplerp_a102": BlendMode(lop_lerp, a_ratio=1.02),
"loplerp_a105": BlendMode(lop_lerp, a_ratio=1.05),
"cosinesimilarity": BlendMode(
cosine_similarity_blend,
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
),
"cosinesimilarity_flat": BlendMode(
cosine_similarity_blend_flat,
start_dim=1,
end_dim=-1,
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
),
"cosinesimilarity_flat_spatdims": BlendMode(
cosine_similarity_blend_flat,
start_dim=-2,
end_dim=-1,
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
),
"cosinesimilarity_avg_spatdims": BlendMode(
cosine_similarity_blend_avg,
dims=(-1, -2),
backend=DEFAULT_COSINE_SIMILARITY_BLEND_BACKEND,
),
"hybrid_lerp_cosinesimilarity": BlendMode(
blend_blend,
blend_mode_a="lerp",
blend_mode_b="cosinesimilarity_flat_spatdims",
blend_mode_blend="lerp",
blend_blend=0.5,
),
# 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),
@@ -631,6 +1282,11 @@ BLENDING_MODES = {
force_rescale=True,
),
"subtract": BlendMode(lambda a, b, t: a * t - b * t, allow_scale=False),
"subtract_b": BlendMode(lambda a, b, t: a - b * t, allow_scale=False),
"subtract_b_scaleup_a": BlendMode(
lambda a, b, t: a * (1.0 + t) - b * t,
allow_scale=False,
),
"vividlight": BlendMode(
lambda a, b, _t: torch.where(
b <= 0.5,
@@ -650,6 +1306,8 @@ BLENDING_MODES |= {
BLENDING_MODES |= {f"rev{k}": v.edited(rev=True) for k, v in BLENDING_MODES.items()}
BLENDING_MODES = BlendingModes(BLENDING_MODES)
BIDERP_MODES = {
k: v.edited(norm_dims=0)
for k, v in BLENDING_MODES.items()
+9 -3
View File
@@ -11,27 +11,33 @@ from . import (
taevid,
)
_blepping_integrations = None
NODE_CLASS_MAPPINGS = {
"BlehBlockCFG": blockCFG.BlockCFGBleh,
"BlehBlockOps": ops.BlehBlockOps,
"BlehCast": misc.BlehCast,
"BlehCFGInitSampler": samplers.BlehCFGInitSampler,
"BlehDeepShrink": deepShrink.DeepShrinkBleh,
"BlehDisableNoise": misc.BlehDisableNoise,
"BlehDiscardPenultimateSigma": misc.DiscardPenultimateSigma,
"BlehEnsurePreviewer": misc.BlehEnsurePreviewer,
"BlehForceSeedSampler": samplers.BlehForceSeedSampler,
"BlehGlobalSageAttention": sageAttention.BlehGlobalSageAttention,
"BlehHyperTile": hyperTile.HyperTileBleh,
"BlehImageAsLatent": misc.BlehImageAsLatent,
"BlehInsaneChainSampler": samplers.BlehInsaneChainSampler,
"BlehLatentAsImage": misc.BlehLatentAsImage,
"BlehLatentBlend": ops.BlehLatentBlend,
"BlehLatentOps": ops.BlehLatentOps,
"BlehLatentScaleBy": ops.BlehLatentScaleBy,
"BlehLatentBlend": ops.BlehLatentBlend,
"BlehModelPatchConditional": modelPatchConditional.ModelPatchConditionalNode,
"BlehModelPatchFastTerminate": misc.BlehModelPatchFastTerminate,
"BlehPlug": misc.BlehPlug,
"BlehRefinerAfter": refinerAfter.BlehRefinerAfter,
"BlehSageAttentionSampler": sageAttention.BlehSageAttentionSampler,
"BlehSetSamplerPreset": samplers.BlehSetSamplerPreset,
"BlehCast": misc.BlehCast,
"BlehSetSigmas": misc.BlehSetSigmas,
"BlehEnsurePreviewer": misc.BlehEnsurePreviewer,
"BlehTAEVideoDecode": taevid.TAEVideoDecode,
"BlehTAEVideoEncode": taevid.TAEVideoEncode,
}
+165 -1
View File
@@ -1,13 +1,18 @@
# ruff: noqa: TID252
from __future__ import annotations
import contextlib
import operator
import random
from decimal import Decimal
from functools import partial
import torch
from comfy import model_management
from comfy.model_management import throw_exception_if_processing_interrupted
from ..better_previews.previewer import ensure_previewer # noqa: TID252
from ..better_previews.previewer import ensure_previewer
from ..latent_utils import normalize_to_scale
class DiscardPenultimateSigma:
@@ -318,3 +323,162 @@ class BlehEnsurePreviewer:
def go(cls, *, any_input):
ensure_previewer()
return (any_input,)
class BlehImageAsLatent:
DESCRIPTION = "This node allows you to rearrange an IMAGE to look like a LATENT. Can be useful if you want to apply some latent operations to an IMAGE. Can be reversed with the BlehLatentAsImage node."
FUNCTION = "go"
CATEGORY = "latent/advanced"
RETURN_TYPES = ("LATENT",)
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"audio": ("IMAGE",),
"rescale": (
"BOOLEAN",
{
"default": True,
"tooltip": "When enabled, will rescale the image values which usually are from 0 to 1 to -1 to 1.",
},
),
},
}
@classmethod
def go(cls, *, image: torch.Tensor, rescale: bool) -> tuple:
image = image.to(device="cpu", dtype=torch.float32, copy=True)
if image.ndim == 3:
image = image[None]
elif image.ndim != 4:
raise ValueError("Unexpected number of dimensions in image")
image = image.movedim(-1, 1)
if rescale:
image = image.sub_(0.5).mul_(2.0)
return ({"samples": image},)
class BlehLatentAsImage:
DESCRIPTION = "This node lets you rearrange a LATENT to look like an IMAGE. Note: It does not respect anything like masks or latent selection metadata (from nodes like LatentFromBatch) that might exist."
FUNCTION = "go"
CATEGORY = "latent/advanced"
RETURN_TYPES = ("IMAGE",)
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"latent": ("LATENT",),
"values_mode": (
("rescale", "rescale_perchannel", "clamp"),
{"default": "rescale"},
),
"channels_into_batch": (
"BOOLEAN",
{
"default": False,
"tooltip": "When enabled, will create a greyscale image for each latent channel.",
},
),
},
}
@classmethod
def go(
cls,
*,
latent: dict,
values_mode: str,
channels_into_batch: bool,
) -> tuple:
samples = latent["samples"]
if samples.ndim != 4:
raise ValueError("Expected a 4D latent but didn't get one")
if channels_into_batch:
samples = samples.reshape(-1, *samples.shape[2:]).unsqueeze(1)
samples = samples.expand(samples.shape[0], 3, *samples.shape[2:])
image = samples.movedim(1, -1).to(
device="cpu",
dtype=torch.float32,
copy=True,
)[..., :4]
if values_mode == "clamp":
return (image.clamp(0.0, 1.0),)
image = normalize_to_scale(
image,
0.0,
1.0,
dim=(2, 3) if values_mode == "rescale_perchannel" else (1, 2, 3),
)
return (image,)
class BlehModelPatchFastTerminate:
RETURN_TYPES = ("MODEL",)
FUNCTION = "go"
CATEGORY = "hacks"
DESCRIPTION = "Patches a model to check if processing is interrupted at the start of every block. Makes interrupting generations more responsive on supported models (mainly useful for video models that might take 40+ second for a step). Should support most existing models."
@classmethod
def INPUT_TYPES(cls):
return {"required": {"model": ("MODEL",)}}
@classmethod
def go(cls, model):
m = model.clone()
def wrap_transformer_forward(orig_forward, *args: list, **kwargs: dict):
throw_exception_if_processing_interrupted()
return orig_forward(*args, **kwargs)
found = 0
for bt in (
"blocks",
"single_blocks",
"double_blocks",
"transformer_blocks",
"vace_blocks",
"double_stream_blocks",
"single_stream_blocks",
):
bn = 0
while True:
k = f"diffusion_model.{bt}.{bn}"
try:
block = model.get_model_object(k)
except AttributeError:
block = None
if block is None:
k = f"diffusion_model.{bt}.block{bn}"
with contextlib.suppress(AttributeError):
block = model.get_model_object(k)
bn += 1
if block is None:
break
orig_forward = getattr(block, "forward", None)
if orig_forward is None:
continue
m.add_object_patch(
f"{k}.forward",
partial(wrap_transformer_forward, orig_forward),
)
found += 1
if found > 0:
# Appears to be a transformer-based model so we're done.
return (m,)
# Fallthough to handling normal SD models.
def input_block_patch(h, _transformer_options):
throw_exception_if_processing_interrupted()
return h
def output_block_patch(h, hsp, _transformer_options):
throw_exception_if_processing_interrupted()
return h, hsp
m.set_model_input_block_patch(input_block_patch)
m.set_model_output_block_patch(output_block_patch)
return (m,)
+64 -34
View File
@@ -1,3 +1,4 @@
# ruff: noqa: PLR6104
from __future__ import annotations
import contextlib
@@ -5,6 +6,7 @@ import importlib
from typing import TYPE_CHECKING
import comfy.ldm.modules.attention as comfyattn
import torch
import yaml
from comfy.samplers import KSAMPLER
@@ -17,13 +19,16 @@ except ImportError:
sageattn_default_head_sizes = None
sageattn_default_function = None
try:
import spas_sage_attn
except ImportError:
spas_sage_attn = None
if TYPE_CHECKING:
import collections
from collections.abc import Callable
import torch
if sageattention is not None:
try:
@@ -43,7 +48,7 @@ else:
sageattn_version = "unknown"
def attention_sage(
def attention_bleh( # noqa: PLR0914
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
@@ -58,9 +63,20 @@ def attention_sage(
**kwargs: dict[str],
) -> torch.Tensor:
old_sageattn = sageattn_version[:2] in {"1.", "un"}
mask = kwargs.get("mask")
skip_reshape = kwargs.get("skip_reshape", False)
skip_output_reshape = kwargs.get("skip_output_reshape", False)
orig_attn_kwargs = {
k: kwargs.pop(k)
for k in ("mask", "skip_reshape", "skip_output_reshape", "attn_precision")
if k in kwargs.copy()
}
bleh_kwargs = {
k: kwargs.pop(k)
for k in kwargs.copy()
if k.startswith("sm_scale_")
or k in {"q_multiplier", "k_multiplier", "v_multiplier", "output_multiplier"}
}
mask = orig_attn_kwargs.get("mask")
skip_reshape = orig_attn_kwargs.get("skip_reshape", False)
skip_output_reshape = orig_attn_kwargs.get("skip_output_reshape", False)
batch = q.shape[0]
dim_head = q.shape[-1] // (1 if skip_reshape else heads)
enabled = sageattn_allow_head_sizes is None or dim_head in sageattn_allow_head_sizes
@@ -68,15 +84,11 @@ def attention_sage(
enabled = all(t.shape == q.shape for t in (k, v))
if sageattn_verbose:
print(
f"\n>> SAGE({enabled}): reshape={not skip_reshape}, output_reshape={not skip_output_reshape}, dim_head={q.shape[-1]}, heads={heads}, adj_heads={dim_head}, q={q.shape}, k={k.shape}, v={v.shape}, args: {kwargs}\n",
f"\n>> SAGE({enabled}): reshape={not skip_reshape}, output_reshape={not skip_output_reshape}, dim_head={q.shape[-1]}, heads={heads}, adj_heads={dim_head}, q={q.shape}, k={k.shape}, v={v.shape}, orig_attn_args={orig_attn_kwargs}, args: {kwargs}\n",
)
if not enabled:
filtered_kwargs = {
k: v
for k, v in kwargs.items()
if k in {"mask", "skip_reshape", "skip_output_reshape", "attn_precision"}
}
return orig_attention(q, k, v, heads, **filtered_kwargs)
return orig_attention(q, k, v, heads, **orig_attn_kwargs)
tensor_layout = kwargs.pop("tensor_layout", None)
if old_sageattn:
tensor_layout = "HND"
@@ -106,18 +118,27 @@ def attention_sage(
do_transpose = skip_output_reshape
if not old_sageattn:
kwargs["tensor_layout"] = tensor_layout
sm_scale_hd = kwargs.pop(f"sm_scale_{dim_head}", None)
sm_scale_hd = bleh_kwargs.pop(f"sm_scale_{dim_head}", None)
if sm_scale_hd is not None:
kwargs["sm_scale"] = sm_scale_hd
result = sageattn_function(
q,
k,
v,
is_causal=False,
attn_mask=mask,
dropout_p=0.0,
**kwargs,
q_multiplier = bleh_kwargs.get("q_multiplier", 1.0)
k_multiplier = bleh_kwargs.get("k_multiplier", 1.0)
v_multiplier = bleh_kwargs.get("v_multiplier", 1.0)
output_multiplier = bleh_kwargs.get("output_multiplier", 1.0)
if q_multiplier != 1.0:
q = q * q_multiplier
if k_multiplier != 1.0:
k = k * k_multiplier
if v_multiplier != 1.0:
v = v * v_multiplier
kwargs = {"is_causal": False, "dropout_p": 0.0, "attn_mask": mask} | kwargs
result = (
torch.zeros_like(q)
if output_multiplier == 0
else sageattn_function(q, k, v, **kwargs)
)
if output_multiplier not in {0, 1}:
result *= output_multiplier
if do_transpose:
result = result.transpose(1, 2)
if not skip_output_reshape:
@@ -138,28 +159,37 @@ def copy_funattrs(fun, dest=None):
return dest
def make_sageattn_wrapper(
def make_attn_wrapper(
*,
orig_attn,
sageattn_function: str = "sageattn",
**kwargs: dict,
):
outer_kwargs = kwargs
sageattn_function = getattr(sageattention, sageattn_function)
if sageattn_function.startswith("sparge") and spas_sage_attn is None:
raise ValueError(
"SpargeAttention is not available, make sure you have the spas_sage_attn Python package installed",
)
if sageattn_function == "sparge":
sageattn_function = spas_sage_attn.spas_sage2_attn_meansim_cuda
elif sageattn_function == "sparge1":
sageattn_function = spas_sage_attn.spas_sage_attn_meansim_cuda
else:
sageattn_function = getattr(sageattention, sageattn_function)
def attn(
*args: list,
_sage_outer_kwargs=outer_kwargs,
_sage_orig_attention=orig_attn,
_sage_sageattn_function=sageattn_function,
_sage_attn=attention_sage,
_bleh_outer_kwargs=outer_kwargs,
_bleh_orig_attention=orig_attn,
_bleh_attn_function=sageattn_function,
_bleh_attn=attention_bleh,
**kwargs: dict,
) -> torch.Tensor:
return _sage_attn(
return _bleh_attn(
*args,
orig_attention=_sage_orig_attention,
sageattn_function=_sage_sageattn_function,
**_sage_outer_kwargs,
orig_attention=_bleh_orig_attention,
sageattn_function=_bleh_attn_function,
**_bleh_outer_kwargs,
**kwargs,
)
@@ -175,7 +205,7 @@ def sageattn_context(
yield None
return
orig_attn = copy_funattrs(comfyattn.optimized_attention)
attn = make_sageattn_wrapper(orig_attn=orig_attn, **kwargs)
attn = make_attn_wrapper(orig_attn=orig_attn, **kwargs)
try:
copy_funattrs(attn, comfyattn.optimized_attention)
yield None
@@ -246,7 +276,7 @@ class BlehGlobalSageAttention:
)
if not cls.orig_attn:
cls.orig_attn = copy_funattrs(comfyattn.optimized_attention)
attn = make_sageattn_wrapper(
attn = make_attn_wrapper(
orig_attn=cls.orig_attn,
**get_yaml_parameters(yaml_parameters),
)
+122 -13
View File
@@ -2,9 +2,10 @@ from __future__ import annotations
import contextlib
import importlib
import math
import random
from copy import deepcopy
from functools import partial
from functools import partial, update_wrapper
from os import environ
from typing import Any, Callable, NamedTuple
@@ -141,20 +142,19 @@ class BlehForceSeedSampler:
FUNCTION = "go"
def go(
self,
sampler: object,
seed_offset: int | None = 1,
) -> tuple[KSAMPLER, SamplerChain]:
def go(self, sampler: object, seed_offset: int | None = 1) -> tuple[KSAMPLER]:
return (
KSAMPLER(
self.sampler_function,
extra_options=sampler.extra_options
| {
"bleh_wrapped_sampler": sampler,
"bleh_seed_offset": seed_offset,
},
inpaint_options=sampler.inpaint_options | {},
update_wrapper(
partial(
self.sampler_function,
bleh_wrapped_sampler=sampler,
bleh_seed_offset=seed_offset,
),
sampler.sampler_function,
),
extra_options=sampler.extra_options.copy(),
inpaint_options=sampler.inpaint_options.copy(),
),
)
@@ -325,3 +325,112 @@ class BlehSetSamplerPreset:
)
BLEH_PRESET[preset] = (deepcopy(sampler), sigmas)
return (any_input,)
class BlehCFGInitSampler:
DESCRIPTION = "Sampler wrapper that allows skipping some number of initial steps, similar to CFGZeroStar zero-init."
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sampler": (
"SAMPLER",
{
"tooltip": "Connect the sampler you want to wrap here. It will be called to sample as normal after the configured number of skipped steps.",
},
),
"steps": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 9999.0,
"tooltip": "Number of steps to skip before sampling. You can use a fractional value here, but it may not work well. Whole skipped steps do not require a model call.",
},
),
"mode": (
("zero", "afs", "afs_flow_hack", "scale_down"),
{
"default": "zero",
"tooltip": "zero: Works like CFGZeroStar zero init (just skips steps).\nafs: Analytical first step mode. A method of scaling down the initial noise to match skipped steps.\nafs_hack: A version of AFS mode that may work better for flow models.\nscale_down: The simplest approach to scaling down the latent to match the skipped steps.",
},
),
},
}
FUNCTION = "go"
def go(
self,
*,
sampler: object,
steps: float,
mode: str,
) -> tuple[KSAMPLER]:
if steps == 0:
return (sampler,)
sampler_function = update_wrapper(
partial(
self.sampler_function,
bleh_ci_wrapped_sampler=sampler,
bleh_ci_mode=mode,
bleh_ci_steps=steps,
),
sampler.sampler_function,
)
return (
KSAMPLER(
sampler_function,
extra_options=sampler.extra_options.copy(),
inpaint_options=sampler.inpaint_options.copy(),
),
)
@staticmethod
def sampler_function(
model: object,
x: torch.Tensor,
sigmas: torch.Tensor,
*args: list[Any],
extra_args: dict[str, Any] | None = None,
bleh_ci_wrapped_sampler: object | None = None,
bleh_ci_mode: str = "zero",
bleh_ci_steps: float = 0,
**kwargs: dict[str, Any],
) -> torch.Tensor:
if not bleh_ci_wrapped_sampler:
raise ValueError("Wrapped sampler missing!")
sigmas = sigmas.clone()
whole_steps = int(bleh_ci_steps)
steps = math.ceil(bleh_ci_steps)
step_fraction = bleh_ci_steps - whole_steps
for idx in range(steps if bleh_ci_mode != "zero" else 0):
sigma, sigma_next = sigmas[idx : idx + 2]
dt = sigma_next - sigma
if idx == whole_steps:
dt *= step_fraction
# From https://arxiv.org/abs/2210.05475
if bleh_ci_mode == "afs":
d = x / (1.0 + sigma**2) ** 0.5
elif bleh_ci_mode == "afs_flow_hack":
d = x / ((1.0 + (sigma * 10.0) ** 2) ** 0.5) / 10.0
elif bleh_ci_mode == "scale_down":
d = x / sigma
else:
raise ValueError("Bad CFG init mode")
x = x + d * dt # noqa: PLR6104
if steps > 0 and step_fraction != 0:
sigmas[whole_steps] += (sigmas[steps] - sigmas[whole_steps]) * step_fraction
return bleh_ci_wrapped_sampler.sampler_function(
model,
x,
sigmas[whole_steps:],
*args,
extra_args=extra_args,
**kwargs,
)
+5
View File
@@ -16,8 +16,13 @@ class Settings:
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)
self.btp_throttle_secs_fallback = btp.get("throttle_secs_fallback")
if self.btp_throttle_secs_fallback is None:
self.btp_throttle_secs_fallback = self.btp_throttle_secs
self.btp_skip_upscale_layers = btp.get("skip_upscale_layers", 0)
self.btp_preview_device = btp.get("preview_device")
# default, keep, float32, float16, bfloat16
self.btp_preview_dtype = btp.get("preview_dtype")
self.btp_maxed_batch_step_mode = btp.get("maxed_batch_step_mode", False)
self.btp_compile_previewer = btp.get("compile_previewer", False)
self.btp_oom_fallback = btp.get("oom_fallback", "latent2rgb")