Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
50ddd2ac47 | ||
|
|
3768f4a768 | ||
|
|
e82622cefe | ||
|
|
793914ced3 | ||
|
|
46540aa7bc | ||
|
|
873034b095 | ||
|
|
7d399643dd | ||
|
|
50033c2622 | ||
|
|
36eb9d592b | ||
|
|
bc50983954 | ||
|
|
0c37c6b124 | ||
|
|
8d7d9281f7 | ||
|
|
e55881afe2 | ||
|
|
0e8e508213 | ||
|
|
ff69ac386d | ||
|
|
5be3c4c4f1 | ||
|
|
81b836d2f1 | ||
|
|
a57c16a624 | ||
|
|
e6cdf03189 | ||
|
|
47e4b5073b | ||
|
|
0778fa4b03 |
Binary file not shown.
|
After Width: | Height: | Size: 2.3 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 2.2 MiB |
+194
-60
@@ -1,45 +1,152 @@
|
||||
import ldm_patched.modules.model_base
|
||||
import inspect
|
||||
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,
|
||||
"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}),
|
||||
"skew": ("FLOAT", {"default": 4.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}),
|
||||
"squash": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "advanced/model"
|
||||
|
||||
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 = [model]
|
||||
visited = set()
|
||||
|
||||
while queue:
|
||||
obj = queue.pop(0)
|
||||
if id(obj) in visited:
|
||||
continue
|
||||
visited.add(id(obj))
|
||||
|
||||
# 1) direct hit on this object ---------------------------------
|
||||
p = _canon(getattr(obj, "parameterization", None)) # k-diffusion
|
||||
if p:
|
||||
return _RAW_TO_ENUM.get(p, PredictionType.UNKNOWN)
|
||||
|
||||
p = _canon(getattr(getattr(obj, "config", None), "prediction_type", None)) # diffusers
|
||||
if p:
|
||||
return _RAW_TO_ENUM.get(p, PredictionType.UNKNOWN)
|
||||
|
||||
p = _canon(getattr(obj, "prediction_type", None)) # rare misc
|
||||
if p:
|
||||
return _RAW_TO_ENUM.get(p, PredictionType.UNKNOWN)
|
||||
|
||||
# 2) enqueue child containers we care about -------------------
|
||||
for attr in ("model_sampling", "model", "diffusion_model", "scheduler"):
|
||||
child = getattr(obj, attr, None)
|
||||
if child is not None:
|
||||
queue.append(child)
|
||||
|
||||
# 3) default ------------------------------------------------------
|
||||
return PredictionType.EPS
|
||||
|
||||
|
||||
def _pre_scale_conditioning(self, x_orig, sigma, cond, uncond):
|
||||
x_div = None
|
||||
eps_cond = cond
|
||||
eps_uncond = uncond
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# v → ε conversion
|
||||
logging.debug(f"NRS._pre_scale_conditioning: generating x_div, cond, and uncond for v-pred")
|
||||
sigma2_1 = (sigma ** 2 + 1.0)
|
||||
x_div = x_orig / sigma2_1
|
||||
root = sigma2_1.sqrt()
|
||||
|
||||
eps_cond = ((x_div - (x_orig - cond)) * root) / (sigma)
|
||||
eps_uncond = ((x_div - (x_orig - uncond)) * root) / (sigma)
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
logging.debug(f"NRS._pre_scale_conditioning: already in eps, no pre-scale needed")
|
||||
pass # already in ε space
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._pre_scale_conditioning: x0-prediction not supported yet.")
|
||||
else:
|
||||
raise RuntimeError("NRS._pre_scale_conditioning: Could not determine prediction type for this model.")
|
||||
|
||||
return x_div, eps_cond, eps_uncond
|
||||
|
||||
def _post_scale_conditioning(self, x_orig, x_div, x_final, sigma):
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# ε → v conversion
|
||||
root = (sigma ** 2 + 1).sqrt()
|
||||
|
||||
logging.debug(f"NRS._post_scale_conditioning: generating cfg_result for v-pred")
|
||||
return x_orig - (x_div - x_final * sigma / root)
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
# already in ε space
|
||||
logging.debug(f"NRS._post_scale_conditioning: already in eps, no post-scale needed")
|
||||
return x_final
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._post_scale_conditioning: x0-prediction not supported yet.")
|
||||
else:
|
||||
raise RuntimeError("NRS._post_scale_conditioning: Could not determine prediction type for this model.")
|
||||
|
||||
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
|
||||
|
||||
def nrs(args):
|
||||
cond = args["cond"]
|
||||
uncond = args["uncond"]
|
||||
sigma = args["sigma"]
|
||||
sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1))
|
||||
x_orig = args["input"]
|
||||
self.__pred_type = self.__pred_type if self.__pred_type is not None else self._get_pred_type(model)
|
||||
|
||||
logging.debug(f"NRS.nrs: Skew: {skew}, Stretch: {stretch}, Squash: {squash}")
|
||||
|
||||
#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")
|
||||
|
||||
x_div, eps_cond, eps_uncond = self._pre_scale_conditioning(x_orig, sigma, cond, uncond)
|
||||
|
||||
x_final = None
|
||||
match "v0.4.5":
|
||||
match "v0.5.0":
|
||||
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)
|
||||
u_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_cond, dim=-1, keepdim=True)
|
||||
u_on_c = (u_dot_c / c_dot_c) * eps_cond
|
||||
u_rej_c = eps_uncond - u_on_c
|
||||
displaced = (eps_cond - skew * u_rej_c)
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
|
||||
# squash displaced vector towards len(cond) based on squash scale
|
||||
@@ -49,18 +156,18 @@ class NRS:
|
||||
logging.debug(f"NRS.nrs: squashed")
|
||||
|
||||
# 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
|
||||
sq_dot_c = torch.sum(squashed * eps_cond, dim=-1, keepdim=True)
|
||||
sq_on_c = (sq_dot_c / c_dot_c) * eps_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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_uncond - u_on_c
|
||||
displaced = eps_cond + stretch * (eps_cond - torch.clamp(u_dot_c / c_dot_c, min=0, max=1) * eps_cond) - skew * u_rej_c
|
||||
logging.debug(f"NRS.nrs: displaced & stretched")
|
||||
|
||||
# squash displaced vector towards len(cond) based on squash scale
|
||||
@@ -70,12 +177,12 @@ class NRS:
|
||||
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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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)
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_uncond - u_on_c
|
||||
displaced = (eps_cond - skew * u_rej_c)
|
||||
logging.debug(f"NRS.nrs: displaced")
|
||||
|
||||
# squash displaced vector towards len(cond) based on squash scale
|
||||
@@ -89,81 +196,81 @@ class NRS:
|
||||
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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_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))
|
||||
x_final = (eps_cond - squash * u_rej_c + stretch * eps_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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_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)
|
||||
stretched = eps_cond + stretch * eps_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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_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)
|
||||
stretched = eps_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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_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)
|
||||
stretched = eps_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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_uncond - u_on_c
|
||||
cond_len = c_dot_c ** 0.5
|
||||
proj_diff = cond - u_on_c
|
||||
proj_diff = eps_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)
|
||||
stretched = eps_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_dot_c = torch.sum(eps_uncond * eps_cond, dim=-1, keepdim=True)
|
||||
c_dot_c = torch.sum(eps_cond * eps_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
|
||||
u_on_c = u_on_c_mag * eps_cond
|
||||
u_rej_c = eps_uncond - u_on_c
|
||||
cond_len = c_dot_c ** 0.5
|
||||
proj_diff = cond - u_on_c
|
||||
proj_diff = eps_cond - u_on_c
|
||||
|
||||
# Amplify Cond based on length compared to projection of uncond
|
||||
stretched = cond + (stretch * proj_diff)
|
||||
stretched = eps_cond + (stretch * proj_diff)
|
||||
|
||||
# Skew/Steer Conf based on rejection of uncond on cond
|
||||
skewed = stretched - skew * u_rej_c
|
||||
@@ -172,8 +279,35 @@ class NRS:
|
||||
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
|
||||
case "v0.5.0":
|
||||
def _dot(a, b):
|
||||
return (a*b).flatten(1).sum(dim=1, keepdim=True) # [B,1]
|
||||
def _nrm2(v):
|
||||
return _dot(v, v)
|
||||
|
||||
eps = eps = torch.finfo(eps_cond.dtype).eps
|
||||
c_dot_c = _nrm2(eps_cond) + eps # [B,1]
|
||||
u_dot_c = _dot(eps_uncond, eps_cond) # [B,1]
|
||||
|
||||
return x_orig - (x - x_final * sigma / (sigma * sigma + 1.0) ** 0.5)
|
||||
u_on_c = (u_dot_c / c_dot_c).unsqueeze(-1) * eps_cond # [B,1,1,1] * [B,C,H,W]
|
||||
u_rej_c = eps_uncond - u_on_c
|
||||
proj_diff = eps_cond - u_on_c
|
||||
|
||||
|
||||
# Amplify Cond based on length compared to projection of uncond
|
||||
stretched = eps_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
|
||||
cond_len = torch.sqrt(c_dot_c) # [B,1]
|
||||
nrs_len = torch.sqrt(_nrm2(skewed) + eps) # [B,1]
|
||||
|
||||
squash_scale = (1 - squash) + squash * (cond_len / nrs_len)
|
||||
x_final = skewed * squash_scale.unsqueeze(-1)
|
||||
|
||||
return self._post_scale_conditioning(x_orig, x_div, x_final, sigma)
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_sampler_cfg_function(nrs, True)
|
||||
|
||||
@@ -1,27 +1,44 @@
|
||||
# Negative Rejection Steering
|
||||
NRS seeks to replace the 'naive' linear interpolation of Classifier Free Guidance with a more nuanced steering of the generation process.
|
||||
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.
|
||||
|
||||
This is accomplised in 3 steps:
|
||||
1. **Displacement**: The conditioned output tensor is displaced in the direction of the rejection of the unconditioned tensor on the conditioned tensor. This lengthens the tensor in a direction perpendicular to it's direction without affecting the positive guidance. The tensor is displaced by the rejection x the Displacement parameter.
|
||||
2. **Squashing**: The displaced tensor is rescaled towards the original length of the conditioned tensor. This means for high displacement scaling values the tensor 'turns' away from the unconditioned direction, which for very negative displacements, it turns towards the unconditioned tensor. 0 displacement outputs the original conditioned tensor.
|
||||
3. **Stretching**: The post-squash 'steered' tensor is stretched towards the direction of the original conditioned tensor. The more sharp the steering the less pronounced the stretch is, with fully aligned tensors being stretched the full stretch scale parameter. 1x stretch adds 100% length to the tensor.
|
||||
#### _**TL;DR**_:
|
||||
1. CFG is a bad 'knob'
|
||||
2. NRS replaces CFG with 3 new knobs.
|
||||
3. NRS lets you to create cooler outputs than CFG.
|
||||
|
||||
# Alpha Release
|
||||
Implements NRS with Skew, Stretch, and Squash parameters.
|
||||
### 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;">
|
||||
|
||||
### 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 multiplied by the Skew 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. The stretch is multiplied by the Stretch parameter.
|
||||
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.
|
||||
|
||||
[Interactive Graph on Math3D.org](https://www.math3d.org/aTJW4UZtCh)
|
||||
</details>
|
||||
|
||||
## Parameters
|
||||
Skew and Stretch are roughly similar to CFG, but decomposed, with `Stretch + Skew = CFG`, roughly.
|
||||
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.
|
||||
|
||||
**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 with 1.0 being the same amplitude, while 0.0 is the unmodified amplitude resulting from the Squash and Stretch functions.
|
||||
- **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.
|
||||
|
||||
## Beginner How-To
|
||||
1. Set Squash to 0.0
|
||||
2. Set Skew & Stretch each to 1/2 your normal CFG Scale setting
|
||||
3. Test some generation. 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 image aspects and colors
|
||||
6. Adjust Squash up to remove artifacts and color burn (these will tend to be replaced by additional or extraneous details and elements)
|
||||
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).
|
||||
|
||||
**Tip**: You can experiment with negative values for Skew and Stretch as well to see what the model 'believes' your negative prompt 'means'.
|
||||
**Tip**: You can experiment with negative values for Skew and Stretch as well, to see how the model is interpeting your negative prompt.
|
||||
|
||||
## 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']
|
||||
@@ -11,28 +11,25 @@ class NRSScript(scripts.Script):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.enabled = False
|
||||
self.skew = 2.0
|
||||
self.skew = 4.0
|
||||
self.stretch = 2.0
|
||||
self.squash = 1.0
|
||||
self.squash = 0.0
|
||||
|
||||
sorting_priority = 5
|
||||
|
||||
def title(self):
|
||||
return "Negative Rejection Steering for reForge"
|
||||
return "Negative Rejection Steering"
|
||||
|
||||
def show(self, is_img2img):
|
||||
return scripts.AlwaysVisible
|
||||
|
||||
def ui(self, *args, **kwargs):
|
||||
with gr.Accordion(open=False, label=self.title()):
|
||||
gr.HTML("<p><i>Adjust the settings for Negative Rejection Steering.</i></p>")
|
||||
enabled = gr.Checkbox(label="Enable NRS", value=self.enabled)
|
||||
gr.HTML("<p><i>Adjust the amount guidance is steered.</i></p>")
|
||||
skew = gr.Slider(label="NRS Skew Scale", minimum=-30.0, maximum=30.0, step=0.01, value=self.skew)
|
||||
gr.HTML("<p><i>Adjust the amount guidance is amplified.</i></p>")
|
||||
stretch = gr.Slider(label="NRS Stretch Scale", minimum=-30.0, maximum=30.0, step=0.01, value=self.stretch)
|
||||
gr.HTML("<p><i>Adjust the amount final guidance is normalized.</i></p>")
|
||||
squash = gr.Slider(label="NRS Squash Multiplier", minimum=0.0, maximum=1.0, step=0.01, value=self.squash)
|
||||
gr.HTML("<p><i>Adjust the settings for Negative Rejection Steering.</i></p>")
|
||||
skew = gr.Slider(label="NRS Skew Scale", info="Adjusts the amount guidance is steered.", minimum=-30.0, maximum=30.0, step=0.01, value=self.skew)
|
||||
stretch = gr.Slider(label="NRS Stretch Scale", info="Adjusts the amount guidance is amplified.", minimum=-30.0, maximum=30.0, step=0.01, value=self.stretch)
|
||||
squash = gr.Slider(label="NRS Squash Multiplier", info="Adjusts the amount final guidance is normalized.", minimum=0.0, maximum=1.0, step=0.01, value=self.squash)
|
||||
|
||||
enabled.change(
|
||||
lambda x: self.update_enabled(x),
|
||||
|
||||
Reference in New Issue
Block a user