From 5dc9743592f7d41eb5efd49c10792f63ab783d57 Mon Sep 17 00:00:00 2001 From: nekoworkshop Date: Sat, 29 Apr 2023 14:41:37 -0400 Subject: [PATCH 1/5] Initial implementation --- modules/generation_parameters_copypaste.py | 2 +- scripts/xyz_grid.py | 42 ++++++++++++++++++++++ 2 files changed, 43 insertions(+), 1 deletion(-) diff --git a/modules/generation_parameters_copypaste.py b/modules/generation_parameters_copypaste.py index 964432d72..b6609e224 100644 --- a/modules/generation_parameters_copypaste.py +++ b/modules/generation_parameters_copypaste.py @@ -316,7 +316,7 @@ infotext_to_setting_name_mapping = [ ('Token merging merge attention', 'token_merging_merge_attention'), ('Token merging merge cross attention', 'token_merging_merge_cross_attention'), ('Token merging merge mlp', 'token_merging_merge_mlp'), - ('Token merging maximum downsampling', 'token_merging_maximum_downsampling'), + ('Token merging maximum downsampling', 'token_merging_maximum_down_sampling'), ('Token merging stride x', 'token_merging_stride_x'), ('Token merging stride y', 'token_merging_stride_y') ] diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index 700cc599e..68c294959 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -145,6 +145,39 @@ def apply_face_restore(p, opt, x): p.restore_faces = is_active +def apply_token_merging1(p, x, xs): + is_active = x.lower() in ('true', 'yes', 'y', '1') + p.override_settings["token_merging"] = is_active + +def apply_token_merging_ratio_hr(p, x, xs): + p.override_settings["token_merging_ratio_hr"] = x + +def apply_token_merging_ratio(p, x, xs): + p.override_settings["token_merging_ratio"] = x + +def apply_token_merging_hr_only(p, x, xs): + is_active = x.lower() in ('true', 'yes', 'y', '1') + p.override_settings["token_merging_hr_only"] = is_active + +def apply_token_merging_random(p, x, xs): + is_active = x.lower() in ('true', 'yes', 'y', '1') + p.override_settings["token_merging_random"] = is_active + +def apply_token_merging_attention(p, x, xs): + is_active = x.lower() in ('true', 'yes', 'y', '1') + p.override_settings["token_merging_merge_attention"] = is_active + +def apply_token_merging_cross_attention(p, x, xs): + is_active = x.lower() in ('true', 'yes', 'y', '1') + p.override_settings["token_merging_merge_cross_attention"] = is_active + +def apply_token_merging_mlp(p, x, xs): + is_active = x.lower() in ('true', 'yes', 'y', '1') + p.override_settings["token_merging_merge_mlp"] = is_active + +def apply_token_merging_maximum_down_sampling (p, x, xs): + p.override_settings["token_merging_maximum_down_sampling"] = x + #opts.data["token_merging_maximum_down_sampling"] = x def format_value_add_label(p, opt, x): if type(x) == float: @@ -226,6 +259,15 @@ axis_options = [ AxisOption("Styles", str, apply_styles, choices=lambda: list(shared.prompt_styles.styles)), AxisOption("UniPC Order", int, apply_uni_pc_order, cost=0.5), AxisOption("Face restore", str, apply_face_restore, format_value=format_value), + AxisOption("Token Merging", str, apply_token_merging1), + AxisOption("Token merging ratio",float,apply_token_merging_ratio), + AxisOption("Token merging ratio for Hires fix",float,apply_token_merging_ratio_hr), + AxisOption("Token merging apply only to Hires fix",str,apply_token_merging_hr_only, choices= lambda: ["Yes","No"]), + AxisOption("Token Merging use random pertubations",str,apply_token_merging_random, choices = lambda: ["Yes","No"]), + AxisOption("Token Merging merge attention", str, apply_token_merging_attention, choices= lambda: ["Yes","No"]), + AxisOption("Token Merging merge cross attention", str, apply_token_merging_cross_attention, choices= lambda: ["Yes","No"]), + AxisOption("Token Merging merge mlp", str, apply_token_merging_mlp, choices= lambda: ["Yes","No"]), + AxisOption("Token Merging maxium down sampling", int, apply_token_merging_maximum_down_sampling, choices= lambda: ["1","2","4","8"]) ] From 0e165ed2ee45238133d0132c018280c97eb0f053 Mon Sep 17 00:00:00 2001 From: nekoworkshop Date: Sat, 29 Apr 2023 17:18:01 -0400 Subject: [PATCH 2/5] Use SharedSettingsStackHelper --- scripts/xyz_grid.py | 54 ++++++++++++++++++++++++++++++++------------- 1 file changed, 39 insertions(+), 15 deletions(-) diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index 68c294959..2d785bf04 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -145,39 +145,40 @@ def apply_face_restore(p, opt, x): p.restore_faces = is_active -def apply_token_merging1(p, x, xs): - is_active = x.lower() in ('true', 'yes', 'y', '1') - p.override_settings["token_merging"] = is_active - def apply_token_merging_ratio_hr(p, x, xs): - p.override_settings["token_merging_ratio_hr"] = x + opts.data["token_merging_ratio_hr"] = x def apply_token_merging_ratio(p, x, xs): - p.override_settings["token_merging_ratio"] = x + opts.data["token_merging_ratio"] = x def apply_token_merging_hr_only(p, x, xs): is_active = x.lower() in ('true', 'yes', 'y', '1') - p.override_settings["token_merging_hr_only"] = is_active + opts.data["token_merging_hr_only"] = is_active def apply_token_merging_random(p, x, xs): is_active = x.lower() in ('true', 'yes', 'y', '1') - p.override_settings["token_merging_random"] = is_active + opts.data["token_merging_random"] = is_active def apply_token_merging_attention(p, x, xs): is_active = x.lower() in ('true', 'yes', 'y', '1') - p.override_settings["token_merging_merge_attention"] = is_active + opts.data["token_merging_merge_attention"] = is_active def apply_token_merging_cross_attention(p, x, xs): is_active = x.lower() in ('true', 'yes', 'y', '1') - p.override_settings["token_merging_merge_cross_attention"] = is_active + opts.data["token_merging_merge_cross_attention"] = is_active def apply_token_merging_mlp(p, x, xs): is_active = x.lower() in ('true', 'yes', 'y', '1') - p.override_settings["token_merging_merge_mlp"] = is_active + opts.data["token_merging_merge_mlp"] = is_active def apply_token_merging_maximum_down_sampling (p, x, xs): - p.override_settings["token_merging_maximum_down_sampling"] = x - #opts.data["token_merging_maximum_down_sampling"] = x + opts.data["token_merging_maximum_down_sampling"] = x + +def apply_token_merging_stride_x(p, x, xs): + opts.data["token_merging_stride_x"] = x + +def apply_token_merging_stride_y(p, x, xs): + opts.data["token_merging_stride_y"] = x def format_value_add_label(p, opt, x): if type(x) == float: @@ -259,7 +260,6 @@ axis_options = [ AxisOption("Styles", str, apply_styles, choices=lambda: list(shared.prompt_styles.styles)), AxisOption("UniPC Order", int, apply_uni_pc_order, cost=0.5), AxisOption("Face restore", str, apply_face_restore, format_value=format_value), - AxisOption("Token Merging", str, apply_token_merging1), AxisOption("Token merging ratio",float,apply_token_merging_ratio), AxisOption("Token merging ratio for Hires fix",float,apply_token_merging_ratio_hr), AxisOption("Token merging apply only to Hires fix",str,apply_token_merging_hr_only, choices= lambda: ["Yes","No"]), @@ -267,7 +267,9 @@ axis_options = [ AxisOption("Token Merging merge attention", str, apply_token_merging_attention, choices= lambda: ["Yes","No"]), AxisOption("Token Merging merge cross attention", str, apply_token_merging_cross_attention, choices= lambda: ["Yes","No"]), AxisOption("Token Merging merge mlp", str, apply_token_merging_mlp, choices= lambda: ["Yes","No"]), - AxisOption("Token Merging maxium down sampling", int, apply_token_merging_maximum_down_sampling, choices= lambda: ["1","2","4","8"]) + AxisOption("Token Merging maxium down sampling", int, apply_token_merging_maximum_down_sampling, choices= lambda: ["1","2","4","8"]), + AxisOption("Token Merging Stride - X", int, apply_token_merging_stride_x, choices= lambda: ["2","4","6","8"]), + AxisOption("Token Merging Stride - Y", int, apply_token_merging_stride_y, choices= lambda: ["2","4","6","8"]) ] @@ -384,11 +386,23 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend class SharedSettingsStackHelper(object): def __enter__(self): + #Save overridden settings so they can be restored later. self.CLIP_stop_at_last_layers = opts.CLIP_stop_at_last_layers self.vae = opts.sd_vae self.uni_pc_order = opts.uni_pc_order + self.token_merging_ratio_hr = opts.token_merging_ratio_hr + self.token_merging_ratio = opts.token_merging_ratio + self.token_merging_hr_only = opts.token_merging_hr_only + self.token_merging_random = opts.token_merging_random + self.token_merging_merge_attention = opts.token_merging_merge_attention + self.token_merging_merge_cross_attention = opts.token_merging_merge_cross_attention + self.token_merging_merge_mlp = opts.token_merging_merge_mlp + self.token_merging_maximum_down_sampling = opts.token_merging_maximum_down_sampling + self.token_merging_stride_x = opts.token_merging_stride_x + self.token_merging_stride_y = opts.token_merging_stride_y def __exit__(self, exc_type, exc_value, tb): + #Restore overriden settings after plot generation. opts.data["sd_vae"] = self.vae opts.data["uni_pc_order"] = self.uni_pc_order sd_models.reload_model_weights() @@ -396,6 +410,16 @@ class SharedSettingsStackHelper(object): opts.data["CLIP_stop_at_last_layers"] = self.CLIP_stop_at_last_layers + opts.data["token_merging_ratio_hr"] = self.token_merging_ratio_hr + opts.data["token_merging_ratio"] = self.token_merging_ratio + opts.data["token_merging_hr_only"] = self.token_merging_hr_only + opts.data["token_merging_random"] = self.token_merging_random + opts.data["token_merging_merge_attention"] = self.token_merging_merge_attention + opts.data["token_merging_merge_cross_attention"] = self.token_merging_merge_cross_attention + opts.data["token_merging_merge_mlp"] = self.token_merging_merge_mlp + opts.data["token_merging_maximum_down_sampling"] = self.token_merging_maximum_down_sampling + opts.data["token_merging_stride_x"] = self.token_merging_stride_x + opts.data["token_merging_stride_y"] = self.token_merging_stride_y re_range = re.compile(r"\s*([+-]?\s*\d+)\s*-\s*([+-]?\s*\d+)(?:\s*\(([+-]\d+)\s*\))?\s*") re_range_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*-\s*([+-]?\s*\d+(?:.\d*)?)(?:\s*\(([+-]\d+(?:.\d*)?)\s*\))?\s*") From e97443ff1b9bc76d48f5eee1e8d8c88a8a6b23f3 Mon Sep 17 00:00:00 2001 From: nekoworkshop Date: Sat, 29 Apr 2023 17:37:58 -0400 Subject: [PATCH 3/5] Adjust axis names to be shorter. --- scripts/xyz_grid.py | 20 ++++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index 2d785bf04..f47f0d167 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -260,16 +260,16 @@ axis_options = [ AxisOption("Styles", str, apply_styles, choices=lambda: list(shared.prompt_styles.styles)), AxisOption("UniPC Order", int, apply_uni_pc_order, cost=0.5), AxisOption("Face restore", str, apply_face_restore, format_value=format_value), - AxisOption("Token merging ratio",float,apply_token_merging_ratio), - AxisOption("Token merging ratio for Hires fix",float,apply_token_merging_ratio_hr), - AxisOption("Token merging apply only to Hires fix",str,apply_token_merging_hr_only, choices= lambda: ["Yes","No"]), - AxisOption("Token Merging use random pertubations",str,apply_token_merging_random, choices = lambda: ["Yes","No"]), - AxisOption("Token Merging merge attention", str, apply_token_merging_attention, choices= lambda: ["Yes","No"]), - AxisOption("Token Merging merge cross attention", str, apply_token_merging_cross_attention, choices= lambda: ["Yes","No"]), - AxisOption("Token Merging merge mlp", str, apply_token_merging_mlp, choices= lambda: ["Yes","No"]), - AxisOption("Token Merging maxium down sampling", int, apply_token_merging_maximum_down_sampling, choices= lambda: ["1","2","4","8"]), - AxisOption("Token Merging Stride - X", int, apply_token_merging_stride_x, choices= lambda: ["2","4","6","8"]), - AxisOption("Token Merging Stride - Y", int, apply_token_merging_stride_y, choices= lambda: ["2","4","6","8"]) + AxisOption("ToMe ratio",float,apply_token_merging_ratio), + AxisOption("ToMe ratio for Hires fix",float,apply_token_merging_ratio_hr), + AxisOption("ToMe apply only to Hires fix",str,apply_token_merging_hr_only, choices= lambda: ["Yes","No"]), + AxisOption("ToMe random pertubations",str,apply_token_merging_random, choices = lambda: ["Yes","No"]), + AxisOption("ToMe merge attention", str, apply_token_merging_attention, choices= lambda: ["Yes","No"]), + AxisOption("ToMe merge cross attention", str, apply_token_merging_cross_attention, choices= lambda: ["Yes","No"]), + AxisOption("ToMe merge mlp", str, apply_token_merging_mlp, choices= lambda: ["Yes","No"]), + AxisOption("ToMe maximum down sampling", int, apply_token_merging_maximum_down_sampling, choices= lambda: ["1","2","4","8"]), + AxisOption("ToMe Stride - X", int, apply_token_merging_stride_x, choices= lambda: ["2","4","6","8"]), + AxisOption("ToMe Stride - Y", int, apply_token_merging_stride_y, choices= lambda: ["2","4","6","8"]) ] From c4936fc92784c452bb067b5097b54476400b1abc Mon Sep 17 00:00:00 2001 From: nekoworkshop Date: Sat, 29 Apr 2023 17:59:00 -0400 Subject: [PATCH 4/5] Extra information in ToMe related settings --- modules/shared.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/modules/shared.py b/modules/shared.py index 0dfefa8c4..dc8ac5067 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -422,14 +422,14 @@ options_templates.update(options_section(('sampler-params', "Sampler parameters" options_templates.update(options_section(('token_merging', 'Token Merging'), { "token_merging": OptionInfo(False, "Enable redundant token merging via tomesd. This can provide significant speed and memory improvements.", gr.Checkbox), - "token_merging_ratio": OptionInfo(0.5, "Merging Ratio", gr.Slider, {"minimum": 0, "maximum": 0.9, "step": 0.1}), + "token_merging_ratio": OptionInfo(0.5, "Merging Ratio. Higher merging ratio = faster generation, smaller VRAM usage, lower quality.", gr.Slider, {"minimum": 0, "maximum": 0.9, "step": 0.1}), "token_merging_hr_only": OptionInfo(True, "Apply only to high-res fix pass. Disabling can yield a ~20-35% speedup on contemporary resolutions.", gr.Checkbox), "token_merging_ratio_hr": OptionInfo(0.5, "Merging Ratio (high-res pass) - If 'Apply only to high-res' is enabled, this will always be the ratio used.", gr.Slider, {"minimum": 0, "maximum": 0.9, "step": 0.1}), "token_merging_random": OptionInfo(False, "Use random perturbations - Can improve outputs for certain samplers. For others, it may cause visual artifacting.", gr.Checkbox), - "token_merging_merge_attention": OptionInfo(True, "Merge attention", gr.Checkbox), - "token_merging_merge_cross_attention": OptionInfo(False, "Merge cross attention", gr.Checkbox), - "token_merging_merge_mlp": OptionInfo(False, "Merge mlp", gr.Checkbox), - "token_merging_maximum_down_sampling": OptionInfo(1, "Maximum down sampling", gr.Dropdown, lambda: {"choices": ["1", "2", "4", "8"]}), + "token_merging_merge_attention": OptionInfo(True, "Merge attention (Recommend on)", gr.Checkbox), + "token_merging_merge_cross_attention": OptionInfo(False, "Merge cross attention (Recommend off)", gr.Checkbox), + "token_merging_merge_mlp": OptionInfo(False, "Merge mlp (Strongly recommend off)", gr.Checkbox), + "token_merging_maximum_down_sampling": OptionInfo(1, "Maximum down sampling", gr.Radio, lambda: {"choices": [1, 2, 4, 8]}), "token_merging_stride_x": OptionInfo(2, "Stride - X", gr.Slider, {"minimum": 2, "maximum": 8, "step": 2}), "token_merging_stride_y": OptionInfo(2, "Stride - Y", gr.Slider, {"minimum": 2, "maximum": 8, "step": 2}) })) From df965a837bb6374b25b4f274b752805d68e38634 Mon Sep 17 00:00:00 2001 From: nekoworkshop Date: Sat, 29 Apr 2023 23:01:31 -0400 Subject: [PATCH 5/5] Remove less useful ToMe options from the xyz plot --- scripts/xyz_grid.py | 48 +-------------------------------------------- 1 file changed, 1 insertion(+), 47 deletions(-) diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index f47f0d167..d216d2fef 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -151,35 +151,10 @@ def apply_token_merging_ratio_hr(p, x, xs): def apply_token_merging_ratio(p, x, xs): opts.data["token_merging_ratio"] = x -def apply_token_merging_hr_only(p, x, xs): - is_active = x.lower() in ('true', 'yes', 'y', '1') - opts.data["token_merging_hr_only"] = is_active - def apply_token_merging_random(p, x, xs): is_active = x.lower() in ('true', 'yes', 'y', '1') opts.data["token_merging_random"] = is_active -def apply_token_merging_attention(p, x, xs): - is_active = x.lower() in ('true', 'yes', 'y', '1') - opts.data["token_merging_merge_attention"] = is_active - -def apply_token_merging_cross_attention(p, x, xs): - is_active = x.lower() in ('true', 'yes', 'y', '1') - opts.data["token_merging_merge_cross_attention"] = is_active - -def apply_token_merging_mlp(p, x, xs): - is_active = x.lower() in ('true', 'yes', 'y', '1') - opts.data["token_merging_merge_mlp"] = is_active - -def apply_token_merging_maximum_down_sampling (p, x, xs): - opts.data["token_merging_maximum_down_sampling"] = x - -def apply_token_merging_stride_x(p, x, xs): - opts.data["token_merging_stride_x"] = x - -def apply_token_merging_stride_y(p, x, xs): - opts.data["token_merging_stride_y"] = x - def format_value_add_label(p, opt, x): if type(x) == float: x = round(x, 8) @@ -262,14 +237,7 @@ axis_options = [ AxisOption("Face restore", str, apply_face_restore, format_value=format_value), AxisOption("ToMe ratio",float,apply_token_merging_ratio), AxisOption("ToMe ratio for Hires fix",float,apply_token_merging_ratio_hr), - AxisOption("ToMe apply only to Hires fix",str,apply_token_merging_hr_only, choices= lambda: ["Yes","No"]), - AxisOption("ToMe random pertubations",str,apply_token_merging_random, choices = lambda: ["Yes","No"]), - AxisOption("ToMe merge attention", str, apply_token_merging_attention, choices= lambda: ["Yes","No"]), - AxisOption("ToMe merge cross attention", str, apply_token_merging_cross_attention, choices= lambda: ["Yes","No"]), - AxisOption("ToMe merge mlp", str, apply_token_merging_mlp, choices= lambda: ["Yes","No"]), - AxisOption("ToMe maximum down sampling", int, apply_token_merging_maximum_down_sampling, choices= lambda: ["1","2","4","8"]), - AxisOption("ToMe Stride - X", int, apply_token_merging_stride_x, choices= lambda: ["2","4","6","8"]), - AxisOption("ToMe Stride - Y", int, apply_token_merging_stride_y, choices= lambda: ["2","4","6","8"]) + AxisOption("ToMe random pertubations",str,apply_token_merging_random, choices = lambda: ["Yes","No"]) ] @@ -392,14 +360,7 @@ class SharedSettingsStackHelper(object): self.uni_pc_order = opts.uni_pc_order self.token_merging_ratio_hr = opts.token_merging_ratio_hr self.token_merging_ratio = opts.token_merging_ratio - self.token_merging_hr_only = opts.token_merging_hr_only self.token_merging_random = opts.token_merging_random - self.token_merging_merge_attention = opts.token_merging_merge_attention - self.token_merging_merge_cross_attention = opts.token_merging_merge_cross_attention - self.token_merging_merge_mlp = opts.token_merging_merge_mlp - self.token_merging_maximum_down_sampling = opts.token_merging_maximum_down_sampling - self.token_merging_stride_x = opts.token_merging_stride_x - self.token_merging_stride_y = opts.token_merging_stride_y def __exit__(self, exc_type, exc_value, tb): #Restore overriden settings after plot generation. @@ -412,14 +373,7 @@ class SharedSettingsStackHelper(object): opts.data["token_merging_ratio_hr"] = self.token_merging_ratio_hr opts.data["token_merging_ratio"] = self.token_merging_ratio - opts.data["token_merging_hr_only"] = self.token_merging_hr_only opts.data["token_merging_random"] = self.token_merging_random - opts.data["token_merging_merge_attention"] = self.token_merging_merge_attention - opts.data["token_merging_merge_cross_attention"] = self.token_merging_merge_cross_attention - opts.data["token_merging_merge_mlp"] = self.token_merging_merge_mlp - opts.data["token_merging_maximum_down_sampling"] = self.token_merging_maximum_down_sampling - opts.data["token_merging_stride_x"] = self.token_merging_stride_x - opts.data["token_merging_stride_y"] = self.token_merging_stride_y re_range = re.compile(r"\s*([+-]?\s*\d+)\s*-\s*([+-]?\s*\d+)(?:\s*\(([+-]\d+)\s*\))?\s*") re_range_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*-\s*([+-]?\s*\d+(?:.\d*)?)(?:\s*\(([+-]\d+(?:.\d*)?)\s*\))?\s*")