Further UI Work

This commit is contained in:
AI-Casanova
2023-11-03 14:55:57 -05:00
parent 6230076780
commit 39a37597b5
3 changed files with 254 additions and 3 deletions
+106
View File
@@ -8,6 +8,7 @@ import torch
import tqdm
import gradio as gr
import safetensors.torch
from sd_meh.merge import merge_models
from modules import shared, images, sd_models, sd_vae, sd_models_config
@@ -239,6 +240,111 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_
shared.state.end()
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
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"
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,
)
ckpt_dir = shared.opts.ckpt_dir or sd_models.model_path
filename = custom_name
filename += "." + checkpoint_format
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
"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,
"interp_method": interp_method,
"multiplier": multiplier,
"save_as_half": save_as_half,
"custom_name": custom_name,
}
metadata["sd_merge_recipe"] = json.dumps(merge_recipe)
def add_model_metadata(checkpoint_info):
checkpoint_info.calculate_shorthash()
metadata["sd_merge_models"][checkpoint_info.sha256] = {
"name": checkpoint_info.name,
"legacy_hash": checkpoint_info.hash,
"sd_merge_recipe": checkpoint_info.metadata.get("sd_merge_recipe", None)
}
metadata["sd_merge_models"].update(checkpoint_info.metadata.get("sd_merge_models", {}))
add_model_metadata(primary_model_info)
if secondary_model_info:
add_model_metadata(secondary_model_info)
if tertiary_model_info:
add_model_metadata(tertiary_model_info)
metadata["sd_merge_models"] = json.dumps(metadata["sd_merge_models"])
_, extension = os.path.splitext(output_modelname)
if extension.lower() == ".safetensors":
safetensors.torch.save_file(theta_0, output_modelname, metadata=metadata)
else:
torch.save(theta_0, output_modelname)
sd_models.list_models()
created_model = next((ckpt for ckpt in sd_models.checkpoints_list.values() if ckpt.name == filename), None)
if created_model:
created_model.calculate_shorthash()
create_config(output_modelname, config_source, primary_model_info, secondary_model_info, tertiary_model_info)
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]
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
+2 -2
View File
@@ -231,7 +231,7 @@ def create_output_panel(tabname):
return result_gallery, generation_info, html_info, html_info_formatted, html_log
def create_refresh_button(refresh_component, refresh_method, refreshed_args, elem_id):
def create_refresh_button(refresh_component, refresh_method, refreshed_args, elem_id, visible: bool = True):
def refresh():
refresh_method()
@@ -241,7 +241,7 @@ def create_refresh_button(refresh_component, refresh_method, refreshed_args, ele
return gr.update(**(args or {}))
from modules.ui_components import ToolButton
refresh_button = ToolButton(value=symbols.refresh, elem_id=elem_id)
refresh_button = ToolButton(value=symbols.refresh, elem_id=elem_id, visible=visible)
refresh_button.click(fn=refresh, inputs=[], outputs=[refresh_component])
return refresh_button
+146 -1
View File
@@ -3,12 +3,16 @@ import json
from datetime import datetime
import gradio as gr
from modules import sd_models, sd_vae, extras
from modules.ui_components import FormRow, ToolButton
from modules.ui_components import FormRow, ToolButton, InputAccordion
from modules.ui_common import create_refresh_button
from modules.call_queue import wrap_gradio_gpu_call
from modules.shared import opts, log, req
import modules.errors
import modules.hashes
from sd_meh import merge_methods
from sd_meh.utils import BETA_METHODS, TRIPLE_METHODS, interpolate
from sd_meh.presets import BLOCK_WEIGHTS_PRESETS
search_metadata_civit = None
@@ -136,6 +140,147 @@ def create_ui():
models_outcome,
]
)
with gr.Tab(label="MEH Merge"):
def sd_model_choices():
return ['None'] + sd_models.checkpoint_tiles()
with gr.Row(equal_height=False):
with gr.Column(variant='compact'):
with FormRow():
custom_name = gr.Textbox(label="New model name")
with FormRow():
merge_mode = gr.Dropdown(choices=merge_methods.__all__, value="weighted_sum", label="Interpolation Method")
with FormRow():
primary_model_name = gr.Dropdown(sd_model_choices(), label="Primary model", value="None")
create_refresh_button(primary_model_name, sd_models.list_models, lambda: {"choices": sd_model_choices()}, "refresh_checkpoint_A")
secondary_model_name = gr.Dropdown(sd_model_choices(), label="Secondary model", value="None")
create_refresh_button(secondary_model_name, sd_models.list_models, lambda: {"choices": sd_model_choices()}, "refresh_checkpoint_B")
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 FormRow():
weights_clip = gr.Checkbox(label="Weights Clip")
prune = gr.Checkbox(label="Prune")
re_basin = gr.Checkbox(label="ReBasin")
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")
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")
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):
try:
results = extras.run_MEHmodelmerger(*args)
except Exception as e:
modules.errors.display(e, 'model merge')
sd_models.list_models() # to remove the potentially missing models from the list
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Error merging checkpoints: {e}"]
return results
def tertiary(mode):
if mode in TRIPLE_METHODS:
return [gr.update(visible=True) for _ in range(2)]
else:
return [gr.update(visible=False) for _ in range(2)]
def beta_visibility(mode):
if mode in BETA_METHODS:
return [gr.update(visible=True) for _ in range(9)]
else:
return [gr.update(visible=False) for _ in range(9)]
def show_iters(show):
if show:
return gr.Slider.update(value=5, visible=True)
else:
return gr.Slider.update(value=None, visible=False)
def preset_visiblility(x):
if len(x) == 2:
return gr.Slider.update(value=0.5, visible=True)
else:
return gr.Slider.update(value=None, visible=False)
def load_presets(presets,ratio):
for i, p in enumerate(presets):
presets[i] = BLOCK_WEIGHTS_PRESETS[p]
if len(presets) == 2:
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.change(fn=preset_visiblility, inputs=preset, outputs=preset_lambda)
beta_preset.change(fn=preset_visiblility, inputs=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])
MEHmodelmerger_merge.click(
fn=wrap_gradio_gpu_call(MEHmodelmerger, extra_outputs=lambda: [gr.update() for _ in range(4)]),
_js='modelmerger',
inputs=[
dummy_component,
primary_model_name,
secondary_model_name,
tertiary_model_name,
merge_mode,
alpha,
beta,
precision,
custom_name,
checkpoint_format,
save_metadata,
preset,
weights_clip,
prune,
re_basin,
re_basin_iterations,
device
],
outputs=[
primary_model_name,
secondary_model_name,
tertiary_model_name,
dummy_component,
models_outcome,
]
)
with gr.Tab(label="Validate"):
model_headers = ['name', 'type', 'filename', 'hash', 'added', 'size', 'metadata']