update advanced merging

This commit is contained in:
Vladimir Mandic
2023-11-20 11:43:30 -05:00
parent 564b67bcd8
commit 00b246c052
8 changed files with 135 additions and 136 deletions
-1
View File
@@ -1 +0,0 @@
4c7792ed011b233cdb6e9e42327085f4d66701f2
+7
View File
@@ -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
+36 -39
View File
@@ -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
View File
@@ -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(),
+3 -3
View File
@@ -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
+1 -1
View File
@@ -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
View File
@@ -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,