From 50ddd2ac47e780ff40d5452fe6c3a5560d56dbea Mon Sep 17 00:00:00 2001 From: Reithan Date: Sat, 19 Jul 2025 17:59:27 -0700 Subject: [PATCH] Update math to 0.5.0 (#10) - [X] detects v-pred/eps and uses appropriate pre/post scaling - [X] supports detection in Forge, Comfy and various loaders/models --- NRS/nodes_NRS.py | 249 ++++++++++++++---- scripts/negative_rejection_steering_script.py | 2 +- 2 files changed, 193 insertions(+), 58 deletions(-) diff --git a/NRS/nodes_NRS.py b/NRS/nodes_NRS.py index 5315fab..885057d 100644 --- a/NRS/nodes_NRS.py +++ b/NRS/nodes_NRS.py @@ -1,5 +1,25 @@ +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 @@ -14,31 +34,119 @@ class NRS: 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 @@ -48,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 @@ -69,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 @@ -88,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 @@ -171,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) diff --git a/scripts/negative_rejection_steering_script.py b/scripts/negative_rejection_steering_script.py index 3b19beb..b566083 100644 --- a/scripts/negative_rejection_steering_script.py +++ b/scripts/negative_rejection_steering_script.py @@ -18,7 +18,7 @@ class NRSScript(scripts.Script): sorting_priority = 5 def title(self): - return "Negative Rejection Steering for reForge" + return "Negative Rejection Steering" def show(self, is_img2img): return scripts.AlwaysVisible