Add BlehContrastiveOrthoCFG node.

More features for the Pythagorean LERP blend mode.
Minor cleanups.
This commit is contained in:
blepping
2026-05-18 21:28:45 -06:00
parent 14db301716
commit 74bd9e87f8
5 changed files with 263 additions and 16 deletions
+5 -1
View File
@@ -2,7 +2,11 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 2026TBD
## 20260518
* Added `BlehConstrastiveOrthoCFG` node.
## 20260515
* Support LTX 2.0, LTX 2.3 and LTX 2.3 wide TAE models.
* Added a `BlehFixGuiderPreviewing` node. Using this is necessary for LTX previews, it can also be used to set the LTX 2.3 wide previewer mode or override the default FPS when generating video previews.
+43 -8
View File
@@ -2562,19 +2562,54 @@ def pythagorean_lerp(
b: torch.Tensor,
t: float | torch.Tensor,
*,
growth_power: float = 0.0,
# 0 disables variance clamping.
max_variance: float = 0.0,
# Zero or negative values disable soft clamping.
stiffness: float = 10.0,
# Generally should be left at 1. Only applies when clamping variance.
magnitude_min: float = 1.0,
# Using non-default power will result in something completely different from
# Pythagorean LERP. These generally should be left on the defaults.
power: float = 2.0,
inv_power: float | None = None,
eps: float = 1e-08,
) -> torch.Tensor:
w_a = 1.0 - t
w_b = t
if inv_power is None:
inv_power = 1.0 / power
if not isinstance(t, torch.Tensor):
t = a.new_tensor(t)
t_inv = 1.0 - t
# Calculate how much the variance would shrink.
if isinstance(w_a, torch.Tensor):
variance_shrink = (w_a**2).add_(w_b**2).sqrt_().clamp_min_(eps)
raw_magnitude = magnitude = (
t_inv.abs()
.pow_(power)
.add_(t.abs().pow_(power))
.pow_(inv_power)
.clamp_min_(eps)
)
if growth_power != 1.0:
magnitude = raw_magnitude**growth_power
if max_variance != 0:
magnitude = soft_clamp(
magnitude,
min_val=magnitude_min,
max_val=max_variance,
stiffness=stiffness,
)
if magnitude is raw_magnitude:
w_a, w_b = t_inv, t
else:
variance_shrink = max(eps, (w_a**2 + w_b**2) ** 0.5)
# Then scale the weights to compensate.
return a.mul(w_a / variance_shrink).add_(b * (w_b / variance_shrink))
factor = (
raw_magnitude
if magnitude is raw_magnitude
else magnitude.div_(
raw_magnitude.abs().clamp_min_(eps).copysign(raw_magnitude),
)
)
w_a, w_b = t_inv * factor, t * factor
return (a * w_a).add_(b * w_b)
def rms_interpolation(
+1
View File
@@ -45,6 +45,7 @@ NODE_CLASS_MAPPINGS = {
"BlehModelProcessLatentOut": misc.BlehModelProcessLatentOut,
"BlehFixGuiderPreviewing": misc.BlehFixGuiderPreviewing,
"BlehBlendConditioning": misc.BlehBlendConditioning,
"BlehContrastiveOrthoCFG": misc.BlehContrastiveOrthoCFG,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+210 -6
View File
@@ -9,13 +9,13 @@ import random
from decimal import Decimal
from functools import partial
from itertools import pairwise
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, NamedTuple, Sequence
import torch
from comfy import model_management
from comfy.model_management import throw_exception_if_processing_interrupted
from .. import latent_utils
from .. import latent_utils as lutils
from ..better_previews.previewer import PREVIEWER_STATE, ensure_previewer
if TYPE_CHECKING:
@@ -113,7 +113,7 @@ class BlehDisableNoise:
)
class Wildcard(str): # noqa: FURB189
class Wildcard(str):
__slots__ = ()
def __ne__(self, _unused):
@@ -412,7 +412,7 @@ class BlehLatentAsImage:
)[..., :4]
if values_mode == "clamp":
return (image.clamp(0.0, 1.0),)
image = latent_utils.normalize_to_scale(
image = lutils.normalize_to_scale(
image,
0.0,
1.0,
@@ -871,7 +871,7 @@ class BlehBlendConditioning:
"conditioning_1": ("CONDITIONING",),
"conditioning_2": ("CONDITIONING",),
"blend_mode": (
tuple(latent_utils.BLENDING_MODES.keys()),
tuple(lutils.BLENDING_MODES.keys()),
{"default": "lerp"},
),
"strength": (
@@ -946,7 +946,7 @@ class BlehBlendConditioning:
) -> tuple:
blend_tensor_list = (s.strip() for s in blend_tensors.split(","))
blend_tensor_set = {s for s in blend_tensor_list if s}
blend_function = latent_utils.BLENDING_MODES[blend_mode]
blend_function = lutils.BLENDING_MODES[blend_mode]
mdbm = metadata_base_mode
# Initialize our stateful blender
@@ -1044,3 +1044,207 @@ class BlehBlendConditioning:
blend_full_items=True,
)
return (result,)
class ContrastiveOrthoCFG(NamedTuple):
start_sigma: float = 99999.0
end_sigma: float = 0.0
positive_scale: float = 0.2
negative_scale: float = 1.0
start_dim: int = 1
end_dim: int = 1
use_noise: bool = False
base_denoised: bool = False
def extract_common(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
return lutils.contrastive_ortho_cfg_base_a(
cond,
uncond,
1.0,
a_ortho_scale=-1.0,
b_ortho_scale=1.0,
start_dim=self.start_dim,
end_dim=self.end_dim,
)
def extract_parts(
self,
cond: torch.Tensor,
uncond: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
common = self.extract_common(cond, uncond)
unique_cond = cond - common
unique_uncond = uncond - common
if self.negative_scale != 1:
unique_uncond *= self.negative_scale
if self.positive_scale != 1:
unique_cond *= self.positive_scale
return common, unique_cond, unique_uncond
def check_sigma(self, sigma: torch.Tensor) -> bool:
sigma_f = sigma.mean().detach().cpu().item()
return self.end_sigma <= sigma_f <= self.start_sigma
def pad_sigma(self, sigma: torch.Tensor, ndim: int) -> torch.Tensor:
return (
sigma
if sigma.ndim == 0
else sigma.reshape(sigma.shape[0], *((1,) * (ndim - 1)))
)
def precfg_patch(self, args: dict[str, Any]) -> Sequence[torch.Tensor]:
x = args["input"]
sigma = args["sigma"]
conds_out = args["conds_out"]
if len(conds_out) < 2 or conds_out[1] is None or not self.check_sigma(sigma):
return conds_out
sigma = self.pad_sigma(sigma, x.ndim)
cond, uncond = conds_out[:2]
if self.use_noise:
cond = (x - cond).div_(sigma)
uncond = (x - uncond).div_(sigma)
common, cond_new, uncond_new = self.extract_parts(cond, uncond)
cond_new += common
uncond_new += common
if self.use_noise:
cond_new = cond_new.mul_(-sigma).add_(x)
uncond_new = uncond_new.mul_(-sigma).add_(x)
return conds_out.__class__((cond_new, uncond_new, *conds_out[2:]))
def postcfg_patch(self, args: dict[str, Any]) -> torch.Tensor:
x = args["input"]
cond = args["denoised"] if self.base_denoised else args["cond_denoised"]
uncond = args.get("uncond_denoised")
sigma = args["sigma"]
denoised = args["denoised"]
if not self.check_sigma(sigma) or uncond is cond or uncond is None:
return denoised
sigma = self.pad_sigma(sigma, x.ndim)
if self.use_noise:
cond, uncond, denoised = (
(x - t).div_(sigma) for t in (cond, uncond, denoised)
)
common, cond_unique, uncond_unique = self.extract_parts(cond, uncond)
denoised = cond_unique.sub_(uncond_unique).add_(
common if self.base_denoised else denoised,
)
return denoised.mul_(-sigma).add_(x) if self.use_noise else denoised
class BlehContrastiveOrthoCFG:
DESCRIPTION = "Contrastive orthogonal CFG. CFG variant that allows you to control the scale of negative and positive features individually. For the most predictable effects, use positive scales and pre_cfg mode with use_noise disabled. This is just a slower version of CFG if you use the default parameters and positive/negative scales of 1.0."
FUNCTION = "go"
OUTPUT_NODE = False
CATEGORY = "advanced/guidance"
RETURN_TYPES = ("MODEL",)
@classmethod
def INPUT_TYPES(cls):
dc = ContrastiveOrthoCFG()
return {
"required": {
"model": ("MODEL",),
},
"optional": {
"start_sigma": (
"FLOAT",
{
"default": dc.start_sigma,
"max": 99999.0,
"min": 0.0,
},
),
"end_sigma": (
"FLOAT",
{
"default": dc.end_sigma,
"max": 99999.0,
"min": 0.0,
},
),
"positive_scale": (
"FLOAT",
{
"default": dc.positive_scale,
"min": -9999.0,
"max": 9999.0,
"tooltip": "Scale for features unique to the positive prompt (cond). In pre-CFG mode, this gets multiplied by the CFG scale. For example, at CFG 5, the default of 0.2 would result in roughly the same strength as CFG 1.",
},
),
"negative_scale": (
"FLOAT",
{
"default": dc.negative_scale,
"min": -9999.0,
"max": 9999.0,
"tooltip": "Scale for features unique to the negative prompt (uncond). This gets subtracted. In pre-CFG mode the scale is effectively multiplied by CFG.",
},
),
"patch_mode": (
("pre_cfg", "post_cfg", "post_cfg_base_denoised"),
{
"default": "pre_cfg",
"tooltip": "pre_cfg: This mode applies the positive change to cond and the negative change to uncond and lets the CFG function take care of subtracting the negative part and adding the positive part. Since CFG 1 is normally just cond, at CFG one you will get the positive change applied as expected but the negative side will have no effect. Or in other words, result is basically common + unique_positive * CFG - unique_negative * (CFG - 1).\n\npost_cfg: The unique negative features are subtracted at exactly the scale you specify and the unique positive features are added in the same way. However, these are relative to the original cond/uncond generations but are applied to the result of CFG.\n\npost_cfg_base_denoised: This is like post_cfg mode except the positive side is what's unique to denoised (the result of CFG) and both parts are added to a common base. The results can be weird if you have other CFG type effects (I.E. CFG++) running afterward because what's unique to denoised will also include the negative side of uncond. Experimental mode, generally not recommended.",
},
),
"start_dim": (
"INT",
{
"default": dc.start_dim,
"min": -999,
"max": 999,
"tooltip": "Start dimension (zero-based) for determining orthogonal features. Image models typically use dimensions BATCH, CHANNELS, HEIGHT, WIDTH. Video models insert a FRAMES dimension after CHANNELS. The default is to normalize over channels.",
},
),
"end_dim": (
"INT",
{
"default": dc.end_dim,
"min": -999,
"max": 999,
"tooltip": "End dimension (zero-based) for determining orthogonal features. Image models typically use dimensions BATCH, CHANNELS, HEIGHT, WIDTH. Video models insert a FRAMES dimension after CHANNELS. The default is to normalize over channels.",
},
),
"use_noise": (
"BOOLEAN",
{
"default": False,
"tooltip": "Apply CFG to the noise prediction instead of the clean image. CFG normally uses the clean image. Experimental option and generally it's harder to separate out what's orthogonal from noise compared to a clean latent.",
},
),
"force_uncond_generation": (
"BOOLEAN",
{
"default": False,
"tooltip": "Disables the normal optimization that skip generating uncond (negative prompt) when CFG is 1.",
},
),
},
}
@classmethod
def go(
cls,
*,
model,
patch_mode: str = "pre_cfg",
force_uncond_generation: bool = False,
**kwargs: Any,
) -> tuple:
patch_object = ContrastiveOrthoCFG(
base_denoised=patch_mode.endswith("_base_denoised"),
**kwargs,
)
model = model.clone()
if patch_mode == "pre_cfg":
model.set_model_sampler_pre_cfg_function(
patch_object.precfg_patch,
disable_cfg1_optimization=force_uncond_generation,
)
else:
model.set_model_sampler_post_cfg_function(
patch_object.postcfg_patch,
disable_cfg1_optimization=force_uncond_generation,
)
return (model,)
+4 -1
View File
@@ -7,7 +7,7 @@ import random
from copy import deepcopy
from functools import partial, update_wrapper
from os import environ
from typing import Any, Callable, NamedTuple
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
from comfy.samplers import KSAMPLER, KSampler, k_diffusion_sampling
@@ -15,6 +15,9 @@ from tqdm import tqdm
from .misc import Wildcard
if TYPE_CHECKING:
from collections.abc import Callable
BLEH_PRESET_LIMIT = 16
BLEH_PRESET_COUNT = 1
with contextlib.suppress(Exception):