mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
model component merge
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+12
-1
@@ -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
@@ -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; }
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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')
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user