From 78dde5adaad4987ba91ec56467906452feaed4c1 Mon Sep 17 00:00:00 2001 From: Reithan Date: Fri, 21 Mar 2025 19:58:30 -0700 Subject: [PATCH] Alpha implementation --- NRS/nodes_NRS.py | 173 ++++++++++++++---- scripts/negative_rejection_steering_script.py | 39 ++-- 2 files changed, 165 insertions(+), 47 deletions(-) diff --git a/NRS/nodes_NRS.py b/NRS/nodes_NRS.py index adc17f4..6206480 100644 --- a/NRS/nodes_NRS.py +++ b/NRS/nodes_NRS.py @@ -14,16 +14,15 @@ class NRS: CATEGORY = "advanced/model" - def patch(self, model, squash, stretch): + def patch(self, model, skew, stretch, squash): def nrs(args): cond = args["cond"] uncond = args["uncond"] - cond_scale = args["cond_scale"] sigma = args["sigma"] sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1)) x_orig = args["input"] - logging.debug(f"NRS.nrs: CFG: {cond_scale}, Squash: {squash}, Stretch: {stretch}") + 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) @@ -32,40 +31,146 @@ class NRS: logging.debug(f"NRS.nrs: generated cond and uncond") x_final = None - if False: - # 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 - cond_scale * u_rej_c) - logging.debug(f"NRS.nrs: displaced") + 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") - # squash displaced vector towards len(cond) based on squash scale - sq_len = torch.sum(displaced * displaced, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * ((c_dot_c/sq_len) ** 0.5) - squashed = displaced * squash_scale - logging.debug(f"NRS.nrs: squashed") + # 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") - # 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") - else: - # 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 + stretch * (cond - torch.clamp(u_dot_c / c_dot_c, min=0, max=1) * cond) - cond_scale * u_rej_c - logging.debug(f"NRS.nrs: displaced & stretched") + # 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 displaced vector towards len(cond) based on squash scale - sq_len = torch.sum(displaced * displaced, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * ((c_dot_c/sq_len) ** 0.5) - x_final = displaced * squash_scale - logging.debug(f"NRS.nrs: final") + # 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 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) diff --git a/scripts/negative_rejection_steering_script.py b/scripts/negative_rejection_steering_script.py index f245341..9bbe8b1 100644 --- a/scripts/negative_rejection_steering_script.py +++ b/scripts/negative_rejection_steering_script.py @@ -11,8 +11,9 @@ class NRSScript(scripts.Script): def __init__(self): super().__init__() self.enabled = False - self.squash = 0.5 - self.stretch = 1.0 + self.skew = 2.0 + self.stretch = 2.0 + self.squash = 1.0 sorting_priority = 5 @@ -26,23 +27,27 @@ class NRSScript(scripts.Script): with gr.Accordion(open=False, label=self.title()): gr.HTML("

Adjust the settings for Negative Rejection Steering.

") enabled = gr.Checkbox(label="Enable NRS", value=self.enabled) + gr.HTML("

Adjust the amount guidance is steered.

") + skew = gr.Slider(label="NRS Skew Scale", minimum=-30.0, maximum=30.0, step=0.01, value=self.skew) + gr.HTML("

Adjust the amount guidance is amplified.

") + stretch = gr.Slider(label="NRS Stretch Scale", minimum=-30.0, maximum=30.0, step=0.01, value=self.stretch) + gr.HTML("

Adjust the amount final guidance is normalized.

") squash = gr.Slider(label="NRS Squash Multiplier", minimum=0.0, maximum=1.0, step=0.01, value=self.squash) - stretch = gr.Slider(label="NRS Stretch Multiplier", minimum=-1.0, maximum=30.0, step=0.01, value=self.stretch) enabled.change( lambda x: self.update_enabled(x), inputs=[enabled] ) - return (enabled, squash, stretch) + return (enabled, skew, stretch, squash) def update_enabled(self, value): self.enabled = value def process_before_every_sampling(self, p, *args, **kwargs): - if len(args) >= 3: - self.enabled, self.squash, self.stretch = args[:3] + if len(args) >= 4: + self.enabled, self.skew, self.stretch, self.squash = args[:4] else: logging.warning("Not enough arguments provided to process_before_every_sampling") return @@ -50,10 +55,12 @@ class NRSScript(scripts.Script): xyz = getattr(p, "_nrs_xyz", {}) if "enabled" in xyz: self.enabled = xyz["enabled"] == "True" - if "squash" in xyz: - self.squash = xyz["squash"] + if "skew" in xyz: + self.skew = xyz["skew"] if "stretch" in xyz: self.stretch = xyz["stretch"] + if "squash" in xyz: + self.squash = xyz["squash"] # Always start with a fresh clone of the original unet unet = p.sd_model.forge_objects.unet.clone() @@ -63,16 +70,17 @@ class NRSScript(scripts.Script): p.sd_model.forge_objects.unet = unet return - unet = NRS().patch(unet, self.squash, self.stretch)[0] + unet = NRS().patch(unet, self.skew, self.stretch, self.squash)[0] p.sd_model.forge_objects.unet = unet p.extra_generation_params.update({ "NRS_enabled": True, - "NRS_squash": self.squash, + "NRS_skew": self.skew, "NRS_stretch": self.stretch, + "NRS_squash": self.squash, }) - logging.debug(f"NRS: Enabled: {self.enabled}, Squash: {self.squash}, Stretch: {self.stretch}") + logging.debug(f"NRS: Enabled: {self.enabled}, Squash: {self.skew}, Stretch: {self.stretch}, Squash: {self.squash}") return @@ -99,15 +107,20 @@ def make_axis_on_xyz_grid(): choices=lambda: ["True", "False"] ), xyz_grid.AxisOption( - "(NRS) Squash", + "(NRS) Skew", float, - partial(set_value, field="squash"), + partial(set_value, field="skew"), ), xyz_grid.AxisOption( "(NRS) Stretch", float, partial(set_value, field="stretch"), ), + xyz_grid.AxisOption( + "(NRS) Squash", + float, + partial(set_value, field="squash"), + ), ] if not any(x.label.startswith("(NRS)") for x in xyz_grid.axis_options):