Added Tabs

This commit is contained in:
AI-Casanova
2023-11-05 07:10:05 -06:00
parent 39a37597b5
commit 685f392967
2 changed files with 155 additions and 96 deletions
+61 -56
View File
@@ -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):
+94 -40
View File
@@ -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,