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:
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user