From 685f3929676ddf82505f6a05a362d112d6ef0e75 Mon Sep 17 00:00:00 2001 From: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> Date: Sun, 5 Nov 2023 07:10:05 -0600 Subject: [PATCH] Added Tabs --- modules/extras.py | 117 +++++++++++++++++++------------------ modules/ui_models.py | 134 ++++++++++++++++++++++++++++++------------- 2 files changed, 155 insertions(+), 96 deletions(-) diff --git a/modules/extras.py b/modules/extras.py index 0d5fc2c76..ae7b9af8a 100644 --- a/modules/extras.py +++ b/modules/extras.py @@ -12,7 +12,6 @@ from sd_meh.merge import merge_models from modules import shared, images, sd_models, sd_vae, sd_models_config - checkpoint_dict_skip_on_merge = ["cond_stage_model.transformer.text_model.embeddings.position_ids"] @@ -32,6 +31,7 @@ def create_config(ckpt_result, config_source, a, b, c): def config(x): res = sd_models_config.find_checkpoint_config_near_filename(x) if x else None return res if res != shared.sd_default_config else None + if config_source == 0: cfg = config(a) or config(b) or config(c) elif config_source == 1: @@ -54,7 +54,9 @@ def to_half(tensor, enable): return tensor -def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_model_name, interp_method, multiplier, save_as_half, custom_name, checkpoint_format, config_source, bake_in_vae, discard_weights, save_metadata): # pylint: disable=unused-argument +def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_model_name, interp_method, multiplier, + save_as_half, custom_name, checkpoint_format, config_source, bake_in_vae, discard_weights, + save_metadata): # pylint: disable=unused-argument shared.state.begin('merge') save_as_half = save_as_half == 0 @@ -148,14 +150,18 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_ # have another 4 channels for unmasked picture's latent space, plus one channel for mask, for a total of 9 if a.shape != b.shape and a.shape[0:1] + a.shape[2:] == b.shape[0:1] + b.shape[2:]: if a.shape[1] == 4 and b.shape[1] == 9: - raise RuntimeError("When merging inpainting model with a normal one, A must be the inpainting model.") + raise RuntimeError( + "When merging inpainting model with a normal one, A must be the inpainting model.") if a.shape[1] == 4 and b.shape[1] == 8: - raise RuntimeError("When merging instruct-pix2pix model with a normal one, A must be the instruct-pix2pix model.") - if a.shape[1] == 8 and b.shape[1] == 4:#If we have an Instruct-Pix2Pix model... - theta_0[key][:, 0:4, :, :] = theta_func2(a[:, 0:4, :, :], b, multiplier)#Merge only the vectors the models have in common. Otherwise we get an error due to dimension mismatch. + raise RuntimeError( + "When merging instruct-pix2pix model with a normal one, A must be the instruct-pix2pix model.") + if a.shape[1] == 8 and b.shape[1] == 4: # If we have an Instruct-Pix2Pix model... + theta_0[key][:, 0:4, :, :] = theta_func2(a[:, 0:4, :, :], b, + multiplier) # Merge only the vectors the models have in common. Otherwise we get an error due to dimension mismatch. result_is_instruct_pix2pix_model = True else: - assert a.shape[1] == 9 and b.shape[1] == 4, f"Bad dimensions for merged layer {key}: A={a.shape}, B={b.shape}" + assert a.shape[1] == 9 and b.shape[ + 1] == 4, f"Bad dimensions for merged layer {key}: A={a.shape}, B={b.shape}" theta_0[key][:, 0:4, :, :] = theta_func2(a[:, 0:4, :, :], b, multiplier) result_is_inpainting_model = True else: @@ -193,7 +199,7 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_ if save_metadata: metadata = {"format": "pt", "sd_merge_models": {}} merge_recipe = { - "type": "webui", # indicate this model was merged with webui's built-in merger + "type": "webui", # indicate this model was merged with webui's built-in merger "primary_model_hash": primary_model_info.sha256, "secondary_model_hash": secondary_model_info.sha256 if secondary_model_info else None, "tertiary_model_hash": tertiary_model_info.sha256 if tertiary_model_info else None, @@ -238,70 +244,68 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_ shared.log.info(f"Model merge saved: {output_modelname}.") shared.state.textinfo = "Checkpoint saved" shared.state.end() - return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], "Checkpoint saved to " + output_modelname] + return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], + "Checkpoint saved to " + output_modelname] -def run_MEHmodelmerger(id_task, primary_model_name, secondary_model_name, tertiary_model_name, merge_mode,base_alpha, - base_beta, weights_alpha, - weights_beta, precision, custom_name, checkpoint_format, save_metadata, preset, weights_clip, prune, - re_basin, re_basin_iterations, device): # pylint: disable=unused-argument +def run_MEHmodelmerger(id_task, **kwargs): # pylint: disable=unused-argument shared.state.begin('model-merge') - models = { - "model_a": sd_models.checkpoints_list[primary_model_name].filename, - "model_b": sd_models.checkpoints_list[secondary_model_name].filename, - } - if tertiary_model_name is not None: - models |= {"model_c": sd_models.checkpoints_list[tertiary_model_name].filename} - work_device = device - threads = 1 - preset = preset if preset != "None" else None - block_weights_preset_alpha = block_weights_preset_beta = block_weights_preset_alpha_b = block_weights_preset_beta_b = preset if preset is not None else None - presets_alpha_lambda = presets_beta_lambda = None - logging_level = "INFO" + kwargs["models"] = { + "model_a": sd_models.checkpoints_list[kwargs.get("primary_model_name", None)].filename, + "model_b": sd_models.checkpoints_list[kwargs.get("secondary_model_name", None)].filename, + } + del kwargs["primary_model_name"] + del kwargs["secondary_model_name"] + if kwargs.get("tertiary_model_name", None) is not None: + kwargs["models"] |= {"model_c": sd_models.checkpoints_list[kwargs.get("tertiary_model_name", None)].filename} + del kwargs["tertiary_model_name"] + try: + alpha = [float(x) for x in [kwargs["alpha_base"]] + kwargs["alpha_in_blocks"].split(",") + [kwargs["alpha_mid_block"]] + kwargs["alpha_out_blocks"].split(",")] + assert len(alpha) == 26 or len(alpha) == 20, "Alpha Block Weights are wrong length (26 or 20 for SDXL) falling back" + kwargs["alpha"] = alpha + except Exception as e: + kwargs["alpha"] = kwargs.get("alpha_preset", kwargs["alpha"]) + print(e) + finally: + kwargs.pop("alpha_base", None) + kwargs.pop("alpha_in_blocks", None) + kwargs.pop("alpha_mid_block", None) + kwargs.pop("alpha_out_blocks", None) + kwargs.pop("alpha_preset", None) + if kwargs.get("beta", False): + try: + beta = [float(x) for x in [kwargs["beta_base"]] + kwargs["beta_in_blocks"].split(",") + [kwargs["beta_mid_block"]] + kwargs["beta_out_blocks"].split(",")] + assert len(beta) == 26 or len(beta) == 20, "Beta Block Weights are wrong length (26 or 20 for SDXL) falling back" + kwargs["beta"] = beta + except Exception as e: + kwargs["beta"] = kwargs.get("beta_preset", kwargs["beta"]) + print(e) + finally: + kwargs.pop("beta_base", None) + kwargs.pop("beta_in_blocks", None) + kwargs.pop("beta_mid_block", None) + kwargs.pop("beta_out_blocks", None) + kwargs.pop("beta_preset", None) + + return [*[gr.update() for _ in range(4)], f"{kwargs}"] def fail(message): shared.state.textinfo = message shared.state.end() return [*[gr.update() for _ in range(4)], message] - - - theta_0 = main( - model_a, - model_b, - model_c, - merge_mode, - weights_clip, - precision, - str(weights_alpha), - base_alpha, - str(weights_beta), - base_beta, - re_basin, - re_basin_iterations, - device, - work_device, - prune, - block_weights_preset_alpha, - block_weights_preset_beta, - threads, - block_weights_preset_alpha_b, - block_weights_preset_beta_b, - presets_alpha_lambda, - presets_beta_lambda, - logging_level, - ) + theta_0 = merge_models(**kwargs) ckpt_dir = shared.opts.ckpt_dir or sd_models.model_path filename = custom_name - filename += "." + checkpoint_format + filename += "." + kwargs.get("checkpoint_format", None) output_modelname = os.path.join(ckpt_dir, filename) shared.state.textinfo = "Saving" metadata = None if save_metadata: metadata = {"format": "pt", "sd_merge_models": {}} merge_recipe = { - "type": "webui", # indicate this model was merged with webui's built-in merger + "type": "SDNext", # indicate this model was merged with webui's built-in merger "primary_model_hash": primary_model_info.sha256, "secondary_model_hash": secondary_model_info.sha256 if secondary_model_info else None, "tertiary_model_hash": tertiary_model_info.sha256 if tertiary_model_info else None, @@ -345,8 +349,9 @@ def run_MEHmodelmerger(id_task, primary_model_name, secondary_model_name, tertia return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], "Checkpoint saved to " + output_modelname] -def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_name, unet_conv, text_encoder_conv, vae_conv, others_conv, fix_clip): +def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_name, unet_conv, text_encoder_conv, + vae_conv, others_conv, fix_clip): # position_ids in clip is int64. model_ema.num_updates is int32 dtypes_to_fp16 = {torch.float32, torch.float64, torch.bfloat16} dtypes_to_bf16 = {torch.float32, torch.float64, torch.float16} @@ -384,7 +389,6 @@ def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_nam state_dict = m["state_dict"] if "state_dict" in m else m return state_dict - def fix_model(model, fix_clip=False): # code from model-toolkit nai_keys = { @@ -446,6 +450,7 @@ def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_nam ok[wk] = t elif conv_t == "delete": return + shared.log.info("Model convert: running") if conv_type == "ema-only": for k in tqdm.tqdm(state_dict): diff --git a/modules/ui_models.py b/modules/ui_models.py index e5115662c..abcb9bce1 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -1,5 +1,6 @@ import os import json +import inspect from datetime import datetime import gradio as gr from modules import sd_models, sd_vae, extras @@ -140,9 +141,9 @@ def create_ui(): models_outcome, ] ) - with gr.Tab(label="MEH Merge"): + with gr.Tab(label="Advanced Merge"): def sd_model_choices(): - return ['None'] + sd_models.checkpoint_tiles() + return ['None'] + sd_models.checkpoint_tiles() with gr.Row(equal_height=False): with gr.Column(variant='compact'): with FormRow(): @@ -157,31 +158,35 @@ def create_ui(): tertiary_model_name = gr.Dropdown(sd_model_choices(), label="Tertiary model", value="None", visible=False) tertiary_refresh = create_refresh_button(tertiary_model_name, sd_models.list_models, lambda: {"choices": sd_model_choices()}, "refresh_checkpoint_C",visible=False) with FormRow(): - alpha = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Alpha Ratio', value=0.5) - beta = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Beta Ratio', value=None, visible=False) - with InputAccordion(False, label="Block Merge", elem_id=f"block_merge") as block_accordion: - with FormRow(): - alpha_label = gr.Markdown("# Alpha") - with FormRow(): - preset = gr.Dropdown(choices=["None"]+list(BLOCK_WEIGHTS_PRESETS.keys()), value=None, label="Block Weight Preset", multiselect=True, max_choices=2) - preset_lambda = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Preset Interpolation Ratio', value=None, visible=False) - apply_preset = ToolButton('⇩', visible=True) - with FormRow(): - base = gr.Textbox(value=None, label="Base", scale=1) - in_blocks = gr.Textbox(value=None, label="In Blocks", scale=10) - mid_block = gr.Textbox(value=None, label="Mid Block", scale=1) - out_blocks = gr.Textbox(value=None, label="Out Block", scale=10) - with FormRow(): - beta_label = gr.Markdown("# Beta", visible=False) - with FormRow(): - beta_preset = gr.Dropdown(choices=["None"]+list(BLOCK_WEIGHTS_PRESETS.keys()), value=None, label="Block Weight Preset", multiselect=True, max_choices=2, interactive=True, visible=False) - beta_preset_lambda = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Preset Interpolation Ratio', value=None, interactive=True, visible=False) - beta_apply_preset = ToolButton('⇩', interactive=True, visible=False) - with FormRow(): - beta_base = gr.Textbox(value=None, label="Base", scale=1, interactive=True, visible=False) - beta_in_blocks = gr.Textbox(value=None, label="In Blocks", interactive=True, scale=10, visible=False) - beta_mid_block = gr.Textbox(value=None, label="Mid Block", interactive=True, scale=1, visible=False) - beta_out_blocks = gr.Textbox(value=None, label="Out Block", interactive=True, scale=10, visible=False) + with gr.Tabs() as tabs: + with gr.TabItem(label="Simple Merge", id=0): + with FormRow(): + alpha = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Alpha Ratio', value=0.5) + beta = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Beta Ratio', value=None, visible=False) + with gr.TabItem(label="Preset Block Merge", id=1): + with FormRow(): + alpha_preset = gr.Dropdown(choices=["None"]+list(BLOCK_WEIGHTS_PRESETS.keys()), value=None, label="ALPHA Block Weight Preset", multiselect=True, max_choices=2) + alpha_preset_lambda = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Preset Interpolation Ratio', value=None, visible=False) + apply_preset = ToolButton('⇩', visible=True) + with FormRow(): + beta_preset = gr.Dropdown(choices=["None"]+list(BLOCK_WEIGHTS_PRESETS.keys()), value=None, label="BETA Block Weight Preset", multiselect=True, max_choices=2, interactive=True, visible=False) + beta_preset_lambda = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Preset Interpolation Ratio', value=None, interactive=True, visible=False) + beta_apply_preset = ToolButton('⇩', interactive=True, visible=False) + with gr.TabItem(label="Manual Block Merge", id=2): + with FormRow(): + alpha_label = gr.Markdown("# Alpha") + with FormRow(): + alpha_base = gr.Textbox(value=None, label="Base", min_width=70, scale=1) + alpha_in_blocks = gr.Textbox(value=None, label="In Blocks", scale=15) + alpha_mid_block = gr.Textbox(value=None, label="Mid Block", min_width=80, scale=1) + alpha_out_blocks = gr.Textbox(value=None, label="Out Block", scale=15) + with FormRow(): + beta_label = gr.Markdown("# Beta", visible=False) + with FormRow(): + beta_base = gr.Textbox(value=None, label="Base", min_width=70, scale=1, interactive=True, visible=False) + beta_in_blocks = gr.Textbox(value=None, label="In Blocks", interactive=True, scale=15, visible=False) + beta_mid_block = gr.Textbox(value=None, label="Mid Block", min_width=80, interactive=True, scale=1, visible=False) + beta_out_blocks = gr.Textbox(value=None, label="Out Block", interactive=True, scale=15, visible=False) with FormRow(): weights_clip = gr.Checkbox(label="Weights Clip") prune = gr.Checkbox(label="Prune") @@ -189,22 +194,59 @@ def create_ui(): with FormRow(): re_basin_iterations = gr.Slider(minimum=0, maximum=25, step=1, label='Number of ReBasin Iterations', value=None, visible=False) with FormRow(): - checkpoint_format = gr.Radio(choices=["ckpt", "safetensors"], value="safetensors", label="Model format") + checkpoint_format = gr.Radio(choices=["ckpt", "safetensors"], value="safetensors", visible=False, label="Model format") with FormRow(): precision = gr.Radio(choices=["fp16", "fp32"], value="fp16", label="Model precision") with FormRow(): device = gr.Radio(choices=["cpu", "cuda"], value="cpu", label="Device") with FormRow(): - bake_in_vae = gr.Dropdown(choices=["None"] + list(sd_vae.vae_dict), value="None", label="Bake in VAE") + bake_in_vae = gr.Dropdown(choices=["None"] + list(sd_vae.vae_dict), value="None", interactive=True, label="Bake in VAE") create_refresh_button(bake_in_vae, sd_vae.refresh_vae_list, lambda: {"choices": ["None"] + list(sd_vae.vae_dict)}, "modelmerger_refresh_bake_in_vae") with FormRow(): save_metadata = gr.Checkbox(value=True, label="Save metadata") with gr.Row(): MEHmodelmerger_merge = gr.Button(value="Merge", variant='primary') - def MEHmodelmerger(*args): + def MEHmodelmerger(dummy_component, + primary_model_name, + secondary_model_name, + tertiary_model_name, + merge_mode, + alpha, + beta, + alpha_preset, + alpha_preset_lambda, + alpha_base, + alpha_in_blocks, + alpha_mid_block, + alpha_out_blocks, + beta_preset, + beta_preset_lambda, + beta_base, + beta_in_blocks, + beta_mid_block, + beta_out_blocks, + precision, + custom_name, + checkpoint_format, + save_metadata, + weights_clip, + prune, + re_basin, + re_basin_iterations, + device, + bake_in_vae): + kwargs = {} + for x in inspect.getfullargspec(MEHmodelmerger)[0]: + kwargs[x] = locals()[x] + for key in list(kwargs.keys()): + if kwargs[key] in [None,"None","",0,[]]: + del kwargs[key] + + # return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"{kwargs}"] + try: - results = extras.run_MEHmodelmerger(*args) + results = extras.run_MEHmodelmerger(dummy_component, **kwargs) except Exception as e: modules.errors.display(e, 'model merge') sd_models.list_models() # to remove the potentially missing models from the list @@ -240,17 +282,17 @@ def create_ui(): preset = interpolate(presets, ratio) else: preset = presets[0] - preset = [str(x) for x in preset] - preset = [preset[0],",".join(preset[1:13]),preset[13],",".join(preset[14:])] - print(preset) - return [gr.update(value=x) for x in preset] + preset = ['%.3f' % x for x in preset] + preset = [preset[0], ",".join(preset[1:13]),preset[13], ",".join(preset[14:])] + return [gr.update(value=x) for x in preset]+[gr.update(selected=2)] - preset.change(fn=preset_visiblility, inputs=preset, outputs=preset_lambda) - beta_preset.change(fn=preset_visiblility, inputs=preset, outputs=beta_preset_lambda) + alpha_preset.change(fn=preset_visiblility, inputs=alpha_preset, outputs=alpha_preset_lambda) + beta_preset.change(fn=preset_visiblility, inputs=alpha_preset, outputs=beta_preset_lambda) merge_mode.input(fn=tertiary, inputs=merge_mode, outputs=[tertiary_model_name, tertiary_refresh]) merge_mode.input(fn=beta_visibility, inputs=merge_mode, outputs=[beta, alpha_label, beta_label, beta_apply_preset, beta_preset, beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks]) re_basin.change(fn=show_iters, inputs=re_basin,outputs=re_basin_iterations) - apply_preset.click(fn=load_presets,inputs=[preset, preset_lambda], outputs=[base,in_blocks,mid_block,out_blocks]) + apply_preset.click(fn=load_presets,inputs=[alpha_preset, alpha_preset_lambda], outputs=[alpha_base,alpha_in_blocks,alpha_mid_block,alpha_out_blocks,tabs]) + beta_apply_preset.click(fn=load_presets,inputs=[beta_preset, beta_preset_lambda], outputs=[beta_base,beta_in_blocks,beta_mid_block,beta_out_blocks,tabs]) MEHmodelmerger_merge.click( fn=wrap_gradio_gpu_call(MEHmodelmerger, extra_outputs=lambda: [gr.update() for _ in range(4)]), _js='modelmerger', @@ -262,16 +304,28 @@ def create_ui(): merge_mode, alpha, beta, + alpha_preset, + alpha_preset_lambda, + alpha_base, + alpha_in_blocks, + alpha_mid_block, + alpha_out_blocks, + beta_preset, + beta_preset_lambda, + beta_base, + beta_in_blocks, + beta_mid_block, + beta_out_blocks, precision, custom_name, checkpoint_format, save_metadata, - preset, weights_clip, prune, re_basin, re_basin_iterations, - device + device, + bake_in_vae, ], outputs=[ primary_model_name,