Implement extra Rescale Classifier-Free Guidance features and UI options (#55)

* Implement RCFG

Source: https://github.com/ashen-sensored/sd-dynamic-thresholding-rcfg

* Update scripts

Synchronise changes from master branch

* Fixes

* Fixes

* Fixes

* Missing UI options

* cleanup some PR complexity

* small bit more cleaning

* cleanup and reformat and fix

arg position is sensitive in API, so add to end. Mark 'clean path' separate from 'special path', reformat UI a bit

---------

Co-authored-by: Alex "mcmonkey" Goodwin <git_commits@alexgoodwin.dev>
This commit is contained in:
Convexity-ai
2023-07-29 20:32:27 -07:00
committed by GitHub
co-authored by Alex mcmonkey Goodwin
parent 27700fddf8
commit 96739fab6e
2 changed files with 68 additions and 16 deletions
+40 -11
View File
@@ -3,7 +3,7 @@ import torch, math
######################### DynThresh Core #########################
class DynThresh:
def __init__(self, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, experiment_mode, maxSteps):
def __init__(self, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, experiment_mode, maxSteps, separate_feature_channels, scaling_startpoint, variability_measure, interpolate_phi):
self.mimic_scale = mimic_scale
self.threshold_percentile = threshold_percentile
self.mimic_mode = mimic_mode
@@ -13,6 +13,10 @@ class DynThresh:
self.mimic_scale_min = mimic_scale_min
self.experiment_mode = experiment_mode
self.sched_val = sched_val
self.sep_feat_channels = separate_feature_channels
self.scaling_startpoint = scaling_startpoint
self.variability_measure = variability_measure
self.interpolate_phi = interpolate_phi
def interpretScale(self, scale, mode, min):
scale -= min
@@ -72,20 +76,45 @@ class DynThresh:
mim_centered = mim_flattened - mim_means
cfg_centered = cfg_flattened - cfg_means
### Get the maximum value of all datapoints (with an optional threshold percentile on the uncond)
mim_max = mim_centered.abs().max(dim=2).values.unsqueeze(2)
cfg_max = torch.quantile(cfg_centered.abs(), self.threshold_percentile, dim=2).unsqueeze(2)
actualMax = torch.maximum(cfg_max, mim_max)
if self.sep_feat_channels:
if self.variability_measure == 'STD':
min_scaleref = mim_centered.std(dim=2).unsqueeze(2)
cfg_scaleref = cfg_centered.std(dim=2).unsqueeze(2)
else: # 'AD'
min_scaleref = mim_centered.abs().max(dim=2).values.unsqueeze(2)
cfg_scaleref = torch.quantile(cfg_centered.abs(), self.threshold_percentile, dim=2).unsqueeze(2)
### Clamp to the max
cfg_clamped = cfg_centered.clamp(-actualMax, actualMax)
### Now shrink from the max to normalize and grow to the mimic scale (instead of the CFG scale)
cfg_renormalized = (cfg_clamped / actualMax) * mim_max
else:
if self.variability_measure == 'STD':
min_scaleref = mim_centered.std()
cfg_scaleref = cfg_centered.std()
else: # 'AD'
min_scaleref = mim_centered.abs().max()
cfg_scaleref = torch.quantile(cfg_centered.abs(), self.threshold_percentile)
if self.scaling_startpoint == 'ZERO':
scaling_factor = min_scaleref / cfg_scaleref
result = cfg_flattened * scaling_factor
else: # 'MEAN'
if self.variability_measure == 'STD':
cfg_renormalized = (cfg_centered / cfg_scaleref) * min_scaleref
else: # 'AD'
### Get the maximum value of all datapoints (with an optional threshold percentile on the uncond)
max_scaleref = torch.maximum(min_scaleref, cfg_scaleref)
### Clamp to the max
cfg_clamped = cfg_centered.clamp(-max_scaleref, max_scaleref)
### Now shrink from the max to normalize and grow to the mimic scale (instead of the CFG scale)
cfg_renormalized = (cfg_clamped / max_scaleref) * min_scaleref
### Now add it back onto the averages to get into real scale again and return
result = cfg_renormalized + cfg_means
### Now add it back onto the averages to get into real scale again and return
result = cfg_renormalized + cfg_means
actualRes = result.unflatten(2, mim_target.shape[2:])
if self.interpolate_phi != 1.0:
actualRes = actualRes * self.interpolate_phi + cfg_target * (1.0 - self.interpolate_phi)
if self.experiment_mode == 1:
num = actualRes.cpu().numpy()
for y in range(0, 64):
+28 -5
View File
@@ -42,13 +42,19 @@ class Script(scripts.Script):
gr.HTML(value=f"<br>View <a style=\"border-bottom: 1px #00ffff dotted;\" href=\"https://github.com/mcmonkeyprojects/sd-dynamic-thresholding/wiki/Usage-Tips\">the wiki for usage tips.</a><br><br>", elem_id='dynthres_wiki_link')
mimic_scale = gr.Slider(minimum=1.0, maximum=30.0, step=0.5, label='Mimic CFG Scale', value=7.0, elem_id='dynthres_mimic_scale')
with gr.Accordion("Dynamic Thresholding Advanced Options", open=False, elem_id='dynthres_advanced_opts'):
threshold_percentile = gr.Slider(minimum=90.0, value=100.0, maximum=100.0, step=0.05, label='Top percentile of latents to clamp', elem_id='dynthres_threshold_percentile')
with gr.Row():
threshold_percentile = gr.Slider(minimum=90.0, value=100.0, maximum=100.0, step=0.05, label='Top percentile of latents to clamp', elem_id='dynthres_threshold_percentile')
interpolate_phi = gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label="Interpolate Phi", value=1.0, elem_id='dynthres_interpolate_phi')
with gr.Row():
mimic_mode = gr.Dropdown(VALID_MODES, value="Constant", label="Mimic Scale Scheduler", elem_id='dynthres_mimic_mode')
cfg_mode = gr.Dropdown(VALID_MODES, value="Constant", label="CFG Scale Scheduler", elem_id='dynthres_cfg_mode')
mimic_scale_min = gr.Slider(minimum=0.0, maximum=30.0, step=0.5, visible=False, label="Minimum value of the Mimic Scale Scheduler", elem_id='dynthres_mimic_scale_min')
cfg_scale_min = gr.Slider(minimum=0.0, maximum=30.0, step=0.5, visible=False, label="Minimum value of the CFG Scale Scheduler", elem_id='dynthres_cfg_scale_min')
sched_val = gr.Slider(minimum=0.0, maximum=40.0, step=0.5, value=4.0, visible=False, label="Scheduler Value", info="Value unique to the scheduler mode - for Power Up/Down, this is the power. For Linear/Cosine Repeating, this is the number of repeats per image.", elem_id='dynthres_sched_val')
with gr.Row():
separate_feature_channels = gr.Checkbox(value=True, label="Separate Feature Channels", elem_id='dynthres_separate_feature_channels')
scaling_startpoint = gr.Radio(["ZERO", "MEAN"], value="MEAN", label="Scaling Startpoint", elem_id='dynthres_scaling_startpoint')
variability_measure = gr.Radio(["STD", "AD"], value="AD", label="Variability Measure", elem_id='dynthres_variability_measure')
def shouldShowSchedulerValue(cfgMode, mimicMode):
sched_vis = cfgMode in MODES_WITH_VALUE or mimicMode in MODES_WITH_VALUE
return vis_change(sched_vis), vis_change(mimicMode != "Constant"), vis_change(cfgMode != "Constant")
@@ -63,17 +69,21 @@ class Script(scripts.Script):
(enabled, lambda d: gr.Checkbox.update(value="Dynamic thresholding enabled" in d)),
(accordion, lambda d: gr.Accordion.update(visible="Dynamic thresholding enabled" in d)),
(mimic_scale, "Mimic scale"),
(separate_feature_channels, "Separate Feature Channels"),
(scaling_startpoint, lambda d: gr.Radio.update(value=d.get("Scaling Startpoint", "MEAN"))),
(variability_measure, lambda d: gr.Radio.update(value=d.get("Variability Measure", "AD"))),
(interpolate_phi, "Interpolate Phi"),
(threshold_percentile, "Threshold percentile"),
(mimic_scale_min, "Mimic scale minimum"),
(mimic_mode, lambda d: gr.Dropdown.update(value=d.get("Mimic mode", "Constant"))),
(cfg_mode, lambda d: gr.Dropdown.update(value=d.get("CFG mode", "Constant"))),
(cfg_scale_min, "CFG scale minimum"),
(sched_val, "Scheduler value"))
return [enabled, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val]
return [enabled, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, separate_feature_channels, scaling_startpoint, variability_measure, interpolate_phi]
last_id = 0
def process_batch(self, p, enabled, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, batch_number, prompts, seeds, subseeds):
def process_batch(self, p, enabled, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, separate_feature_channels, scaling_startpoint, variability_measure, interpolate_phi, batch_number, prompts, seeds, subseeds):
enabled = getattr(p, 'dynthres_enabled', enabled)
if not enabled:
return
@@ -83,6 +93,10 @@ class Script(scripts.Script):
if orig_sampler_name == 'UniPC' and p.enable_hr:
raise RuntimeError(f"UniPC does not support Hires Fix. Auto WebUI silently swaps to DDIM for this, which DynThresh does not support. Please swap to a sampler capable of img2img processing for HR Fix to work.")
mimic_scale = getattr(p, 'dynthres_mimic_scale', mimic_scale)
separate_feature_channels = getattr(p, 'dynthres_separate_feature_channels', separate_feature_channels)
scaling_startpoint = getattr(p, 'dynthres_scaling_startpoint', scaling_startpoint)
variability_measure = getattr(p, 'dynthres_variability_measure', variability_measure)
interpolate_phi = getattr(p, 'dynthres_interpolate_phi', interpolate_phi)
threshold_percentile = getattr(p, 'dynthres_threshold_percentile', threshold_percentile)
mimic_mode = getattr(p, 'dynthres_mimic_mode', mimic_mode)
mimic_scale_min = getattr(p, 'dynthres_mimic_scale_min', mimic_scale_min)
@@ -92,6 +106,10 @@ class Script(scripts.Script):
sched_val = getattr(p, 'dynthres_scheduler_val', sched_val)
p.extra_generation_params["Dynamic thresholding enabled"] = True
p.extra_generation_params["Mimic scale"] = mimic_scale
p.extra_generation_params["Separate Feature Channels"] = separate_feature_channels
p.extra_generation_params["Scaling Startpoint"] = scaling_startpoint
p.extra_generation_params["Variability Measure"] = variability_measure
p.extra_generation_params["Interpolate Phi"] = interpolate_phi
p.extra_generation_params["Threshold percentile"] = threshold_percentile
p.extra_generation_params["Sampler"] = orig_sampler_name
if mimic_mode != "Constant":
@@ -109,7 +127,7 @@ class Script(scripts.Script):
threshold_percentile *= 0.01
# Make a placeholder sampler
sampler = sd_samplers.all_samplers_map[orig_sampler_name]
dtData = dynthres_core.DynThresh(mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, experiment_mode, p.steps)
dtData = dynthres_core.DynThresh(mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, experiment_mode, p.steps, separate_feature_channels, scaling_startpoint, variability_measure, interpolate_phi)
if orig_sampler_name == "UniPC":
def uniPCConstructor(model):
return CustomVanillaSDSampler(dynthres_unipc.CustomUniPCSampler, model, dtData)
@@ -129,7 +147,7 @@ class Script(scripts.Script):
if p.sampler is not None:
p.sampler = sd_samplers.create_sampler(fixed_sampler_name, p.sd_model)
def postprocess_batch(self, p, enabled, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, batch_number, images):
def postprocess_batch(self, p, enabled, mimic_scale, threshold_percentile, mimic_mode, mimic_scale_min, cfg_mode, cfg_scale_min, sched_val, separate_feature_channels, scaling_startpoint, variability_measure, interpolate_phi, batch_number, images):
if not enabled or not hasattr(p, 'orig_sampler_name'):
return
p.sampler_name = p.orig_sampler_name
@@ -188,6 +206,11 @@ def make_axis_options():
raise RuntimeError(f"Unknown Scheduler: {x}")
extra_axis_options = [
xyz_grid.AxisOption("[DynThres] Mimic Scale", float, apply_mimic_scale),
xyz_grid.AxisOption("[DynThres] Separate Feature Channels", int,
xyz_grid.apply_field("dynthres_separate_feature_channels")),
xyz_grid.AxisOption("[DynThres] Scaling Startpoint", str, xyz_grid.apply_field("dynthres_scaling_startpoint"), choices=lambda:['ZERO', 'MEAN']),
xyz_grid.AxisOption("[DynThres] Variability Measure", str, xyz_grid.apply_field("dynthres_variability_measure"), choices=lambda:['STD', 'AD']),
xyz_grid.AxisOption("[DynThres] Interpolate Phi", float, xyz_grid.apply_field("dynthres_interpolate_phi")),
xyz_grid.AxisOption("[DynThres] Threshold Percentile", float, xyz_grid.apply_field("dynthres_threshold_percentile")),
xyz_grid.AxisOption("[DynThres] Mimic Scheduler", str, xyz_grid.apply_field("dynthres_mimic_mode"), confirm=confirm_scheduler, choices=lambda: VALID_MODES),
xyz_grid.AxisOption("[DynThres] Mimic minimum", float, xyz_grid.apply_field("dynthres_mimic_scale_min")),