From 5dc9743592f7d41eb5efd49c10792f63ab783d57 Mon Sep 17 00:00:00 2001 From: nekoworkshop Date: Sat, 29 Apr 2023 14:41:37 -0400 Subject: [PATCH] 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"]) ]