Alpha implementation
This commit is contained in:
+139
-34
@@ -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)
|
||||
|
||||
|
||||
@@ -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("<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)
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user