model component merge

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-01-19 14:36:09 -05:00
parent de1063275c
commit 311d402b0c
11 changed files with 833 additions and 245 deletions
+12 -1
View File
@@ -2,15 +2,26 @@
## Update for 2025-01-19
- **Model Merge**
- replace model components and merge LoRAs
in addition to existing model weights merge support
now also having ability to replace model components and merge LoRAs
you can also test merges in-memory without needing to save to disk at all
and you can also use it to convert diffusers to safetensors if you want
*example*: replace vae in your favorite model with a fixed one? replace text encoder? etc.
*note*: limited to sdxl for now, additional models can be added depending on popularity
- **Detailer**:
- in addition as standard behavior of detect & run-generate, it can now also run face-restore models
- included models are: *CodeFormer, RestoreFormer, GFPGan, GPEN-BFR*
- **Other**:
- **ipex**: update supported torch versions
- **gallery**: add http fallback for slow/unreliable links
- **upscale**: code refactor to unify latent, resize and model based upscalers
- **splash**: add legacy mode indicator on splash screen
- **network**: extract thumbnail from model metadata if present
- **Refactor**:
- **upscale**: code refactor to unify latent, resize and model based upscalers
- **loader**: ability to run in-memory models
- **schedulers**: ability to create model-less schedulers
- **Fixes**:
- non-full vae decode
- send-to image transfer
+1 -1
View File
@@ -77,7 +77,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
#extensions .info { margin: 0; }
#extensions .date { opacity: 0.85; font-size: 90%; }
/* extra networks */
/* networks */
.extra-networks > div { margin: 0; border-bottom: none !important; }
.extra-networks .second-line { display: flex; width: -moz-available; width: -webkit-fill-available; gap: 0.3em; box-shadow: var(--input-shadow); margin-bottom: 2px; }
.extra-networks .search { flex: 1; }
+5 -2
View File
@@ -207,7 +207,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
#extensions .info { margin: 0; }
#extensions .date { opacity: 0.85; font-size: var(--text-sm); }
/* extra networks */
/* networks */
.extra_networks_root { width: 0; position: absolute; height: auto; right: 0; top: 13em; z-index: 100; } /* default is sidebar view */
.extra-networks { background: var(--background-color); padding: var(--block-label-padding); }
.extra-networks > div { margin: 0; border-bottom: none !important; gap: 0.3em 0; }
@@ -269,11 +269,14 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
.ar-dropdown div { margin: 0; background: var(--background-color)}
#txt2img_sampler_timesteps, #img2img_sampler_timesteps { max-width: calc(var(--left-column) - 50px); }
/* extras */
/* models */
.extras { gap: 0.2em 1em !important }
#extras_generate, #extras_interrupt, #extras_skip { display: block !important; position: relative; height: 36px; }
#extras_upscale { margin-top: 10px }
#pnginfo_html_info .gradio-html > div { margin: 0.5em; }
#models_image, #models_image > div { min-height: 0; }
#models_error { font-family: monospace; color: var(--body-text-color-subdued) }
/* log monitor */
.log-monitor { display: none; justify-content: unset !important; overflow: hidden; padding: 0; margin-top: auto; font-family: monospace; font-size: var(--text-xxs); }
+75 -141
View File
@@ -4,14 +4,12 @@ import json
import time
import shutil
from PIL import Image
import torch
import tqdm
import gradio as gr
import safetensors.torch
from modules.merging.merge import merge_models
from modules.merging.merge_utils import TRIPLE_METHODS
from modules import shared, images, sd_models, sd_vae, sd_models_config, devices
from modules.merging import merge, merge_utils, modules_sdxl
from modules import shared, images, sd_models, sd_vae, sd_samplers, sd_models_config, devices
def run_pnginfo(image):
@@ -73,9 +71,9 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
if kwargs.get("secondary_model_name", None) in [None, 'None']:
return fail("Failed: Merging requires a secondary model.")
secondary_model_info = sd_models.get_closet_checkpoint_match(kwargs.get("secondary_model_name", None))
if kwargs.get("tertiary_model_name", None) in [None, 'None'] and kwargs.get("merge_mode", None) in TRIPLE_METHODS:
if kwargs.get("tertiary_model_name", None) in [None, 'None'] and kwargs.get("merge_mode", None) in merge_utils.TRIPLE_METHODS:
return fail(f"Failed: Interpolation method ({kwargs.get('merge_mode', None)}) requires a tertiary model.")
tertiary_model_info = sd_models.get_closet_checkpoint_match(kwargs.get("tertiary_model_name", None)) if kwargs.get("merge_mode", None) in TRIPLE_METHODS else None
tertiary_model_info = sd_models.get_closet_checkpoint_match(kwargs.get("tertiary_model_name", None)) if kwargs.get("merge_mode", None) in merge_utils.TRIPLE_METHODS else None
del kwargs["primary_model_name"]
del kwargs["secondary_model_name"]
@@ -128,7 +126,7 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
sd_models.unload_model_weights()
try:
theta_0 = merge_models(**kwargs)
theta_0 = merge.merge_models(**kwargs)
except Exception as e:
return fail(f"{e}")
@@ -205,144 +203,80 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) 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,
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}
def run_model_modules(model_type:str, model_name:str, custom_name:str,
comp_unet:str, comp_vae:str, comp_te1:str, comp_te2:str,
precision:str, comp_scheduler:str, comp_prediction:str,
comp_lora:str, comp_fuse:float,
meta_author:str, meta_version:str, meta_license:str, meta_desc:str, meta_hint:str, meta_thumbnail:Image.Image,
create_diffusers:bool, create_safetensors:bool, debug:bool):
def conv_fp16(t: torch.Tensor):
return t.half() if t.dtype in dtypes_to_fp16 else t
def conv_bf16(t: torch.Tensor):
return t.bfloat16() if t.dtype in dtypes_to_bf16 else t
def conv_full(t):
return t
_g_precision_func = {
"full": conv_full,
"fp32": conv_full,
"fp16": conv_fp16,
"bf16": conv_bf16,
}
def check_weight_type(k: str) -> str:
if k.startswith("model.diffusion_model"):
return "unet"
elif k.startswith("first_stage_model"):
return "vae"
elif k.startswith("cond_stage_model"):
return "clip"
return "other"
def load_model(path):
if path.endswith(".safetensors"):
m = safetensors.torch.load_file(path, device="cpu")
status = ''
def msg(text, err:bool=False):
nonlocal status
if err:
shared.log.error(f'Modules merge: {text}')
else:
m = torch.load(path, map_location="cpu")
state_dict = m["state_dict"] if "state_dict" in m else m
return state_dict
shared.log.info(f'Modules merge: {text}')
status += text + '<br>'
return status
def fix_model(model, fix_clip=False):
# code from model-toolkit
nai_keys = {
'cond_stage_model.transformer.embeddings.': 'cond_stage_model.transformer.text_model.embeddings.',
'cond_stage_model.transformer.encoder.': 'cond_stage_model.transformer.text_model.encoder.',
'cond_stage_model.transformer.final_layer_norm.': 'cond_stage_model.transformer.text_model.final_layer_norm.'
}
for k in list(model.keys()):
for r in nai_keys:
if type(k) == str and k.startswith(r):
new_key = k.replace(r, nai_keys[r])
model[new_key] = model[k]
del model[k]
shared.log.warning(f"Model convert: fixed NovelAI error key: {k}")
break
if fix_clip:
i = "cond_stage_model.transformer.text_model.embeddings.position_ids"
if i in model:
correct = torch.Tensor([list(range(77))]).to(torch.int64)
now = model[i].to(torch.int64)
if model_type != 'sdxl':
yield msg("only SDXL models are supported", err=True)
return
if len(custom_name) == 0:
yield msg("output name is required", err=True)
return
checkpoint_info = sd_models.get_closet_checkpoint_match(model_name)
if checkpoint_info is None:
yield msg("input model not found", err=True)
return
fn = checkpoint_info.filename
shared.state.begin('Merge')
yield msg("modules merge starting")
yield msg("unload current model")
sd_models.unload_model_weights(op='model')
broken = correct.ne(now)
broken = [i for i in range(77) if broken[0][i]]
model[i] = correct
if len(broken) != 0:
shared.log.warning(f"Model convert: fixed broken CLiP: {broken}")
modules_sdxl.recipe.name = custom_name
modules_sdxl.recipe.author = meta_author
modules_sdxl.recipe.version = meta_version
modules_sdxl.recipe.desc = meta_desc
modules_sdxl.recipe.hint = meta_hint
modules_sdxl.recipe.license = meta_license
modules_sdxl.recipe.thumbnail = meta_thumbnail
modules_sdxl.recipe.base = fn
modules_sdxl.recipe.unet = comp_unet
modules_sdxl.recipe.vae = comp_vae
modules_sdxl.recipe.te1 = comp_te1
modules_sdxl.recipe.te2 = comp_te2
modules_sdxl.recipe.prediction = comp_prediction
modules_sdxl.recipe.diffusers = create_diffusers
modules_sdxl.recipe.safetensors = create_safetensors
modules_sdxl.recipe.fuse = float(comp_fuse)
modules_sdxl.recipe.debug = debug
return model
if model == "":
return "Error: you must choose a model"
if len(checkpoint_formats) == 0:
return "Error: at least choose one model save format"
extra_opt = {
"unet": unet_conv,
"clip": text_encoder_conv,
"vae": vae_conv,
"other": others_conv
}
shared.state.begin('Convert')
model_info = sd_models.checkpoints_list[model]
shared.state.textinfo = f"Load {model_info.filename}..."
shared.log.info(f"Model convert loading: {model_info.filename}")
state_dict = load_model(model_info.filename)
ok = {} # {"state_dict": {}}
conv_func = _g_precision_func[precision]
def _hf(wk: str, t: torch.Tensor):
if not isinstance(t, torch.Tensor):
return
w_t = check_weight_type(wk)
conv_t = extra_opt[w_t]
if conv_t == "convert":
ok[wk] = conv_func(t)
elif conv_t == "copy":
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):
ema_k = "___"
try:
ema_k = "model_ema." + k[6:].replace(".", "")
except Exception:
pass
if ema_k in state_dict:
_hf(k, state_dict[ema_k])
elif not k.startswith("model_ema.") or k in ["model_ema.num_updates", "model_ema.decay"]:
_hf(k, state_dict[k])
elif conv_type == "no-ema":
for k, v in tqdm.tqdm(state_dict.items()):
if "model_ema." not in k:
_hf(k, v)
loras = [l.strip() if ':' in l else f'{l.strip()}:1.0' for l in comp_lora.split(',') if len(l.strip()) > 0]
for lora, strength in [l.split(':') for l in loras]:
modules_sdxl.recipe.lora[lora] = float(strength)
scheduler = sd_samplers.create_sampler(comp_scheduler, None)
modules_sdxl.recipe.scheduler = scheduler.__class__.__name__ if scheduler is not None else None
if precision == 'fp32':
modules_sdxl.recipe.precision = torch.float32
elif precision == 'bf16':
modules_sdxl.recipe.precision = torch.bfloat16
else:
for k, v in tqdm.tqdm(state_dict.items()):
_hf(k, v)
modules_sdxl.recipe.precision = torch.float16
ok = fix_model(ok, fix_clip=fix_clip)
output = ""
ckpt_dir = shared.cmd_opts.ckpt_dir or sd_models.model_path
save_name = f"{model_info.model_name}-{precision}"
if conv_type != "disabled":
save_name += f"-{conv_type}"
if custom_name != "":
save_name = custom_name
for fmt in checkpoint_formats:
ext = ".safetensors" if fmt == "safetensors" else ".ckpt"
_save_name = save_name + ext
save_path = os.path.join(ckpt_dir, _save_name)
shared.log.info(f"Model convert saving: {save_path}")
if fmt == "safetensors":
safetensors.torch.save_file(ok, save_path)
else:
torch.save({"state_dict": ok}, save_path)
output += f"Checkpoint saved to {save_path}<br>"
modules_sdxl.status = status
yield from modules_sdxl.merge()
status = modules_sdxl.status
devices.torch_gc(force=True)
yield msg("modules merge complete")
if modules_sdxl.pipeline is not None:
checkpoint_info = sd_models.CheckpointInfo(filename='None')
shared.sd_model = modules_sdxl.pipeline
sd_models.set_defaults(shared.sd_model, checkpoint_info)
sd_models.set_diffuser_options(shared.sd_model, offload=False)
sd_models.set_diffuser_offload(shared.sd_model)
yield msg("pipeline loaded")
shared.state.end()
return output
+297
View File
@@ -0,0 +1,297 @@
import io
import os
import re
import hashlib
import torch
from safetensors.torch import load_file, save_file
unet_conversion_map = [
# (stable-diffusion, HF Diffusers)
("time_embed.0.weight", "time_embedding.linear_1.weight"),
("time_embed.0.bias", "time_embedding.linear_1.bias"),
("time_embed.2.weight", "time_embedding.linear_2.weight"),
("time_embed.2.bias", "time_embedding.linear_2.bias"),
("input_blocks.0.0.weight", "conv_in.weight"),
("input_blocks.0.0.bias", "conv_in.bias"),
("out.0.weight", "conv_norm_out.weight"),
("out.0.bias", "conv_norm_out.bias"),
("out.2.weight", "conv_out.weight"),
("out.2.bias", "conv_out.bias"),
# the following are for sdxl
("label_emb.0.0.weight", "add_embedding.linear_1.weight"),
("label_emb.0.0.bias", "add_embedding.linear_1.bias"),
("label_emb.0.2.weight", "add_embedding.linear_2.weight"),
("label_emb.0.2.bias", "add_embedding.linear_2.bias"),
]
unet_conversion_map_resnet = [
# (stable-diffusion, HF Diffusers)
("in_layers.0", "norm1"),
("in_layers.2", "conv1"),
("out_layers.0", "norm2"),
("out_layers.3", "conv2"),
("emb_layers.1", "time_emb_proj"),
("skip_connection", "conv_shortcut"),
]
unet_conversion_map_layer = []
# hardcoded number of downblocks and resnets/attentions...
# would need smarter logic for other networks.
for i in range(3):
# loop over downblocks/upblocks
for j in range(2):
# loop over resnets/attentions for downblocks
hf_down_res_prefix = f"down_blocks.{i}.resnets.{j}."
sd_down_res_prefix = f"input_blocks.{3*i + j + 1}.0."
unet_conversion_map_layer.append((sd_down_res_prefix, hf_down_res_prefix))
if i > 0:
hf_down_atn_prefix = f"down_blocks.{i}.attentions.{j}."
sd_down_atn_prefix = f"input_blocks.{3*i + j + 1}.1."
unet_conversion_map_layer.append((sd_down_atn_prefix, hf_down_atn_prefix))
for j in range(4):
# loop over resnets/attentions for upblocks
hf_up_res_prefix = f"up_blocks.{i}.resnets.{j}."
sd_up_res_prefix = f"output_blocks.{3*i + j}.0."
unet_conversion_map_layer.append((sd_up_res_prefix, hf_up_res_prefix))
if i < 2:
# no attention layers in up_blocks.0
hf_up_atn_prefix = f"up_blocks.{i}.attentions.{j}."
sd_up_atn_prefix = f"output_blocks.{3 * i + j}.1."
unet_conversion_map_layer.append((sd_up_atn_prefix, hf_up_atn_prefix))
if i < 3:
# no downsample in down_blocks.3
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0.conv."
sd_downsample_prefix = f"input_blocks.{3*(i+1)}.0.op."
unet_conversion_map_layer.append((sd_downsample_prefix, hf_downsample_prefix))
# no upsample in up_blocks.3
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
sd_upsample_prefix = f"output_blocks.{3*i + 2}.{1 if i == 0 else 2}."
unet_conversion_map_layer.append((sd_upsample_prefix, hf_upsample_prefix))
unet_conversion_map_layer.append(("output_blocks.2.2.conv.", "output_blocks.2.1.conv."))
hf_mid_atn_prefix = "mid_block.attentions.0."
sd_mid_atn_prefix = "middle_block.1."
unet_conversion_map_layer.append((sd_mid_atn_prefix, hf_mid_atn_prefix))
for j in range(2):
hf_mid_res_prefix = f"mid_block.resnets.{j}."
sd_mid_res_prefix = f"middle_block.{2*j}."
unet_conversion_map_layer.append((sd_mid_res_prefix, hf_mid_res_prefix))
def convert_unet_state_dict(unet_state_dict):
# buyer beware: this is a *brittle* function,
# and correct output requires that all of these pieces interact in
# the exact order in which I have arranged them.
mapping = {k: k for k in unet_state_dict.keys()}
for sd_name, hf_name in unet_conversion_map:
mapping[hf_name] = sd_name
for k, v in mapping.items():
if "resnets" in k:
for sd_part, hf_part in unet_conversion_map_resnet:
v = v.replace(hf_part, sd_part)
mapping[k] = v
for k, v in mapping.items():
for sd_part, hf_part in unet_conversion_map_layer:
v = v.replace(hf_part, sd_part)
mapping[k] = v
new_state_dict = {sd_name: unet_state_dict[hf_name] for hf_name, sd_name in mapping.items()}
return new_state_dict
vae_conversion_map = [
# (stable-diffusion, HF Diffusers)
("nin_shortcut", "conv_shortcut"),
("norm_out", "conv_norm_out"),
("mid.attn_1.", "mid_block.attentions.0."),
]
for i in range(4):
# down_blocks have two resnets
for j in range(2):
hf_down_prefix = f"encoder.down_blocks.{i}.resnets.{j}."
sd_down_prefix = f"encoder.down.{i}.block.{j}."
vae_conversion_map.append((sd_down_prefix, hf_down_prefix))
if i < 3:
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0."
sd_downsample_prefix = f"down.{i}.downsample."
vae_conversion_map.append((sd_downsample_prefix, hf_downsample_prefix))
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
sd_upsample_prefix = f"up.{3-i}.upsample."
vae_conversion_map.append((sd_upsample_prefix, hf_upsample_prefix))
# up_blocks have three resnets
# also, up blocks in hf are numbered in reverse from sd
for j in range(3):
hf_up_prefix = f"decoder.up_blocks.{i}.resnets.{j}."
sd_up_prefix = f"decoder.up.{3-i}.block.{j}."
vae_conversion_map.append((sd_up_prefix, hf_up_prefix))
# this part accounts for mid blocks in both the encoder and the decoder
for i in range(2):
hf_mid_res_prefix = f"mid_block.resnets.{i}."
sd_mid_res_prefix = f"mid.block_{i+1}."
vae_conversion_map.append((sd_mid_res_prefix, hf_mid_res_prefix))
vae_conversion_map_attn = [
# (stable-diffusion, HF Diffusers)
("norm.", "group_norm."),
# the following are for SDXL
("q.", "to_q."),
("k.", "to_k."),
("v.", "to_v."),
("proj_out.", "to_out.0."),
]
def reshape_weight_for_sd(w):
# convert HF linear weights to SD conv2d weights
if not w.ndim == 1:
return w.reshape(*w.shape, 1, 1)
else:
return w
def convert_vae_state_dict(vae_state_dict):
mapping = {k: k for k in vae_state_dict.keys()}
for k, v in mapping.items():
for sd_part, hf_part in vae_conversion_map:
v = v.replace(hf_part, sd_part)
mapping[k] = v
for k, v in mapping.items():
if "attentions" in k:
for sd_part, hf_part in vae_conversion_map_attn:
v = v.replace(hf_part, sd_part)
mapping[k] = v
new_state_dict = {v: vae_state_dict[k] for k, v in mapping.items()}
weights_to_convert = ["q", "k", "v", "proj_out"]
for k, v in new_state_dict.items():
for weight_name in weights_to_convert:
if f"mid.attn_1.{weight_name}.weight" in k:
new_state_dict[k] = reshape_weight_for_sd(v)
return new_state_dict
textenc_conversion_lst = [
# (stable-diffusion, HF Diffusers)
("transformer.resblocks.", "text_model.encoder.layers."),
("ln_1", "layer_norm1"),
("ln_2", "layer_norm2"),
(".c_fc.", ".fc1."),
(".c_proj.", ".fc2."),
(".attn", ".self_attn"),
("ln_final.", "text_model.final_layer_norm."),
("token_embedding.weight", "text_model.embeddings.token_embedding.weight"),
("positional_embedding", "text_model.embeddings.position_embedding.weight"),
]
protected = {re.escape(x[1]): x[0] for x in textenc_conversion_lst}
textenc_pattern = re.compile("|".join(protected.keys()))
# Ordering is from https://github.com/pytorch/pytorch/blob/master/test/cpp/api/modules.cpp
code2idx = {"q": 0, "k": 1, "v": 2}
def convert_openclip_text_enc_state_dict(text_enc_dict):
new_state_dict = {}
capture_qkv_weight = {}
capture_qkv_bias = {}
for k, v in text_enc_dict.items():
if (
k.endswith(".self_attn.q_proj.weight")
or k.endswith(".self_attn.k_proj.weight")
or k.endswith(".self_attn.v_proj.weight")
):
k_pre = k[: -len(".q_proj.weight")]
k_code = k[-len("q_proj.weight")]
if k_pre not in capture_qkv_weight:
capture_qkv_weight[k_pre] = [None, None, None]
capture_qkv_weight[k_pre][code2idx[k_code]] = v
continue
if (
k.endswith(".self_attn.q_proj.bias")
or k.endswith(".self_attn.k_proj.bias")
or k.endswith(".self_attn.v_proj.bias")
):
k_pre = k[: -len(".q_proj.bias")]
k_code = k[-len("q_proj.bias")]
if k_pre not in capture_qkv_bias:
capture_qkv_bias[k_pre] = [None, None, None]
capture_qkv_bias[k_pre][code2idx[k_code]] = v
continue
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k)
new_state_dict[relabelled_key] = v
for k_pre, tensors in capture_qkv_weight.items():
if None in tensors:
raise RuntimeError("CORRUPTED MODEL: one of the q-k-v values for the text encoder was missing")
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k_pre)
new_state_dict[relabelled_key + ".in_proj_weight"] = torch.cat(tensors)
for k_pre, tensors in capture_qkv_bias.items():
if None in tensors:
raise RuntimeError("CORRUPTED MODEL: one of the q-k-v values for the text encoder was missing")
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k_pre)
new_state_dict[relabelled_key + ".in_proj_bias"] = torch.cat(tensors)
return new_state_dict
def convert_openai_text_enc_state_dict(text_enc_dict):
return text_enc_dict
def calculate_model_hash(state_dict):
func = hashlib.sha256()
for module in state_dict.values():
buffer = io.BytesIO()
torch.save(module, buffer)
func.update(buffer.getvalue())
return func.hexdigest()
def convert(model_path:str, checkpoint_path:str, metadata:dict={}):
unet_path = os.path.join(model_path, "unet", "diffusion_pytorch_model.safetensors")
vae_path = os.path.join(model_path, "vae", "diffusion_pytorch_model.safetensors")
text_enc_path = os.path.join(model_path, "text_encoder", "model.safetensors")
text_enc_2_path = os.path.join(model_path, "text_encoder_2", "model.safetensors")
unet_state_dict = load_file(unet_path, device="cpu")
vae_state_dict = load_file(vae_path, device="cpu")
text_enc_dict = load_file(text_enc_path, device="cpu")
text_enc_2_dict = load_file(text_enc_2_path, device="cpu")
unet_state_dict = convert_unet_state_dict(unet_state_dict)
unet_state_dict = {"model.diffusion_model." + k: v for k, v in unet_state_dict.items()}
vae_state_dict = convert_vae_state_dict(vae_state_dict)
vae_state_dict = {"first_stage_model." + k: v for k, v in vae_state_dict.items()}
text_enc_dict = convert_openai_text_enc_state_dict(text_enc_dict)
text_enc_dict = {"conditioner.embedders.0.transformer." + k: v for k, v in text_enc_dict.items()}
text_enc_2_dict = convert_openclip_text_enc_state_dict(text_enc_2_dict)
text_enc_2_dict = {"conditioner.embedders.1.model." + k: v for k, v in text_enc_2_dict.items()}
text_enc_2_dict["conditioner.embedders.1.model.text_projection"] = text_enc_2_dict.pop("conditioner.embedders.1.model.text_projection.weight").T.contiguous()
state_dict = {
**unet_state_dict,
**vae_state_dict,
**text_enc_dict,
**text_enc_2_dict
}
if metadata.get('modelspec.hash_sha256', None) is not None:
metadata['modelspec.hash_sha256'] = calculate_model_hash(state_dict)
save_file(state_dict, checkpoint_path, metadata=metadata)
return metadata
+310
View File
@@ -0,0 +1,310 @@
import io
import os
import json
import base64
from datetime import datetime
from PIL import Image
import torch
from safetensors.torch import load_file
import diffusers
import transformers
from modules import shared, devices
class Recipe:
author = ''
name = ''
version = ''
desc = ''
hint = ''
license = ''
prediction = ''
thumbnail = None
base = None
unet = None
vae = None
te1 = None
te2 = None
scheduler = 'UniPCMultistepScheduler'
dtype = torch.float16
diffusers = True
safetensors = True
debug = False
lora = {
}
fuse = 1.0
class Test:
generate = True
prompt = 'astronaut in a diner drinking coffee with burger and french fries on the table'
negative = 'ugly, blurry'
width = 1024
height = 1024
guidance = 4
steps = 20
recipe = Recipe()
test = Test()
pipeline: diffusers.StableDiffusionXLPipeline = None
status = ''
def msg(text, err:bool=False):
global status # pylint: disable=global-statement
if err:
shared.log.error(f'Modules merge: {text}')
else:
shared.log.info(f'Modules merge: {text}')
status += text + '<br>'
return status
def load_base(override:str=None):
global pipeline # pylint: disable=global-statement
fn = override or recipe.base
yield msg(f'base={fn}')
if os.path.isfile(fn):
pipeline = diffusers.StableDiffusionXLPipeline.from_single_file(fn, cache_dir=shared.opts.hfcache_dir, torch_dtype=recipe.dtype, add_watermarker=False)
elif os.path.isdir(fn):
pipeline = diffusers.StableDiffusionXLPipeline.from_pretrained(fn, cache_dir=shared.opts.hfcache_dir, torch_dtype=recipe.dtype, add_watermarker=False)
else:
yield msg('base: not found')
return None
pipeline.vae.register_to_config(force_upcast = False)
def load_unet(pipe: diffusers.StableDiffusionXLPipeline, override:str=None):
if (recipe.unet is None or len(recipe.unet) == 0) and override is None:
return None
fn = override or recipe.unet
if not os.path.isabs(fn):
fn = os.path.join(shared.opts.unet_dir, fn)
if not fn.endswith('.safetensors'):
fn += '.safetensors'
yield msg(f'unet={fn}')
if recipe.debug:
yield msg(f'config={pipe.unet.config}')
try:
unet = diffusers.UNet2DConditionModel.from_config(pipe.unet.config).to(recipe.dtype)
state_dict = load_file(fn)
unet.load_state_dict(state_dict)
pipe.unet = unet.to(device=devices.device, dtype=recipe.dtype)
except Exception as e:
yield msg(f'unet: {e}')
def load_scheduler(pipe: diffusers.StableDiffusionXLPipeline, override:str=None):
if recipe.scheduler is None and override is None:
return None
config = pipe.scheduler.config.__dict__
scheduler = override or recipe.scheduler
yield msg(f'scheduler={scheduler}')
if recipe.debug:
yield msg(f'config={config}')
try:
pipe.scheduler = getattr(diffusers, scheduler).from_config(config)
except Exception as e:
yield msg(f'scheduler: {e}')
def load_vae(pipe: diffusers.StableDiffusionXLPipeline, override:str=None):
if (recipe.vae is None or len(recipe.vae) == 0)and override is None:
return None
fn = override or recipe.vae
if not os.path.isabs(fn):
fn = os.path.join(shared.opts.vae_dir, fn)
if not fn.endswith('.safetensors'):
fn += '.safetensors'
try:
vae = diffusers.AutoencoderKL.from_single_file(fn, cache_dir=shared.opts.hfcache_dir, torch_dtype=recipe.dtype)
vae.config.force_upcast = False
vae.config.scaling_factor = 0.13025
vae.config.sample_size = 1024
yield msg(f'vae={fn}')
if recipe.debug:
yield msg(f'config={pipe.vae.config}')
pipe.vae = vae.to(device=devices.device, dtype=recipe.dtype)
except Exception as e:
yield msg(f'vae: {e}')
def load_te1(pipe: diffusers.StableDiffusionXLPipeline, override:str=None):
if (recipe.te1 is None or len(recipe.te1) == 0) and override is None:
return None
config = pipe.text_encoder.config.__dict__
pretrained_config = transformers.PretrainedConfig.from_dict(config)
fn = override or recipe.te1
if not os.path.isabs(fn):
fn = os.path.join(shared.opts.te_dir, fn)
if not fn.endswith('.safetensors'):
fn += '.safetensors'
yield msg(f'te1={fn}')
if recipe.debug:
yield msg(f'config={config}')
try:
state_dict = load_file(fn)
te1 = transformers.CLIPTextModel.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=pretrained_config, cache_dir=shared.opts.hfcache_dir)
pipe.text_encoder = te1.to(device=devices.device, dtype=recipe.dtype)
except Exception as e:
yield msg(f'te1: {e}')
def load_te2(pipe: diffusers.StableDiffusionXLPipeline, override:str=None):
if (recipe.te2 is None or len(recipe.te2) == 0) and override is None:
return None
config = pipe.text_encoder_2.config.__dict__
pretrained_config = transformers.PretrainedConfig.from_dict(config)
fn = override or recipe.te2
if not os.path.isabs(fn):
fn = os.path.join(shared.opts.te_dir, fn)
if not fn.endswith('.safetensors'):
fn += '.safetensors'
yield msg(f'te2={recipe.te2}')
if recipe.debug:
yield msg(f'config={config}')
try:
state_dict = load_file(fn)
te2 = transformers.CLIPTextModelWithProjection.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=pretrained_config, cache_dir=shared.opts.hfcache_dir)
pipe.text_encoder_2 = te2.to(device=devices.device, dtype=recipe.dtype)
except Exception as e:
yield msg(f'te2: {e}')
def load_lora(pipe: diffusers.StableDiffusionXLPipeline, override: dict=None, fuse: float=None):
if recipe.lora is None and override is None:
return None
names = []
pipe.unfuse_lora()
pipe.unload_lora_weights()
loras = override or recipe.lora
for lora, weight in loras.items():
try:
fn = lora
if not os.path.isabs(fn):
fn = os.path.join(shared.opts.lora_dir, fn)
if not fn.endswith('.safetensors'):
fn += '.safetensors'
yield msg(f'lora={fn} weight={weight} fuse={fuse or recipe.fuse}')
name = os.path.splitext(os.path.basename(lora))[0].replace('.', '').replace(' ', '').replace('-', '').replace('_', '')
names.append(name)
pipe.load_lora_weights(fn, name)
except Exception as e:
yield msg(f'lora: {e}')
if len(names) > 0:
pipe.set_adapters(adapter_names=names, adapter_weights=list(loras.values()))
pipe.fuse_lora(adapter_names=names, lora_scale=fuse or recipe.fuse, components=["unet", "text_encoder", "text_encoder_2"])
pipe.unload_lora_weights()
def test_model(pipe: diffusers.StableDiffusionXLPipeline, fn: str, **kwargs):
if not test.generate:
return
try:
generator = torch.Generator(devices.device).manual_seed(int(4242))
args = {
'prompt': test.prompt,
'negative_prompt': test.negative,
'num_inference_steps': test.steps,
'width': test.width,
'height': test.height,
'guidance_scale': test.guidance,
'generator': generator,
}
args.update(kwargs)
yield msg(f'test={args}')
image = pipe(**args).images[0]
yield msg(f'image={fn} {image}')
image.save(fn)
except Exception as e:
yield msg(f'test: {e}')
def get_thumbnail():
if recipe.thumbnail is None:
return ''
image = Image.open(recipe.thumbnail)
image = image.convert('RGB')
image.thumbnail((512, 512), resample=Image.Resampling.LANCZOS)
buffer = io.BytesIO()
image.save(buffer, format="JPEG", quality=50)
b64encoded = base64.b64encode(buffer.getvalue()).decode("utf-8")
return f'data:image/jpeg;base64,{b64encoded}'
def get_metadata():
return {
"modelspec.sai_model_spec": "1.0.0",
"modelspec.architecture": "stable-diffusion-xl-v1-base",
"modelspec.implementation": "diffusers",
"modelspec.title": recipe.name,
"modelspec.version": recipe.version,
"modelspec.description": recipe.desc,
"modelspec.author": recipe.author,
"modelspec.date": datetime.now().isoformat(timespec='minutes'),
"modelspec.license": recipe.license,
"modelspec.usage_hint": recipe.hint,
"modelspec.prediction_type": recipe.prediction,
"modelspec.dtype": str(recipe.dtype).split('.')[1],
"modelspec.hash_sha256": "",
"modelspec.thumbnail": get_thumbnail(),
"recipe": json.dumps({
"base": os.path.basename(recipe.base) if recipe.base else "default",
"unet": os.path.basename(recipe.unet) if recipe.unet else "default",
"vae": os.path.basename(recipe.vae) if recipe.vae else "default",
"te1": os.path.basename(recipe.te1) if recipe.te1 else "default",
"te2": os.path.basename(recipe.te2) if recipe.te2 else "default",
"scheduler": recipe.scheduler or "default",
"lora": [f'{os.path.basename(k)}:{v}' for k, v in recipe.lora.items()],
}),
}
def save_model(pipe: diffusers.StableDiffusionXLPipeline):
author = recipe.author if len(recipe.author) > 0 else 'anonymous'
folder = os.path.join(shared.opts.diffusers_dir, f'models--{author}--{recipe.name}')
if len(recipe.version) > 0:
folder += f'-{recipe.version}'
if not recipe.diffusers or recipe.safetensors:
return
try:
yield msg('save')
yield msg(f'pretrained={folder}')
pipe.save_pretrained(folder, safe_serialization=True, push_to_hub=False)
with open(os.path.join(folder, 'vae', 'config.json'), 'r', encoding='utf8') as f:
vae_config = json.load(f)
vae_config['force_upcast'] = False
vae_config['scaling_factor'] = 0.13025
vae_config['sample_size'] = 1024
with open(os.path.join(folder, 'vae', 'config.json'), 'w', encoding='utf8') as f:
json.dump(vae_config, f, indent=2)
if recipe.safetensors:
fn = recipe.name
if len(recipe.version) > 0:
fn += f'-{recipe.version}'
if not os.path.isabs(fn):
fn = os.path.join(shared.opts.ckpt_dir, fn)
if not fn.endswith('.safetensors'):
fn += '.safetensors'
yield msg(f'safetensors={fn}')
from modules.merging import convert_sdxl
metadata = convert_sdxl(model_path=folder, checkpoint_path=fn, metadata=get_metadata())
if 'modelspec.thumbnail' in metadata:
metadata['modelspec.thumbnail'] = f"{metadata['modelspec.thumbnail'].split(',')[0]}:{len(metadata['modelspec.thumbnail'])}"
yield msg(f'metadata={metadata}')
except Exception as e:
yield msg(f'save: {e}')
def merge():
global pipeline # pylint: disable=global-statement
yield from load_base()
if pipeline is None:
return
pipeline = pipeline.to(device=devices.device, dtype=recipe.dtype)
yield from load_scheduler(pipeline)
yield from load_unet(pipeline)
yield from load_vae(pipeline)
yield from load_te1(pipeline)
yield from load_te2(pipeline)
yield from load_lora(pipeline)
yield from save_model(pipeline)
# pipeline = pipeline.to(device=devices.device, dtype=recipe.dtype)
# test_model(pipeline, '/tmp/merge.png')
+7 -2
View File
@@ -52,7 +52,12 @@ class CheckpointInfo:
relname, ext = os.path.splitext(relname)
ext = ext.lower()[1:]
if os.path.isfile(filename): # ckpt or safetensor
if filename.lower() == 'none':
self.name = 'none'
self.relname = 'none'
self.sha256 = None
self.type = 'unknown'
elif os.path.isfile(filename): # ckpt or safetensor
self.name = relname
self.filename = filename
self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{relname}")
@@ -173,7 +178,7 @@ def update_model_hashes():
return txt
def get_closet_checkpoint_match(s: str):
def get_closet_checkpoint_match(s: str) -> CheckpointInfo:
if s.startswith('https://huggingface.co/'):
model_name = s.replace('https://huggingface.co/', '')
checkpoint_info = CheckpointInfo(model_name) # create a virutal model info
+18 -14
View File
@@ -437,6 +437,23 @@ def load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_con
return sd_model
def set_defaults(sd_model, checkpoint_info):
sd_model.sd_model_hash = checkpoint_info.calculate_shorthash() # pylint: disable=attribute-defined-outside-init
sd_model.sd_checkpoint_info = checkpoint_info # pylint: disable=attribute-defined-outside-init
sd_model.sd_model_checkpoint = checkpoint_info.filename # pylint: disable=attribute-defined-outside-init
if hasattr(sd_model, "prior_pipe"):
sd_model.default_scheduler = copy.deepcopy(sd_model.prior_pipe.scheduler) if hasattr(sd_model.prior_pipe, "scheduler") else None
else:
sd_model.default_scheduler = copy.deepcopy(sd_model.scheduler) if hasattr(sd_model, "scheduler") else None
sd_model.is_sdxl = False # a1111 compatibility item
sd_model.is_sd2 = hasattr(sd_model, 'cond_stage_model') and hasattr(sd_model.cond_stage_model, 'model') # a1111 compatibility item
sd_model.is_sd1 = not sd_model.is_sd2 # a1111 compatibility item
sd_model.logvar = sd_model.logvar.to(devices.device) if hasattr(sd_model, 'logvar') else None # fix for training
shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256
if hasattr(sd_model, "set_progress_bar_config"):
sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=80, colour='#327fba')
def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model', revision=None): # pylint: disable=unused-argument
if timer is None:
timer = Timer()
@@ -515,20 +532,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
shared.log.error(f'Load {op}: name="{checkpoint_info.name if checkpoint_info is not None else None}" not loaded')
return
sd_model.sd_model_hash = checkpoint_info.calculate_shorthash() # pylint: disable=attribute-defined-outside-init
sd_model.sd_checkpoint_info = checkpoint_info # pylint: disable=attribute-defined-outside-init
sd_model.sd_model_checkpoint = checkpoint_info.filename # pylint: disable=attribute-defined-outside-init
if hasattr(sd_model, "prior_pipe"):
sd_model.default_scheduler = copy.deepcopy(sd_model.prior_pipe.scheduler) if hasattr(sd_model.prior_pipe, "scheduler") else None
else:
sd_model.default_scheduler = copy.deepcopy(sd_model.scheduler) if hasattr(sd_model, "scheduler") else None
sd_model.is_sdxl = False # a1111 compatibility item
sd_model.is_sd2 = hasattr(sd_model, 'cond_stage_model') and hasattr(sd_model.cond_stage_model, 'model') # a1111 compatibility item
sd_model.is_sd1 = not sd_model.is_sd2 # a1111 compatibility item
sd_model.logvar = sd_model.logvar.to(devices.device) if hasattr(sd_model, 'logvar') else None # fix for training
shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256
if hasattr(sd_model, "set_progress_bar_config"):
sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=80, colour='#327fba')
set_defaults(sd_model, checkpoint_info)
if "Kandinsky" in sd_model.__class__.__name__: # need a special case
sd_model.scheduler.name = 'DDIM'
+21 -19
View File
@@ -7,7 +7,6 @@ from modules.sd_samplers_common import samples_to_image_grid, sample_to_image #
debug = shared.log.trace if os.environ.get('SD_SAMPLER_DEBUG', None) is not None else lambda *args, **kwargs: None
debug('Trace: SAMPLER')
all_samplers = []
all_samplers = []
all_samplers_map = {}
samplers = all_samplers
samplers_for_img2img = all_samplers
@@ -49,7 +48,7 @@ def visible_sampler_names():
def create_sampler(name, model):
if name is None or name == 'None':
return model.scheduler
return model.scheduler if model is not None else None
try:
current = model.scheduler.__class__.__name__
except Exception:
@@ -86,28 +85,31 @@ def create_sampler(name, model):
if not any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' in name:
shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} flow-match scheduler unsupported')
return None
# if any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' not in name:
# shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} linear scheduler unsupported')
# return None
sampler = config.constructor(model)
if sampler is None:
sampler = config.constructor(model)
if sampler is None or sampler.sampler is None:
model.scheduler = copy.deepcopy(model.default_scheduler)
if model is not None:
if sampler is None or sampler.sampler is None:
model.scheduler = copy.deepcopy(model.default_scheduler)
else:
model.scheduler = sampler.sampler
if not hasattr(model, 'scheduler_config'):
model.scheduler_config = sampler.sampler.config.copy() if hasattr(sampler, 'sampler') and hasattr(sampler.sampler, 'config') else {}
if hasattr(model, "prior_pipe") and hasattr(model.prior_pipe, "scheduler"):
model.prior_pipe.scheduler = sampler.sampler
model.prior_pipe.scheduler.config.clip_sample = False
if "flow" in model.scheduler.__class__.__name__.lower():
shared.state.prediction_type = "flow_prediction"
elif hasattr(model.scheduler, "config") and hasattr(model.scheduler.config, "prediction_type"):
shared.state.prediction_type = model.scheduler.config.prediction_type
if model is not None:
clean_config = {k: v for k, v in model.scheduler.config.items() if not k.startswith('_') and v is not None and v is not False}
cls = model.scheduler.__class__.__name__
else:
model.scheduler = sampler.sampler
if not hasattr(model, 'scheduler_config'):
model.scheduler_config = sampler.sampler.config.copy() if hasattr(sampler, 'sampler') and hasattr(sampler.sampler, 'config') else {}
if hasattr(model, "prior_pipe") and hasattr(model.prior_pipe, "scheduler"):
model.prior_pipe.scheduler = sampler.sampler
model.prior_pipe.scheduler.config.clip_sample = False
if "flow" in model.scheduler.__class__.__name__.lower():
shared.state.prediction_type = "flow_prediction"
elif hasattr(model.scheduler, "config") and hasattr(model.scheduler.config, "prediction_type"):
shared.state.prediction_type = model.scheduler.config.prediction_type
clean_config = {k: v for k, v in model.scheduler.config.items() if not k.startswith('_') and v is not None and v is not False}
clean_config = {k: v for k, v in sampler.sampler.config.items() if not k.startswith('_') and v is not None and v is not False}
cls = sampler.sampler.__class__.__name__
name = sampler.name if sampler is not None and sampler.sampler is not None else 'Default'
shared.log.debug(f'Sampler: "{name}" class={model.scheduler.__class__.__name__} config={clean_config}')
shared.log.debug(f'Sampler: "{name}" class={cls} config={clean_config}')
return sampler.sampler
else:
return None
+7 -10
View File
@@ -15,26 +15,20 @@ try:
CMStochasticIterativeScheduler,
UniPCMultistepScheduler,
DDIMScheduler,
EulerDiscreteScheduler,
EulerAncestralDiscreteScheduler,
EDMEulerScheduler,
FlowMatchEulerDiscreteScheduler,
DEISMultistepScheduler,
SASolverScheduler,
DPMSolverSinglestepScheduler,
DPMSolverMultistepScheduler,
EDMDPMSolverMultistepScheduler,
CosineDPMSolverMultistepScheduler,
DPMSolverSDEScheduler,
HeunDiscreteScheduler,
FlowMatchHeunDiscreteScheduler,
LCMScheduler,
PNDMScheduler,
IPNDMScheduler,
DDPMScheduler,
@@ -172,14 +166,17 @@ class DiffusionSampler:
return
self.name = name
self.config = {}
if not hasattr(model, 'scheduler'):
return
if getattr(model, "default_scheduler", None) is None: # sanity check
self.sampler = None
# if not hasattr(model, 'scheduler'):
# return
if getattr(model, "default_scheduler", None) is None and (model is not None): # sanity check
model.default_scheduler = copy.deepcopy(model.scheduler)
for key, value in config.get('All', {}).items(): # apply global defaults
self.config[key] = value
debug_log(f'Sampler: all="{self.config}"')
if hasattr(model.default_scheduler, 'scheduler_config'): # find model defaults
if model is None:
orig_config = {}
elif hasattr(model.default_scheduler, 'scheduler_config'): # find model defaults
orig_config = model.default_scheduler.scheduler_config
else:
orig_config = model.default_scheduler.config
+80 -55
View File
@@ -4,18 +4,16 @@ import json
import inspect
from datetime import datetime
import gradio as gr
from modules import sd_models, sd_vae, extras
from modules import errors, sd_models, sd_vae, extras, sd_samplers, ui_symbols, hashes
from modules.ui_components import ToolButton
from modules.ui_common import create_refresh_button
from modules.call_queue import wrap_gradio_gpu_call
from modules.shared import opts, log, req, readfile, max_workers, native
import modules.ui_symbols
import modules.errors
import modules.hashes
from modules.merging import merge_methods
from modules.merging.merge_utils import BETA_METHODS, TRIPLE_METHODS, interpolate
from modules.merging.merge_presets import BLOCK_WEIGHTS_PRESETS, SDXL_BLOCK_WEIGHTS_PRESETS
search_metadata_civit = None
extra_ui = []
@@ -32,9 +30,6 @@ def create_ui():
with gr.Column(elem_id='models_input_container', scale=3):
def gr_show(visible=True):
return {"visible": visible, "__type__": "update"}
with gr.Tab(label="Current"):
def analyze():
from modules import modelstats
@@ -57,45 +52,6 @@ def create_ui():
model_analyze.click(fn=analyze, inputs=[], outputs=[model_desc, model_modules, model_meta])
with gr.Tab(label="Convert"):
with gr.Row():
model_name = gr.Dropdown(sd_models.checkpoint_titles(), label="Original model")
create_refresh_button(model_name, sd_models.list_models, lambda: {"choices": sd_models.checkpoint_titles()}, "refresh_checkpoint_Z")
with gr.Row():
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", label="Model pruning methods")
with gr.Row():
checkpoint_formats = gr.CheckboxGroup(choices=["ckpt", "safetensors"], value=["safetensors"], label="Model Format")
with gr.Row():
show_extra_options = gr.Checkbox(label="Show extra options", value=False)
fix_clip = gr.Checkbox(label="Fix clip", value=False)
with gr.Row(visible=False) as extra_options:
specific_part_conv = ["copy", "convert", "delete"]
unet_conv = gr.Dropdown(specific_part_conv, value="convert", label="unet")
text_encoder_conv = gr.Dropdown(specific_part_conv, value="convert", label="text encoder")
vae_conv = gr.Dropdown(specific_part_conv, value="convert", label="vae")
others_conv = gr.Dropdown(specific_part_conv, value="convert", label="others")
show_extra_options.change(fn=lambda x: gr_show(x), inputs=[show_extra_options], outputs=[extra_options])
model_converter_convert = gr.Button(label="Convert", variant='primary')
model_converter_convert.click(
fn=extras.run_modelconvert,
inputs=[
model_name,
checkpoint_formats,
precision, m_type, custom_name,
unet_conv,
text_encoder_conv,
vae_conv,
others_conv,
fix_clip
],
outputs=[models_outcome]
)
with gr.Tab(label="Merge"):
def sd_model_choices():
return ['None'] + sd_models.checkpoint_titles()
@@ -222,7 +178,7 @@ def create_ui():
try:
results = extras.run_modelmerger(dummy_component, **kwargs)
except Exception as e:
modules.errors.display(e, 'Merge')
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_titles()) for _ in range(4)], f"Error merging checkpoints: {e}"]
return results
@@ -334,6 +290,76 @@ def create_ui():
]
)
with gr.Tab(label="Modules"):
with gr.Row():
with gr.Column(scale=3):
model_type = gr.Dropdown(label="Model type", choices=['sd15', 'sdxl', 'sd21', 'sd35', 'flux.1'], value='sdxl', interactive=False)
with gr.Column(scale=5):
with gr.Row():
model_name = gr.Dropdown(sd_models.checkpoint_titles(), label="Input model")
create_refresh_button(model_name, sd_models.list_models, lambda: {"choices": sd_models.checkpoint_titles()}, "refresh_checkpoint_Z")
with gr.Column(scale=5):
custom_name = gr.Textbox(label="Output model", placeholder="Output model path")
with gr.Row():
with gr.Column(scale=3):
gr.HTML('Model components<br><span style="color: var(--body-text-color-subdued)">Specify the components to include<br>Paths can be relative or absolute</span><br>')
with gr.Column(scale=5):
comp_unet = gr.Textbox(placeholder="UNet model", show_label=False)
comp_vae = gr.Textbox(placeholder="VAE model", show_label=False)
with gr.Column(scale=5):
comp_te1 = gr.Textbox(placeholder="Text encoder 1", show_label=False)
comp_te2 = gr.Textbox(placeholder="Text encoder 2", show_label=False)
with gr.Row():
with gr.Column(scale=3):
gr.HTML('Model settings<br>')
with gr.Column(scale=10):
with gr.Row():
precision = gr.Dropdown(label="Model precision", choices=["fp32", "fp16", "bf16"], value="fp16")
comp_scheduler = gr.Dropdown(label="Sampler", choices=[s.name for s in sd_samplers.samplers if s.constructor is not None])
comp_prediction = gr.Dropdown(Label="Prediction type", choices=["epsilon", "v"], value="epsilon")
with gr.Row():
with gr.Column(scale=3):
gr.HTML('Merge LoRA<br>')
with gr.Column(scale=9):
comp_lora = gr.Textbox(label="Comma separated list with optional strength per LoRA", placeholder="LoRA models")
with gr.Column(scale=1):
comp_fuse = gr.Number(label="Fuse strength", value=1.0)
with gr.Row():
gr.HTML('<br>')
with gr.Row():
with gr.Column(scale=2):
gr.HTML('Model metadata<br>')
with gr.Column(scale=5):
meta_author = gr.Textbox(placeholder="Author name", show_label=False)
meta_version = gr.Textbox(placeholder="Model version", show_label=False)
meta_license = gr.Textbox(placeholder="Model license", show_label=False)
with gr.Column(scale=5):
meta_desc = gr.Textbox(placeholder="Model description", lines=3, show_label=False)
meta_hint = gr.Textbox(placeholder="Model hint", lines=3, show_label=False)
with gr.Column(scale=3):
meta_thumbnail = gr.Image(label="Thumbnail", type='pil', source='upload')
with gr.Row():
gr.HTML('Note: Save is optional as you can merge in-memory and use newly created model immediately')
with gr.Row():
create_diffusers = gr.Checkbox(label="Save diffusers", value=True)
create_safetensors = gr.Checkbox(label="Save safetensors", value=True)
debug = gr.Checkbox(label="Debug info", value=False)
model_modules_btn = gr.Button(label="Modules", variant='primary')
model_modules_btn.click(
fn=extras.run_model_modules,
inputs=[
model_type, model_name, custom_name,
comp_unet, comp_vae, comp_te1, comp_te2,
precision, comp_scheduler, comp_prediction,
comp_lora, comp_fuse,
meta_author, meta_version, meta_license, meta_desc, meta_hint, meta_thumbnail,
create_diffusers, create_safetensors, debug,
],
outputs=[models_outcome]
)
with gr.Tab(label="Validate"):
model_headers = ['name', 'type', 'filename', 'hash', 'added', 'size', 'metadata']
model_data = []
@@ -407,7 +433,7 @@ def create_ui():
gr.HTML('<h2>Search for models</h2>Select a model from the search results to download<br><br>')
with gr.Row():
hf_search_text = gr.Textbox('', label='Search models', placeholder='search huggingface models')
hf_search_btn = ToolButton(value=modules.ui_symbols.search, label="Search")
hf_search_btn = ToolButton(value=ui_symbols.search, label="Search")
with gr.Row():
with gr.Column(scale=2):
with gr.Row():
@@ -562,7 +588,7 @@ def create_ui():
found = True
break
if not found and rehash and os.stat(item['filename']).st_size < (1024 * 1024 * 1024):
sha = modules.hashes.calculate_sha256(item['filename'], quiet=True)[:10]
sha = hashes.calculate_sha256(item['filename'], quiet=True)[:10]
r = req(f'https://civitai.com/api/v1/model-versions/by-hash/{sha}')
log.debug(f'CivitAI search: name="{item["name"]}" hash={sha} status={r.status_code}')
if r.status_code == 200:
@@ -622,7 +648,7 @@ def create_ui():
with gr.Row():
civit_search_text = gr.Textbox('', label='Search models', placeholder='keyword')
civit_search_tag = gr.Textbox('', label='', placeholder='tags')
civit_search_btn = ToolButton(value=modules.ui_symbols.search, label="Search", interactive=True)
civit_search_btn = ToolButton(value=ui_symbols.search, label="Search", interactive=True)
with gr.Row():
civit_search_res = gr.HTML('')
with gr.Row():
@@ -718,13 +744,12 @@ def create_ui():
def civit_update_metadata():
nonlocal update_data
log.debug('CivitAI update metadata: models')
from modules.ui_extra_networks import get_pages
from modules.modelloader import download_civit_meta
from modules import ui_extra_networks, modelloader
res = []
pages = get_pages('Model')
pages = ui_extra_networks.get_pages('Model')
if len(pages) == 0:
return 'CivitAI update metadata: no models found'
page: modules.ui_extra_networks.ExtraNetworksPage = pages[0]
page: ui_extra_networks.ExtraNetworksPage = pages[0]
table_data = []
update_data.clear()
all_hashes = [(item.get('hash', None) or 'XXXXXXXX').upper()[:8] for item in page.list_items()]
@@ -738,7 +763,7 @@ def create_ui():
if r.status_code == 200:
d = r.json()
model.id = d['modelId']
download_civit_meta(model.fn, model.id)
modelloader.download_civit_meta(model.fn, model.id)
fn = os.path.splitext(item['filename'])[0] + '.json'
model.meta = readfile(fn, silent=True)
model.name = model.meta.get('name', model.name)