From 96739fab6e67346961059e880411e0dc2b614c5a Mon Sep 17 00:00:00 2001 From: Convexity-ai <133716788+Convexity-ai@users.noreply.github.com> Date: Sun, 30 Jul 2023 04:32:27 +0100 Subject: [PATCH] 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 --- dynthres_core.py | 51 ++++++++++++++++++++++++++------- scripts/dynamic_thresholding.py | 33 +++++++++++++++++---- 2 files changed, 68 insertions(+), 16 deletions(-) diff --git a/dynthres_core.py b/dynthres_core.py index f5f2994..ab387a9 100644 --- a/dynthres_core.py +++ b/dynthres_core.py @@ -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): diff --git a/scripts/dynamic_thresholding.py b/scripts/dynamic_thresholding.py index 667b63c..827a02d 100644 --- a/scripts/dynamic_thresholding.py +++ b/scripts/dynamic_thresholding.py @@ -42,13 +42,19 @@ class Script(scripts.Script): gr.HTML(value=f"
View the wiki for usage tips.

", 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")),