mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 09:38:23 +02:00
update advanced merging
This commit is contained in:
@@ -1 +0,0 @@
|
||||
4c7792ed011b233cdb6e9e42327085f4d66701f2
|
||||
@@ -27,6 +27,13 @@
|
||||
Change-in-behavior: new line in prompt now means *BREAK*
|
||||
- Add alternative Lora loading algorithm, triggered if `SD_LORA_DIFFUSERS` is set
|
||||
- Update to `diffusers==0.23.0`
|
||||
- **Model merge**
|
||||
- completely redesigned, now based on best-of-class `meh` by @s1dlx
|
||||
and heavily modified for additional functionality and fully integrated by @AI-Casanova (thanks!)
|
||||
- merge SD or SD-XL models using *simple merge* (12 methods),
|
||||
using one of *presets* (20 built-in presets) or custom block merge values
|
||||
- merge with ReBasin permuatations and/or clipping protection
|
||||
- fully multithreaded for fastest merge possible
|
||||
- **Extra networks**
|
||||
- Use multi-threading for 5x load speedup
|
||||
- Better Lora trigger words support
|
||||
|
||||
Submodule extensions-builtin/sd-webui-controlnet updated: c1f3d6f850...0efcf8c993
+36
-39
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import html
|
||||
import json
|
||||
import time
|
||||
import shutil
|
||||
|
||||
import torch
|
||||
@@ -54,6 +55,7 @@ def to_half(tensor, enable):
|
||||
|
||||
def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
shared.state.begin('merge')
|
||||
t0 = time.time()
|
||||
|
||||
def fail(message):
|
||||
shared.state.textinfo = message
|
||||
@@ -81,43 +83,39 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
kwargs["models"] |= {"model_c": sd_models.get_closet_checkpoint_match(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 KeyError as ke:
|
||||
shared.log.warn(f"Merge: Malformed manual block weight at {ke} falling back")
|
||||
if hasattr(kwargs, "alpha_base") and hasattr(kwargs, "alpha_in_blocks") and hasattr(kwargs, "alpha_mid_block") and hasattr(kwargs, "alpha_out_blocks"):
|
||||
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 KeyError as ke:
|
||||
shared.log.warning(f"Merge: Malformed manual block weight: {ke}")
|
||||
elif hasattr(kwargs, "alpha_preset") or hasattr(kwargs, "alpha"):
|
||||
kwargs["alpha"] = kwargs.get("alpha_preset", kwargs["alpha"])
|
||||
except AssertionError as e:
|
||||
shared.log.warn(f"Merge: {e}")
|
||||
kwargs["alpha"] = kwargs.get("alpha_preset", kwargs["alpha"])
|
||||
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):
|
||||
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 hasattr(kwargs, "beta_base") and hasattr(kwargs, "beta_in_blocks") and hasattr(kwargs, "beta_mid_block") and hasattr(kwargs, "beta_out_blocks"):
|
||||
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 KeyError as ke:
|
||||
shared.log.warn(f"Merge: Malformed manual block weight at {ke} falling back")
|
||||
kwargs["beta"] = kwargs.get("beta_preset", kwargs["beta"])
|
||||
except AssertionError as e:
|
||||
shared.log.warn(f"Merge: {e}")
|
||||
kwargs["beta"] = kwargs.get("beta_preset", kwargs["beta"])
|
||||
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)
|
||||
shared.log.warning(f"Merge: Malformed manual block weight: {ke}")
|
||||
elif hasattr(kwargs, "beta_preset") or hasattr(kwargs, "beta"):
|
||||
kwargs["beta"] = kwargs.get("beta_preset", kwargs["beta"])
|
||||
|
||||
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)
|
||||
|
||||
if kwargs["device"] == "gpu":
|
||||
kwargs["device"] = devices.device
|
||||
@@ -141,8 +139,8 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
|
||||
bake_in_vae_filename = sd_vae.vae_dict.get(kwargs.get("bake_in_vae", None), None)
|
||||
if bake_in_vae_filename is not None:
|
||||
shared.log.info(f"Merge: baking in VAE: {bake_in_vae_filename}")
|
||||
shared.state.textinfo = 'Baking in VAE'
|
||||
shared.log.info(f"Merge VAE='{bake_in_vae_filename}'")
|
||||
shared.state.textinfo = 'Merge VAE'
|
||||
vae_dict = sd_vae.load_vae_dict(bake_in_vae_filename)
|
||||
for key in vae_dict.keys():
|
||||
theta_0_key = 'first_stage_model.' + key
|
||||
@@ -154,7 +152,7 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
filename = kwargs.get("custom_name", "Unnamed_Merge")
|
||||
filename += "." + kwargs.get("checkpoint_format", None)
|
||||
output_modelname = os.path.join(ckpt_dir, filename)
|
||||
shared.state.textinfo = "Saving"
|
||||
shared.state.textinfo = "merge saving"
|
||||
metadata = None
|
||||
if kwargs.get("save_metadata", False):
|
||||
metadata = {"format": "pt", "sd_merge_models": {}}
|
||||
@@ -189,23 +187,22 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
|
||||
_, extension = os.path.splitext(output_modelname)
|
||||
|
||||
|
||||
if os.path.exists(output_modelname) and not kwargs.get("overwrite", False):
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Model alredy exists: {output_modelname}"]
|
||||
if extension.lower() == ".safetensors":
|
||||
safetensors.torch.save_file(theta_0, output_modelname, metadata=metadata)
|
||||
else:
|
||||
torch.save(theta_0, output_modelname)
|
||||
|
||||
t1 = time.time()
|
||||
shared.log.info(f"Merge complete: saved='{output_modelname}' time={t1-t0:.2f}")
|
||||
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()
|
||||
if kwargs["device"].type != "cpu":
|
||||
devices.torch_gc(force=True)
|
||||
shared.log.info(f"Merge saved: {output_modelname}.")
|
||||
shared.state.textinfo = "Checkpoint saved"
|
||||
devices.torch_gc(force=True)
|
||||
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)], f"Model saved to {output_modelname}"]
|
||||
|
||||
|
||||
def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_name, unet_conv, text_encoder_conv,
|
||||
|
||||
+30
-46
@@ -5,10 +5,9 @@ from typing import Dict, Optional, Tuple
|
||||
import safetensors.torch
|
||||
import torch
|
||||
from tensordict import TensorDict
|
||||
from tqdm import tqdm
|
||||
import modules.memstats
|
||||
import modules.devices as devices
|
||||
from modules.shared import log
|
||||
from modules.shared import log, console
|
||||
from modules.sd_models import read_state_dict
|
||||
from modules.merging import merge_methods
|
||||
from modules.merging.merge_utils import WeightClass
|
||||
@@ -69,7 +68,7 @@ def restore_sd_model(original_model: Dict, merged_model: Dict) -> Dict:
|
||||
|
||||
|
||||
def log_vram(txt=""):
|
||||
log.debug(f"{txt} VRAM: {modules.memstats.memory_stats()}")
|
||||
log.debug(f"Merge {txt}: {modules.memstats.memory_stats()}")
|
||||
|
||||
|
||||
def load_thetas(
|
||||
@@ -78,7 +77,6 @@ def load_thetas(
|
||||
device: torch.device,
|
||||
precision: str,
|
||||
) -> Dict:
|
||||
log_vram("before loading models")
|
||||
if prune:
|
||||
thetas = {k: prune_sd_model(TensorDict.from_dict(read_state_dict(m, "cpu"))) for k, m in models.items()}
|
||||
else:
|
||||
@@ -104,13 +102,11 @@ def merge_models(
|
||||
device: torch.device = None,
|
||||
work_device: torch.device = None,
|
||||
prune: bool = False,
|
||||
threads: int = 1,
|
||||
threads: int = 4,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
iterations = kwargs.get("re_basin_iterations", 1)
|
||||
thetas = load_thetas(models, prune, device, precision)
|
||||
|
||||
log.info(f"start merging with {merge_mode} method")
|
||||
# log.info(f'Merge start: models={models.values()} precision={precision} clip={weights_clip} rebasin={re_basin} prune={prune} threads={threads}')
|
||||
weight_matcher = WeightClass(thetas["model_a"], **kwargs)
|
||||
if re_basin:
|
||||
merged = rebasin_merge(
|
||||
@@ -119,7 +115,7 @@ def merge_models(
|
||||
merge_mode,
|
||||
precision=precision,
|
||||
weights_clip=weights_clip,
|
||||
iterations=iterations,
|
||||
iterations=kwargs.get("re_basin_iterations", 1),
|
||||
device=device,
|
||||
work_device=work_device,
|
||||
threads=threads,
|
||||
@@ -148,10 +144,9 @@ def un_prune_model(
|
||||
precision: str,
|
||||
) -> Dict:
|
||||
if prune:
|
||||
log.info("Un-pruning merged model")
|
||||
log.info("Merge restoring pruned keys")
|
||||
del thetas
|
||||
devices.torch_gc(force=True)
|
||||
log_vram("remove thetas")
|
||||
devices.torch_gc(force=False)
|
||||
original_a = TensorDict.from_dict(read_state_dict(models["model_a"], device))
|
||||
unpruned = 0
|
||||
for key in original_a.keys():
|
||||
@@ -163,10 +158,9 @@ def un_prune_model(
|
||||
if precision == "fp16":
|
||||
merged.update({key: merged[key].half()})
|
||||
if unpruned > 248: # VAE has 248 keys, and we are purposely restoring it here
|
||||
log.info(f"Merge: {unpruned - 248} unmerged keys restored from Primary Model")
|
||||
log.debug(f"Merge restored from primary model: keys={unpruned - 248}")
|
||||
unpruned = 0
|
||||
del original_a
|
||||
devices.torch_gc(force=True)
|
||||
original_b = TensorDict.from_dict(read_state_dict(models["model_b"], device))
|
||||
for key in original_b.keys():
|
||||
if KEY_POSITION_IDS in key:
|
||||
@@ -177,8 +171,9 @@ def un_prune_model(
|
||||
if precision == "fp16":
|
||||
merged.update({key: merged[key].half()})
|
||||
if unpruned != 0:
|
||||
log.info(f"Merge: {unpruned} unmerged keys restored from Secondary Model")
|
||||
log.debug(f"Merge restored from secondary model: keys={unpruned}")
|
||||
del original_b
|
||||
devices.torch_gc(force=False)
|
||||
|
||||
return fix_clip(merged)
|
||||
|
||||
@@ -191,15 +186,19 @@ def simple_merge(
|
||||
weights_clip: bool = False,
|
||||
device: torch.device = None,
|
||||
work_device: torch.device = None,
|
||||
threads: int = 1,
|
||||
threads: int = 4,
|
||||
) -> Dict:
|
||||
futures = []
|
||||
with tqdm(thetas["model_a"].keys(), desc="stage 1") as progress:
|
||||
# with tqdm(thetas["model_a"].keys(), desc="Merge") as progress:
|
||||
import rich.progress as p
|
||||
with p.Progress(p.TextColumn('[cyan]{task.description}'), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TextColumn('[cyan]keys={task.fields[keys]}'), console=console) as progress:
|
||||
task = progress.add_task(description="Merging", total=len(thetas["model_a"].keys()), keys=len(thetas["model_a"].keys()))
|
||||
with ThreadPoolExecutor(max_workers=threads) as executor:
|
||||
for key in thetas["model_a"].keys():
|
||||
future = executor.submit(
|
||||
simple_merge_key,
|
||||
progress,
|
||||
task,
|
||||
key,
|
||||
thetas,
|
||||
weight_matcher,
|
||||
@@ -214,17 +213,15 @@ def simple_merge(
|
||||
for res in futures:
|
||||
res.result()
|
||||
|
||||
log_vram("after stage 1")
|
||||
|
||||
for key in tqdm(thetas["model_b"].keys(), desc="stage 2"):
|
||||
if KEY_POSITION_IDS in key:
|
||||
continue
|
||||
if "model" in key and key not in thetas["model_a"].keys():
|
||||
thetas["model_a"].update({key: thetas["model_b"][key]})
|
||||
if precision == "fp16":
|
||||
thetas["model_a"].update({key: thetas["model_a"][key].half()})
|
||||
|
||||
log_vram("after stage 2")
|
||||
if len(thetas["model_b"]) > 0:
|
||||
log.debug(f'Merge update thetas: keys={len(thetas["model_b"])}')
|
||||
for key in thetas["model_b"].keys():
|
||||
if KEY_POSITION_IDS in key:
|
||||
continue
|
||||
if "model" in key and key not in thetas["model_a"].keys():
|
||||
thetas["model_a"].update({key: thetas["model_b"][key]})
|
||||
if precision == "fp16":
|
||||
thetas["model_a"].update({key: thetas["model_a"][key].half()})
|
||||
|
||||
return fix_clip(thetas["model_a"])
|
||||
|
||||
@@ -240,13 +237,12 @@ def rebasin_merge(
|
||||
work_device: torch.device = None,
|
||||
threads: int = 1,
|
||||
):
|
||||
# WARNING: not sure how this does when 3 models are involved...
|
||||
|
||||
# not sure how this does when 3 models are involved...
|
||||
model_a = thetas["model_a"].clone()
|
||||
perm_spec = sdunet_permutation_spec()
|
||||
|
||||
for it in range(iterations):
|
||||
log_vram(f"Rebasin iteration {it}")
|
||||
log_vram(f"rebasin: iteration={it}")
|
||||
weight_matcher.set_it(it)
|
||||
|
||||
# normal block merge we already know and love
|
||||
@@ -261,8 +257,6 @@ def rebasin_merge(
|
||||
threads,
|
||||
)
|
||||
|
||||
log_vram("simple merge done")
|
||||
|
||||
# find permutations
|
||||
perm_1, y = weight_matching(
|
||||
perm_spec,
|
||||
@@ -273,13 +267,8 @@ def rebasin_merge(
|
||||
usefp16=precision == "fp16",
|
||||
device=device,
|
||||
)
|
||||
|
||||
log_vram("weight matching #1 done")
|
||||
|
||||
thetas["model_a"] = apply_permutation(perm_spec, perm_1, thetas["model_a"])
|
||||
|
||||
log_vram("apply perm 1 done")
|
||||
|
||||
perm_2, z = weight_matching(
|
||||
perm_spec,
|
||||
thetas["model_b"],
|
||||
@@ -290,8 +279,6 @@ def rebasin_merge(
|
||||
device=device,
|
||||
)
|
||||
|
||||
log_vram("weight matching #2 done")
|
||||
|
||||
new_alpha = torch.nn.functional.normalize(
|
||||
torch.sigmoid(torch.Tensor([y, z])), p=1, dim=0
|
||||
).tolist()[0]
|
||||
@@ -299,8 +286,6 @@ def rebasin_merge(
|
||||
perm_spec, perm_2, thetas["model_a"], new_alpha
|
||||
)
|
||||
|
||||
log_vram("model a updated")
|
||||
|
||||
if weights_clip:
|
||||
clip_thetas = thetas.copy()
|
||||
clip_thetas["model_a"] = model_a
|
||||
@@ -309,12 +294,11 @@ def rebasin_merge(
|
||||
return thetas["model_a"]
|
||||
|
||||
|
||||
def simple_merge_key(progress, key, thetas, *args, **kwargs):
|
||||
def simple_merge_key(progress, task, key, thetas, *args, **kwargs):
|
||||
with merge_key_context(key, thetas, *args, **kwargs) as result:
|
||||
if result is not None:
|
||||
thetas["model_a"].update({key: result.detach().clone()})
|
||||
|
||||
progress.update()
|
||||
progress.update(task, advance=1)
|
||||
|
||||
|
||||
def merge_key( # pylint: disable=inconsistent-return-statements
|
||||
@@ -407,7 +391,7 @@ def get_merge_method_args(
|
||||
|
||||
|
||||
def save_model(model, output_file, file_format) -> None:
|
||||
log.info(f"Saving {output_file}")
|
||||
log.info(f"Merge saving: model='{output_file}'")
|
||||
if file_format == "safetensors":
|
||||
safetensors.torch.save_file(
|
||||
model if type(model) == dict else model.to_dict(),
|
||||
|
||||
@@ -2228,7 +2228,7 @@ def inner_matching(
|
||||
if newL - oldL != 0:
|
||||
linear_sum += abs((newL - oldL).item())
|
||||
number += 1
|
||||
log.info(f" permutation {p}: {newL - oldL}")
|
||||
log.debug(f"Merge Rebasin permutation: {p}={newL-oldL}")
|
||||
|
||||
progress = progress or newL > oldL + 1e-12
|
||||
|
||||
@@ -2262,12 +2262,11 @@ def weight_matching(
|
||||
number = 0
|
||||
|
||||
special_layers = ["P_bg324", "P_bg358", "P_bg337"]
|
||||
for _ in range(max_iter):
|
||||
for _i in range(max_iter):
|
||||
progress = False
|
||||
shuffle(special_layers)
|
||||
for p in special_layers:
|
||||
n = perm_sizes[p]
|
||||
|
||||
linear_sum, number, perm, progress = inner_matching(
|
||||
n,
|
||||
ps,
|
||||
@@ -2281,6 +2280,7 @@ def weight_matching(
|
||||
perm,
|
||||
device,
|
||||
)
|
||||
progress = True
|
||||
if not progress:
|
||||
break
|
||||
|
||||
|
||||
@@ -387,7 +387,7 @@ def read_state_dict(checkpoint_file, map_location=None): # pylint: disable=unuse
|
||||
return None
|
||||
try:
|
||||
pl_sd = None
|
||||
with progress.open(checkpoint_file, 'rb', description=f'[cyan]Loading weights: [yellow]{checkpoint_file}', auto_refresh=True, console=shared.console) as f:
|
||||
with progress.open(checkpoint_file, 'rb', description=f'[cyan]Loading model: [yellow]{checkpoint_file}', auto_refresh=True, console=shared.console) as f:
|
||||
_, extension = os.path.splitext(checkpoint_file)
|
||||
if extension.lower() == ".ckpt" and shared.opts.sd_disable_ckpt:
|
||||
shared.log.warning(f"Checkpoint loading disabled: {checkpoint_file}")
|
||||
|
||||
+57
-45
@@ -38,7 +38,7 @@ def create_ui():
|
||||
create_refresh_button(model_name, sd_models.list_models,
|
||||
lambda: {"choices": sd_models.checkpoint_tiles()}, "refresh_checkpoint_Z")
|
||||
with gr.Row():
|
||||
custom_name = gr.Textbox(label="New model name")
|
||||
custom_name = gr.Textbox(label="Output model name")
|
||||
with gr.Row():
|
||||
precision = gr.Radio(choices=["fp32", "fp16", "bf16"], value="fp16", label="Model precision")
|
||||
m_type = gr.Radio(choices=["disabled", "no-ema", "ema-only"], value="disabled",
|
||||
@@ -149,10 +149,14 @@ def create_ui():
|
||||
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", value=True, visible=False)
|
||||
re_basin = gr.Checkbox(label="ReBasin")
|
||||
overwrite = gr.Checkbox(label="Overwrite model")
|
||||
with FormRow():
|
||||
save_metadata = gr.Checkbox(value=True, label="Save metadata")
|
||||
with FormRow():
|
||||
weights_clip = gr.Checkbox(label="Weights clip")
|
||||
prune = gr.Checkbox(label="Prune", value=True, visible=False)
|
||||
with FormRow():
|
||||
re_basin = gr.Checkbox(label="ReBasin")
|
||||
re_basin_iterations = gr.Slider(minimum=0, maximum=25, step=1,
|
||||
label='Number of ReBasin Iterations', value=None,
|
||||
visible=False)
|
||||
@@ -170,55 +174,62 @@ def create_ui():
|
||||
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():
|
||||
modelmerger_merge = gr.Button(value="Merge", variant='primary')
|
||||
|
||||
def modelmerger(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,
|
||||
unload,
|
||||
bake_in_vae):
|
||||
def modelmerger(dummy_component, # dummy function just to get argspec later
|
||||
overwrite, # pylint: disable=unused-argument
|
||||
primary_model_name, # pylint: disable=unused-argument
|
||||
secondary_model_name, # pylint: disable=unused-argument
|
||||
tertiary_model_name, # pylint: disable=unused-argument
|
||||
merge_mode, # pylint: disable=unused-argument
|
||||
alpha, # pylint: disable=unused-argument
|
||||
beta, # pylint: disable=unused-argument
|
||||
alpha_preset, # pylint: disable=unused-argument
|
||||
alpha_preset_lambda, # pylint: disable=unused-argument
|
||||
alpha_base, # pylint: disable=unused-argument
|
||||
alpha_in_blocks, # pylint: disable=unused-argument
|
||||
alpha_mid_block, # pylint: disable=unused-argument
|
||||
alpha_out_blocks, # pylint: disable=unused-argument
|
||||
beta_preset, # pylint: disable=unused-argument
|
||||
beta_preset_lambda, # pylint: disable=unused-argument
|
||||
beta_base, # pylint: disable=unused-argument
|
||||
beta_in_blocks, # pylint: disable=unused-argument
|
||||
beta_mid_block, # pylint: disable=unused-argument
|
||||
beta_out_blocks, # pylint: disable=unused-argument
|
||||
precision, # pylint: disable=unused-argument
|
||||
custom_name, # pylint: disable=unused-argument
|
||||
checkpoint_format, # pylint: disable=unused-argument
|
||||
save_metadata, # pylint: disable=unused-argument
|
||||
weights_clip, # pylint: disable=unused-argument
|
||||
prune, # pylint: disable=unused-argument
|
||||
re_basin, # pylint: disable=unused-argument
|
||||
re_basin_iterations, # pylint: disable=unused-argument
|
||||
device, # pylint: disable=unused-argument
|
||||
unload, # pylint: disable=unused-argument
|
||||
bake_in_vae): # pylint: disable=unused-argument
|
||||
kwargs = {}
|
||||
for x in inspect.getfullargspec(modelmerger)[0]:
|
||||
kwargs[x] = locals()[x]
|
||||
for key in list(kwargs.keys()):
|
||||
if kwargs[key] in [None, "None", "", 0, []]:
|
||||
del kwargs[key]
|
||||
try:
|
||||
results = extras.run_modelmerger(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
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)],
|
||||
f"Error merging checkpoints: {e}"]
|
||||
return results
|
||||
del kwargs['dummy_component']
|
||||
if kwargs.get("custom_name", None) is None:
|
||||
log.error('Merge: no output model specified')
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], "No output model specified"]
|
||||
elif kwargs.get("primary_model_name", None) is None or kwargs.get("secondary_model_name", None) is None:
|
||||
log.error('Merge: no models selected')
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], "No models selected"]
|
||||
else:
|
||||
log.debug(f'Merge start: {kwargs}')
|
||||
try:
|
||||
results = extras.run_modelmerger(dummy_component, **kwargs)
|
||||
except Exception as e:
|
||||
modules.errors.display(e, '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:
|
||||
@@ -262,7 +273,7 @@ def create_ui():
|
||||
preset = interpolate(presets, ratio)
|
||||
else:
|
||||
preset = presets[0]
|
||||
preset = ['%.3f' % x if int(x) != x else str(x) for x in preset]
|
||||
preset = ['%.3f' % x if int(x) != x else str(x) for x in preset] # pylint: disable=consider-using-f-string
|
||||
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)]
|
||||
|
||||
@@ -291,6 +302,7 @@ def create_ui():
|
||||
_js='modelmerger',
|
||||
inputs=[
|
||||
dummy_component,
|
||||
overwrite,
|
||||
primary_model_name,
|
||||
secondary_model_name,
|
||||
tertiary_model_name,
|
||||
|
||||
Reference in New Issue
Block a user