Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a70d1d09bb | ||
|
|
b034c3f09f | ||
|
|
2583c237f2 | ||
|
|
c83958c457 | ||
|
|
a0b2d99bc7 | ||
|
|
ba145c4722 | ||
|
|
c21dbfe1e5 | ||
|
|
fc38b5c998 | ||
|
|
d26fcf6fc8 | ||
|
|
3d8827f132 | ||
|
|
62bef2e275 | ||
|
|
fecdfe01df | ||
|
|
c5610837e4 | ||
|
|
21ac7cf0cb | ||
|
|
4930d862d5 | ||
|
|
e8b727f914 | ||
|
|
389aedfd17 | ||
|
|
99824b2ee5 | ||
|
|
4bb226aabb | ||
|
|
60b5127cf4 | ||
|
|
e72afd4189 | ||
|
|
98a9b6d656 | ||
|
|
e0988d3b24 | ||
|
|
932c7b2136 | ||
|
|
50ddd2ac47 | ||
|
|
3768f4a768 | ||
|
|
e82622cefe | ||
|
|
793914ced3 |
@@ -0,0 +1,27 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'Reithan' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
submodules: true
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
personal_access_token: ${{ secrets.COMFY_REGISTRY_KEY }}
|
||||
@@ -172,3 +172,8 @@ cython_debug/
|
||||
|
||||
# PyPI configuration file
|
||||
.pypirc
|
||||
|
||||
|
||||
# AI agents
|
||||
.claude/settings.local.json
|
||||
.claude/docs/**
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 33 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 4.8 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 12 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 4.7 MiB |
@@ -1 +0,0 @@
|
||||
from .nodes_NRS import *
|
||||
+257
-151
@@ -1,178 +1,284 @@
|
||||
import logging
|
||||
import torch
|
||||
from enum import Enum, auto
|
||||
from typing import Any
|
||||
|
||||
|
||||
class PredictionType(Enum):
|
||||
EPS = auto() # ε-prediction
|
||||
V = auto() # v-prediction
|
||||
X0 = auto() # x₀-prediction
|
||||
UNKNOWN = auto() # couldn’t detect / new scheduler
|
||||
|
||||
_RAW_TO_ENUM = {
|
||||
"eps": PredictionType.EPS,
|
||||
"epsilon": PredictionType.EPS,
|
||||
"flux": PredictionType.EPS,
|
||||
"chroma": PredictionType.EPS,
|
||||
"flow": PredictionType.EPS, # FLOW models (WAN, etc.) are EPS-compatible
|
||||
"wan": PredictionType.EPS, # WAN21 is FLOW-based
|
||||
"const": PredictionType.EPS, # CONST prediction class used in FLOW models
|
||||
"v": PredictionType.V,
|
||||
"v_prediction": PredictionType.V,
|
||||
"x0": PredictionType.X0,
|
||||
"sample": PredictionType.X0,
|
||||
}
|
||||
|
||||
class NRS:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "model": ("MODEL",),
|
||||
"skew": ("FLOAT", {"default": 2.0, "min": -30.0, "max": 30.0, "step": 0.01}),
|
||||
"stretch": ("FLOAT", {"default": 2.0, "min": -30.0, "max": 30.0, "step": 0.01}),
|
||||
"squash": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
return {"required": { "model": ("MODEL", {"tooltip": "Input model to apply NRS to"}),
|
||||
"skew": ("FLOAT", {"default": 2.00, "min": -30.0, "max": 30.0, "step": 0.01,
|
||||
"tooltip": "Changes the 'direction' of generation, steering away from negative prompt elements. Start with CFG/2."}),
|
||||
"stretch": ("FLOAT", {"default": 5.00, "min": -30.0, "max": 30.0, "step": 0.01,
|
||||
"tooltip": "Intensifies positive prompt elements. Start with your normal CFG value."}),
|
||||
"squash": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 1.0, "step": 0.01,
|
||||
"tooltip": "Softens Skew/Stretch effects, adding micro-detailing. Keep low initially."}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "advanced/model"
|
||||
|
||||
DESCRIPTION = "Negative Rejection Steering (NRS) replaces CFG with more nuanced guidance. IMPORTANT: Set your KSampler CFG to any value (it will be ignored). Connect your model through this node before sampling."
|
||||
|
||||
def _get_pred_type(self, model) -> PredictionType:
|
||||
"""
|
||||
In order to support Comfy, Forge, and possibly other models
|
||||
and various loaders.
|
||||
Walk common wrappers until we find something that looks like a
|
||||
prediction-type flag, then map it to the enum.
|
||||
Defaults to EPS if all else fails.
|
||||
"""
|
||||
def _canon(p):
|
||||
if p is None:
|
||||
return ""
|
||||
if isinstance(p, bytes):
|
||||
p = p.decode(errors="ignore")
|
||||
if isinstance(p, Enum):
|
||||
p = p.name
|
||||
return str(p).strip().lower()
|
||||
|
||||
# Breadth-first search through a few well-known wrappers.
|
||||
queue, seen = [model], set()
|
||||
|
||||
while queue:
|
||||
obj = queue.pop(0)
|
||||
|
||||
# 1) direct hit on this object ---------------------------------
|
||||
for attr in ("model_type", "prediction_type", "parameterization"):
|
||||
p = _canon(getattr(obj, attr, None))
|
||||
if p:
|
||||
pred_type = _RAW_TO_ENUM.get(p, PredictionType.UNKNOWN)
|
||||
if pred_type != PredictionType.UNKNOWN:
|
||||
logging.debug(f"NRS._get_pred_type: Found prediction type '{p}' from attribute '{attr}' -> {pred_type}")
|
||||
return pred_type
|
||||
|
||||
# 2) enqueue child containers we care about -------------------
|
||||
for attr in ("model", "diffusion_model", "config", "scheduler", "inner_model", "model_sampling"):
|
||||
child = getattr(obj, attr, None)
|
||||
if child is not None and id(child) not in seen:
|
||||
seen.add(id(child))
|
||||
queue.append(child)
|
||||
|
||||
# 3) enhanced detection for FLOW models (WAN, Flux, etc.) -------
|
||||
try:
|
||||
# Check model_sampling class type for FLOW models
|
||||
if hasattr(model, "model_sampling") and model.model_sampling is not None:
|
||||
sampling_class_name = type(model.model_sampling).__name__.lower()
|
||||
logging.debug(f"NRS._get_pred_type: Found model_sampling class: {sampling_class_name}")
|
||||
|
||||
# CONST class is used by FLOW models (WAN21, Flux, etc.)
|
||||
if "const" in sampling_class_name:
|
||||
logging.debug(f"NRS._get_pred_type: Detected FLOW model via CONST sampling class -> EPS")
|
||||
return PredictionType.EPS
|
||||
elif "v_prediction" in sampling_class_name:
|
||||
logging.debug(f"NRS._get_pred_type: Detected V-prediction model via sampling class -> V")
|
||||
return PredictionType.V
|
||||
elif "eps" in sampling_class_name:
|
||||
logging.debug(f"NRS._get_pred_type: Detected EPS model via sampling class -> EPS")
|
||||
return PredictionType.EPS
|
||||
|
||||
# Check model.model.model_type enum for newer models
|
||||
if hasattr(model, "model") and hasattr(model.model, "model_type"):
|
||||
model_type_str = _canon(str(model.model.model_type))
|
||||
logging.debug(f"NRS._get_pred_type: Found model.model.model_type: {model_type_str}")
|
||||
|
||||
if "flow" in model_type_str or "flux" in model_type_str:
|
||||
logging.debug(f"NRS._get_pred_type: Detected FLOW/Flux model via model_type -> EPS")
|
||||
return PredictionType.EPS
|
||||
elif "v_prediction" in model_type_str:
|
||||
logging.debug(f"NRS._get_pred_type: Detected V-prediction model via model_type -> V")
|
||||
return PredictionType.V
|
||||
elif "eps" in model_type_str:
|
||||
logging.debug(f"NRS._get_pred_type: Detected EPS model via model_type -> EPS")
|
||||
return PredictionType.EPS
|
||||
|
||||
except Exception as e:
|
||||
logging.debug(f"NRS._get_pred_type: Exception during enhanced detection: {e}")
|
||||
|
||||
# 4) safe default (matches docstring promise) --------------------
|
||||
logging.warning(f"NRS._get_pred_type: Could not determine prediction type for model. Using EPS as fallback.")
|
||||
logging.debug(f"NRS._get_pred_type: Model structure: {[attr for attr in dir(model) if not attr.startswith('_')]}")
|
||||
return PredictionType.EPS
|
||||
|
||||
def _convert_to_eps_space(self, x_orig, sig_root, sigma, cond, uncond):
|
||||
x_div = None
|
||||
eps_cond = cond
|
||||
eps_uncond = uncond
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# v → ε conversion
|
||||
logging.debug(f"NRS._convert_to_eps_space: generating x_div, eps_cond, and eps_uncond for v-pred")
|
||||
x_div = x_orig / (sigma ** 2 + 1)
|
||||
|
||||
eps_cond = ((x_div - (x_orig - cond)) * sig_root) / (sigma)
|
||||
eps_uncond = ((x_div - (x_orig - uncond)) * sig_root) / (sigma)
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
logging.debug(f"NRS._convert_to_eps_space: already in eps, no pre-scale needed")
|
||||
pass # already in ε space
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._convert_to_eps_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS (should not happen with enhanced detection)
|
||||
logging.warning(f"NRS._convert_to_eps_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
pass
|
||||
|
||||
return x_div, eps_cond, eps_uncond
|
||||
|
||||
def _finalize_from_eps_space(self, x_orig, x_div, x_final, sig_root, sigma):
|
||||
nrs_result = x_final
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# ε → v conversion
|
||||
logging.debug(f"NRS._finalize_from_eps_space: generating cfg_result for v-pred")
|
||||
nrs_result = x_orig - (x_div - x_final * sigma / sig_root)
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
# already in ε space
|
||||
logging.debug(f"NRS._finalize_from_eps_space: already in eps, no post-scale needed")
|
||||
pass
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._finalize_from_eps_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS (should not happen with enhanced detection)
|
||||
logging.warning(f"NRS._finalize_from_eps_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
pass
|
||||
return nrs_result
|
||||
|
||||
def _convert_to_v_space(self, x_orig, sig_root, sigma, cond, uncond):
|
||||
x_div = None
|
||||
v_cond = cond
|
||||
v_uncond = uncond
|
||||
if self.__pred_type == PredictionType.V:
|
||||
logging.debug(f"NRS._convert_to_v_space: already in v, no pre-scale needed")
|
||||
pass # already in v space
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
# ε → v conversion
|
||||
logging.debug(f"NRS._convert_to_v_space: generating x_div, v_cond, and v_uncond for eps")
|
||||
x_div = x_orig / (sigma ** 2 + 1)
|
||||
factor = sigma / sig_root
|
||||
|
||||
v_cond = x_orig - (x_div - cond * factor)
|
||||
v_uncond = x_orig - (x_div - uncond * factor)
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._convert_to_v_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS and convert to V-space
|
||||
logging.warning(f"NRS._convert_to_v_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
logging.debug(f"NRS._convert_to_v_space: generating x_div, v_cond, and v_uncond for eps (fallback)")
|
||||
x_div = x_orig / (sigma ** 2 + 1)
|
||||
factor = sigma / sig_root
|
||||
v_cond = x_orig - (x_div - cond * factor)
|
||||
v_uncond = x_orig - (x_div - uncond * factor)
|
||||
|
||||
return x_div, v_cond, v_uncond
|
||||
|
||||
def _finalize_from_v_space(self, x_orig, x_div, x_final, sig_root, sigma):
|
||||
nrs_result = x_final
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# already in v space
|
||||
logging.debug(f"NRS._finalize_from_v_space: already in v, no post-scale needed")
|
||||
pass
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
# v → ε conversion
|
||||
logging.debug(f"NRS._finalize_from_v_space: generating cfg_result for eps")
|
||||
nrs_result = (x_div - (x_orig - x_final)) * (sig_root / sigma)
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._finalize_from_v_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS and convert from V-space
|
||||
logging.warning(f"NRS._finalize_from_v_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
logging.debug(f"NRS._finalize_from_v_space: generating cfg_result for eps (fallback)")
|
||||
nrs_result = (x_div - (x_orig - x_final)) * (sig_root / sigma)
|
||||
return nrs_result
|
||||
|
||||
def patch(self, model, skew, stretch, squash):
|
||||
self.__pred_type = self._get_pred_type(model) if not hasattr(self, "__pred_type") else self.__pred_type
|
||||
self.__OPERATION_SPACE = PredictionType.V
|
||||
|
||||
def nrs(args):
|
||||
logging.debug(f"NRS.nrs: Skew: {skew}, Stretch: {stretch}, Squash: {squash}")
|
||||
# self.__pred_type = self.__pred_type if self.__pred_type is not None else self._get_pred_type(model)
|
||||
cond = args["cond"]
|
||||
uncond = args["uncond"]
|
||||
x_orig = args["input"]
|
||||
|
||||
sigma = args["sigma"]
|
||||
sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1))
|
||||
x_orig = args["input"]
|
||||
sig_root = (sigma ** 2 + 1).sqrt()
|
||||
|
||||
x_div, nrs_cond, nrs_uncond = None, None, None
|
||||
match self.__OPERATION_SPACE:
|
||||
case PredictionType.V:
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_v_space(x_orig, sig_root, sigma, cond, uncond)
|
||||
case PredictionType.EPS:
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_eps_space(x_orig, sig_root, sigma, cond, uncond)
|
||||
case PredictionType.X0:
|
||||
raise RuntimeError("NRS.nrs: x0-prediction not supported yet.")
|
||||
case PredictionType.UNKNOWN:
|
||||
# Fallback: treat as EPS (should not happen with enhanced detection)
|
||||
logging.warning(f"NRS.nrs: Unknown operation space, treating as EPS for conversion")
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_eps_space(x_orig, sig_root, sigma, cond, uncond)
|
||||
case _:
|
||||
raise RuntimeError(f"NRS.nrs: Invalid PredictionType used for operation space: {self.__OPERATION_SPACE}")
|
||||
|
||||
logging.debug(f"NRS.nrs: Skew: {skew}, Stretch: {stretch}, Squash: {squash}")
|
||||
def _dot(a, b):
|
||||
return (a*b).sum(dim=1, keepdim=True) # [B,C,W,H] => [B,1,W,H]
|
||||
|
||||
#rescale cfg has to be done on v-pred model output
|
||||
x = x_orig / (sigma * sigma + 1.0)
|
||||
cond = ((x - (x_orig - cond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma)
|
||||
uncond = ((x - (x_orig - uncond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma)
|
||||
logging.debug(f"NRS.nrs: generated cond and uncond")
|
||||
def _nrm2(v):
|
||||
return _dot(v, v)
|
||||
|
||||
x_final = None
|
||||
match "v0.4.5":
|
||||
case "v1":
|
||||
# displace cond by rejection of uncond on cond
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c = (u_dot_c / c_dot_c) * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
displaced = (cond - skew * u_rej_c)
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
eps = torch.finfo(nrs_cond.dtype).eps
|
||||
c_dot_c = _nrm2(nrs_cond) + eps # [B,1,W,H]
|
||||
u_dot_c = _dot(nrs_uncond, nrs_cond) # [B,1,W,H]
|
||||
u_on_c = (u_dot_c / c_dot_c) * nrs_cond # [B,1,W,H] * [B,C,H,W]
|
||||
|
||||
# Amplify Cond based on length compared to projection of uncond
|
||||
proj_diff = nrs_cond - u_on_c
|
||||
stretched = nrs_cond + (stretch * proj_diff)
|
||||
|
||||
# squash displaced vector towards len(cond) based on squash scale
|
||||
d_len_sq = torch.sum(displaced * displaced, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * ((c_dot_c/d_len_sq) ** 0.5)
|
||||
squashed = displaced * squash_scale
|
||||
logging.debug(f"NRS.nrs: squashed")
|
||||
# Skew/Steer Conf based on rejection of uncond on cond
|
||||
u_rej_c = nrs_uncond - u_on_c
|
||||
skewed = stretched - (skew * u_rej_c)
|
||||
|
||||
# stretch turned vector towards cond based on stretch scale
|
||||
sq_dot_c = torch.sum(squashed * cond, dim=-1, keepdim=True)
|
||||
sq_on_c = (sq_dot_c / c_dot_c) * cond
|
||||
x_final = squashed + sq_on_c * stretch
|
||||
logging.debug(f"NRS.nrs: final")
|
||||
case "v2":
|
||||
# displace cond by rejection of uncond on cond
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
displaced = cond + stretch * (cond - torch.clamp(u_dot_c / c_dot_c, min=0, max=1) * cond) - skew * u_rej_c
|
||||
logging.debug(f"NRS.nrs: displaced & stretched")
|
||||
# Squash final length back down to original length of cond
|
||||
cond_len = nrs_cond.norm(dim=1, keepdim=True)
|
||||
nrs_len = skewed.norm(dim=1, keepdim=True) + eps
|
||||
|
||||
# squash displaced vector towards len(cond) based on squash scale
|
||||
d_len_sq = torch.sum(displaced * displaced, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * ((c_dot_c/d_len_sq) ** 0.5)
|
||||
x_final = displaced * squash_scale
|
||||
logging.debug(f"NRS.nrs: final")
|
||||
case "v3":
|
||||
# displace cond by rejection of uncond on cond
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
displaced = (cond - skew * u_rej_c)
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
squash_scale = (1 - squash) + (squash * (cond_len / nrs_len))
|
||||
x_final = skewed * squash_scale
|
||||
|
||||
# squash displaced vector towards len(cond) based on squash scale
|
||||
d_len_sq = torch.sum(displaced * displaced, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * ((c_dot_c/d_len_sq) ** 0.5)
|
||||
|
||||
# stretch vector towards 2*len(cond) - len(u_on_c)
|
||||
c_len = c_dot_c ** 0.5
|
||||
stretch_scale = (1 - stretch) + stretch * (2 * c_len - u_on_c_mag)/c_len
|
||||
|
||||
x_final = displaced * squash_scale * stretch_scale
|
||||
logging.debug(f"NRS.nrs: final")
|
||||
case "v4":
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
rej_dor_rej = torch.sum(u_rej_c * u_rej_c, dim=-1, keepdim=True)
|
||||
x_final = (cond - squash * u_rej_c + stretch * cond * ((rej_dor_rej/c_dot_c) ** 0.5))
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
case "v0.4.1":
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
rej_dor_rej = torch.sum(u_rej_c * u_rej_c, dim=-1, keepdim=True)
|
||||
stretched = cond + stretch * cond * ((rej_dor_rej/c_dot_c) ** 0.5)
|
||||
skewed = stretched - skew * u_rej_c
|
||||
sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * ((c_dot_c/sk_dot_sk) ** 0.5)
|
||||
x_final = skewed * squash_scale
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
case "v0.4.2":
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
proj_len = torch.sum(u_on_c * u_on_c, dim=-1, keepdim=True) ** 0.5
|
||||
cond_len = c_dot_c ** 0.5
|
||||
stretched = cond * (1 + stretch * torch.abs(cond_len - proj_len) / cond_len)
|
||||
skewed = stretched - skew * u_rej_c
|
||||
sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5)
|
||||
x_final = skewed * squash_scale
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
case "v0.4.3":
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
proj_len = torch.sum(u_on_c * u_on_c, dim=-1, keepdim=True) ** 0.5
|
||||
cond_len = c_dot_c ** 0.5
|
||||
stretched = cond * (1 + stretch * (cond_len - proj_len) / cond_len)
|
||||
skewed = stretched - skew * u_rej_c
|
||||
sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5)
|
||||
x_final = skewed * squash_scale
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
case "v0.4.4":
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
cond_len = c_dot_c ** 0.5
|
||||
proj_diff = cond - u_on_c
|
||||
proj_diff_len = torch.sum(proj_diff * proj_diff, dim=-1, keepdim=True) ** 0.5
|
||||
stretched = cond * (1 + stretch * proj_diff_len / cond_len)
|
||||
skewed = stretched - skew * u_rej_c
|
||||
sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5)
|
||||
x_final = skewed * squash_scale
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
case "v0.4.5":
|
||||
u_dot_c = torch.sum(uncond * cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(cond * cond, dim=-1, keepdim=True)
|
||||
u_on_c_mag = (u_dot_c / c_dot_c)
|
||||
u_on_c = u_on_c_mag * cond
|
||||
u_rej_c = uncond - u_on_c
|
||||
cond_len = c_dot_c ** 0.5
|
||||
proj_diff = cond - u_on_c
|
||||
|
||||
# Amplify Cond based on length compared to projection of uncond
|
||||
stretched = cond + (stretch * proj_diff)
|
||||
|
||||
# Skew/Steer Conf based on rejection of uncond on cond
|
||||
skewed = stretched - skew * u_rej_c
|
||||
|
||||
# Squash final length back down to original length of cond
|
||||
sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True)
|
||||
squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5)
|
||||
x_final = skewed * squash_scale
|
||||
|
||||
return x_orig - (x - x_final * sigma / (sigma * sigma + 1.0) ** 0.5)
|
||||
match self.__OPERATION_SPACE:
|
||||
case PredictionType.V:
|
||||
return self._finalize_from_v_space(x_orig, x_div, x_final, sig_root, sigma)
|
||||
case PredictionType.EPS:
|
||||
return self._finalize_from_eps_space(x_orig, x_div, x_final, sig_root, sigma)
|
||||
case PredictionType.X0:
|
||||
raise RuntimeError("NRS.nrs: x0-prediction not supported yet.")
|
||||
case PredictionType.UNKNOWN:
|
||||
# Fallback: treat as EPS (should not happen with enhanced detection)
|
||||
logging.warning(f"NRS.nrs: Unknown operation space, treating as EPS for finalization")
|
||||
return self._finalize_from_eps_space(x_orig, x_div, x_final, sig_root, sigma)
|
||||
case _:
|
||||
raise RuntimeError(f"NRS.nrs: Invalid PredictionType used for operation space: {self.__OPERATION_SPACE}")
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_sampler_cfg_function(nrs, True)
|
||||
@@ -180,4 +286,4 @@ class NRS:
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"NRS": NRS,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
[](https://github.com/Reithan/negative_rejection_steering/actions/workflows/github-code-scanning/codeql)
|
||||
[](https://registry.comfy.org/nodes/negative_rejection_steering)
|
||||
|
||||
# Negative Rejection Steering
|
||||
NRS seeks to replace the 'naive' linear interpolation of Classifier Free Guidance with a more nuanced and composable steering of the generation process with better mathematical basis.
|
||||
|
||||
@@ -6,39 +9,115 @@ NRS seeks to replace the 'naive' linear interpolation of Classifier Free Guidanc
|
||||
2. NRS replaces CFG with 3 new knobs.
|
||||
3. NRS lets you to create cooler outputs than CFG.
|
||||
|
||||
> [!TIP]
|
||||
> Skip to the [Beginner How-To](#beginner-how-to) if you want to just get started.
|
||||
|
||||
### Math Demonstration
|
||||
<details>
|
||||
<summary>Expand for explanation of algorithm</summary>
|
||||
<img align="right" src="https://github.com/user-attachments/assets/01fabaff-8499-45f6-adad-d54b2c2fb7f1" alt="Graph of NRS vs CFG" style="width: 40%; float: right;">
|
||||
<img align="right" src="Examples/NRS_graph.png" alt="Graph of NRS vs CFG" style="width: 40%; float: right;">
|
||||
|
||||
### NRS is Applied in Three Steps:
|
||||
1. **Skewing**: The conditioned output tensor is skewed away from the direction of the rejection of the unconditioned tensor on the conditioned tensor. This lengthens the tensor in a direction perpendicular to its direction without affecting the positive guidance. The tensor is displaced by the rejection x the Displacement parameter.
|
||||
2. **Stretching**: The skewed tensor is stretched towards the direction of the original conditioned tensor based on its difference from the projection of uncond on cond. 1x stretch adds 100% of this difference to the tensor's length.
|
||||
3. **Squashing**: The skewed and stretched tensor is rescaled towards the original length of the conditioned tensor. 100% squashing outputs the original length of the conditioned tensor simple 'steered' towards the skew & squash output.
|
||||
0. ***V-Space**: Optional pre-NRS step* If the model is not using v-prediction, we transform the EPS `cond` and `uncond` into v-prediction space before continuing, then revert to eps-space before return.
|
||||
1. **Skewing**: The conditioned output tensor is skewed away from the direction of the rejection of the unconditioned tensor on the conditioned tensor. This lengthens the tensor in a direction perpendicular to its direction without affecting the positive guidance. The tensor is displaced by the rejection multiplied by the Skew parameter.[^1]
|
||||
2. **Stretching**: The skewed tensor is stretched towards the direction of the original conditioned tensor based on its difference from the projection of uncond on cond. The stretch is multiplied by the Stretch parameter.[^1]
|
||||
3. **Squashing**: The skewed and stretched tensor is rescaled towards the original length of the conditioned tensor. 100% squashing outputs the original length of the conditioned tensor simply 'steered' towards the skewed & squashed version's direction.[^1]
|
||||
[^1]: All operations are done per feature across the step's batch, width, and height.
|
||||
|
||||
[Interactive Graph on Math3D.org](https://www.math3d.org/aTJW4UZtCh)
|
||||
</details>
|
||||
|
||||
## Parameters
|
||||
Skew and Stretch are roughly similar to CFG, but decomposed, with `Stretch + 2 * Skew = 2 * CFG`, roughly.
|
||||
Meaning, if you want to 'replicate' a simliar effect for a given CFG setting, you should set Skew equal to CFG, and Stretch to 1/2 CFG.
|
||||
Squash should initially be set to 0%, then adjusted based on 'burn' of output.
|
||||
## Examples of NRS Effects
|
||||
**Skew**
|
||||

|
||||
**Stretch**
|
||||

|
||||
**Squash**
|
||||

|
||||
<details>
|
||||
<summary><small>Generation details for reproduction</small></summary>
|
||||
|
||||
- **Skew** changes the 'direction' of generation, which should result in changes to the content and composition of the image.
|
||||
- **Stretch** changes to 'amplification' of generation, which should result in stronger prompt representation.
|
||||
- **Squash** 'normalizes' the resulting guidance back towards the original amplitude. This results in a removal of 'burn-in' and artifacting of the output, transforming these defects into alternative guidance.
|
||||
| Prompt | |
|
||||
| ---------- | --- |
|
||||
| Tool | [Stable Diffusion WebUI reForge](https://github.com/Panchovix/stable-diffusion-webui-reForge) |
|
||||
| Sampler | DPM++ 2M |
|
||||
| Scheduler | Align Your Steps |
|
||||
| Steps | 25 |
|
||||
| Dimensions | 912 x 624 |
|
||||
| Seed | `1334103348` |
|
||||
| Model | [Lobotomized Mix v1.5](https://civitai.com/models/1144932) |
|
||||
| Embeddings | [Lazy Embeddings for ALL illustrious NoobAI...](https://civitai.com/models/1302719), [Smooth Embeddings](https://civitai.com/models/1065154) |
|
||||
| Positive | lazypos, [Smooth_Quality\|SmoothNoob_Quality], BREAK<br>very awa, masterpiece, best quality, year 2024, newest, highres, absurdres,<br>1girl, samurai archer, cyberpunk cityscape, rain-soaked rooftop, neon reflection puddles, volumetric mist,<br>photorealistic, digital art,<br>dramatic rim lighting, shallow depth of field, low angle viewpoint |
|
||||
| Negative | lazyloli, lazynsfw, BREAK<br>lazyhand, SmoothNegative_Hands-neg, BREAK<br>[Smooth_Negative-neg\|SmoothNoob_Negative-neg], BREAK<br>lowres, worst quality, worst aesthetic, bad quality, jpeg artifacts, scan artifacts,<br>blurry, deformed anatomy, bad hands, extra fingers, missing fingers, mutated hands,<br>watermark, logo, text, nsfw |
|
||||
</details>
|
||||
|
||||
### Explanation of Effects
|
||||
#### Skew
|
||||
**Skew** changes the 'direction' of your generation, altering the image generation to 'steer' away from negative prompt elements as they conflict with your positive prompt. Increasing Skew will change scene composition, geometry, and scene elements to ensure that the final image aligns with the intention of your prompt pair.
|
||||
#### Stretch
|
||||
**Stretch** changes the intensity of generated elements that align more with your positive prompt than the negative. This 'hits the gas' on any elements that are more strongly aligned with your positive prompt than your negative, and 'hit the brakes' on the opposite.
|
||||
#### Squash
|
||||
**Squash** is the speed limit. At 0.0 Squash, each diffusion step receives the full intensity you set from Skew and Stretch, while 1.0 Squash ensures each step has only the original step size output by the model. This setting will only remove intensity unless you have a non-zero Skew value. Squash will 'soften' the effects of Skew and Stretch as it's raised, but the 'removed' Skew and Stretch intensity is replaced by enhanced micro-detailing and 'burn'. Squash should generally be left low and used as a 'finishing' step after dialing in a decent Skew and Stretch value.
|
||||
|
||||
## Beginner How-To
|
||||
1. Set Skew to your normal CFG Scale setting and Stretch to 1/2 your normal CFG Scale.
|
||||
2. Set Squash to 0.0.
|
||||
3. Test some outputs. Results should be similar in quality to CFG.
|
||||
4. Adjust Skew up/down to change content and composition.
|
||||
5. Adjust Stretch up/down to change strength of positive prompt aspects and colors.
|
||||
6. Adjust Squash up to remove artifacts and color burn (these will tend to be replaced by additional details and elements).
|
||||
1. Set Skew to 1/2 of your normal CFG Scale setting and Stretch to your full normal CFG Scale. Set Squash to 0.0.<br>
|
||||
*Alternatively, try starting with the default of 2/5/0.75, or at 1/1/1 to get a baseline.*
|
||||
2. Test some outputs. Results should be similar in quality to CFG.
|
||||
3. Adjust Skew to change the intensity of your outputs adherence to your positive and negative prompts. This primarily effects composition of the output.
|
||||
4. Adjust Stretch to intensify your positive prompt's aspects and colors where they differ from the negative prompt. This primarily effects color and texture.
|
||||
5. Adjust Squash to soften Skew and Stretch's effects. The intensity removed from Skew and Stretch will generally become additional micro-detailing and elements.
|
||||
|
||||
**Tip**: You can experiment with negative values for Skew and Stretch as well, to see how the model is interpeting your negative prompt.
|
||||
> [!TIP]
|
||||
> You can experiment with negative values for each setting as well. This can be useful to understand how the model interpreting your negative prompt.
|
||||
|
||||
## Examples
|
||||
> [!WARNING]
|
||||
> Don't set NRS values to negatives if there are things in your negative prompt you **actually** don't want to see.
|
||||
|
||||
## Setup & Installation
|
||||
|
||||
### ComfyUI
|
||||
<details>
|
||||
<summary>ComfyUI Setup Instructions</summary>
|
||||
|
||||
#### Installation
|
||||
Install via ComfyUI Manager or manually clone this repository into your `ComfyUI/custom_nodes/` directory.
|
||||
|
||||
#### Usage
|
||||
1. **Important**: Ignore the CFG setting on your KSampler node - NRS replaces CFG entirely
|
||||
2. Connect your model through the **Negative Rejection Steering** node before sampling
|
||||
3. Configure NRS parameters (Skew/Stretch/Squash) instead of using CFG
|
||||
|
||||
#### Basic Workflow
|
||||
```
|
||||
Model → NRS Node → KSampler
|
||||
```
|
||||
|
||||
**Pro tip**: To verify NRS is working correctly, set CFG to an extremely high value (like 30). If your output looks normal, NRS is functioning properly. If the output appears "turbo fried," check your node connections.
|
||||
|
||||
**Sampler Compatibility**: NRS now supports advanced samplers including WanKSamplerAdvanced, RES4LYF samplers, and FLOW models (WAN21, Flux) with enhanced prediction type detection.
|
||||
|
||||

|
||||
</details>
|
||||
|
||||
### Automatic1111 / Forge / reForge
|
||||
<details>
|
||||
<summary>WebUI Setup Instructions</summary>
|
||||
|
||||
#### Installation
|
||||
1. Install the extension through the Extensions tab in your WebUI
|
||||
2. Enable the extension and restart your WebUI
|
||||
|
||||
#### Usage
|
||||
Once installed and enabled, the NRS settings panel will appear in your generation interface. When NRS is active:
|
||||
- **CFG Scale is ignored** - the WebUI may still show the CFG setting, but it has no effect
|
||||
- Use the NRS parameters (Skew/Stretch/Squash) to control generation instead
|
||||
- Follow the same parameter guidelines from the [Beginner How-To](#beginner-how-to) section
|
||||
</details>
|
||||
|
||||
### StabilityMatrix Integration
|
||||
NRS is available as a **natively supported module** in [StabilityMatrix](https://lykos.ai/), providing an easy installation and management option for users of that platform.
|
||||
|
||||
## Submitted User Examples
|
||||
| User | CFG | NRS |
|
||||
|---|---|---|
|
||||
| --- | --- | --- |
|
||||
| Mohnjiles from StabilityMatrix |  |  |
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from .NRS.nodes_NRS import *
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"NRS": NRS}
|
||||
NODE_DISPLAY_NAME_MAPPINS = {"NRS": "Negative Rejection Steering"}
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
@@ -0,0 +1,11 @@
|
||||
# The 'purpose' or 'intention' behind our 3 knobs
|
||||
|
||||
## SKEW
|
||||
This is the primary 'steering' knob. This will 'turn' the 'direction' the current denoising step is traveling in the latent space. If we define the 'default' (cond) direction as 'prior step -> cond' then we're just applying a 'lateral' skew to that direction to 'turn' it 'away' from the 'unintended' direction (uncond)
|
||||
|
||||
## STRETCH
|
||||
This is the sister knob to Skew. This is the accelerator. We want to go 'faster' into the intended direction (cond) the less aligned it is with the unintended direction (uncond). Think of this like a combination of brakes + gas. If we're headed directly for a brick wall (uncond is in the same direction as cond), we want to apply no acceleration, or negative acceleration. If we're traveling directly away from danger (uncond is in the opposite direction of cond) then we want to stomp the gas and get as far away as we can. There's only 1 problem with this BASIC-level description: as we get further into generation, regardless of pos/neg promp, cond & uncon will naturally align to be the same vector[^1]. In the last stop of inference, cond and uncond will be basically identical if nothing has fucked up. So whatever math we apply here needs to take the progressive alignment of cond & uncond into account. That's why were/are scaling only on the projection difference right now, rather than the full projection.
|
||||
[^1]: This is more true in eps than v-pred. Stretch is inherently more powerful in v-pred based models than eps models.
|
||||
|
||||
## SQUASH
|
||||
This is out 'safety' knob. Think of this like a 'limiter' in a car. This sets the 'top speed' we can go to some multiple of the 'default' speed the model would 'like to' go. i.e. whatever length of directional vector the model produces prior to any skewing or stretching is treated as the 'default' length with Squash=1.0 ensuring we only every go that 'speed' and no more, while Squash=0.0 lets us go any speed we want based on the other 2 knobs. GENERALLY we'll be leaving Squash at 0.0 unless we need it for specific generations.
|
||||
@@ -0,0 +1,16 @@
|
||||
[project]
|
||||
name = "negative_rejection_steering"
|
||||
description = "NRS seeks to replace the 'naive' linear interpolation of Classifier Free Guidance with a more nuanced and composable steering of the generation process with better mathematical basis."
|
||||
authors = [{name = "Bryan O'Malley", email = "bo122081@hotmail.com"}]
|
||||
version = "0.7.4"
|
||||
license = {file = "LICENSE"}
|
||||
readme = "README.md"
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/Reithan/negative_rejection_steering"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "reithan"
|
||||
DisplayName = "Negative Rejection Steering"
|
||||
Icon = "https://raw.githubusercontent.com/Reithan/negative_rejection_steering/main/icon.png"
|
||||
@@ -11,14 +11,14 @@ class NRSScript(scripts.Script):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.enabled = False
|
||||
self.skew = 2.0
|
||||
self.stretch = 2.0
|
||||
self.squash = 1.0
|
||||
self.skew = 2.00
|
||||
self.stretch = 5.00
|
||||
self.squash = 0.75
|
||||
|
||||
sorting_priority = 5
|
||||
|
||||
def title(self):
|
||||
return "Negative Rejection Steering for reForge"
|
||||
return "Negative Rejection Steering"
|
||||
|
||||
def show(self, is_img2img):
|
||||
return scripts.AlwaysVisible
|
||||
|
||||
Reference in New Issue
Block a user