From 311d402b0ca8874d3e6645de85c8a30f313461f6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 19 Jan 2025 14:36:09 -0500 Subject: [PATCH] model component merge Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 13 +- javascript/base.css | 2 +- javascript/sdnext.css | 7 +- modules/extras.py | 216 ++++++++------------- modules/merging/convert_sdxl.py | 297 +++++++++++++++++++++++++++++ modules/merging/modules_sdxl.py | 310 +++++++++++++++++++++++++++++++ modules/sd_checkpoint.py | 9 +- modules/sd_models.py | 32 ++-- modules/sd_samplers.py | 40 ++-- modules/sd_samplers_diffusers.py | 17 +- modules/ui_models.py | 135 ++++++++------ 11 files changed, 833 insertions(+), 245 deletions(-) create mode 100644 modules/merging/convert_sdxl.py create mode 100644 modules/merging/modules_sdxl.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 56a2732b5..eabd560f5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/javascript/base.css b/javascript/base.css index 0a0b621c2..8f89685c2 100644 --- a/javascript/base.css +++ b/javascript/base.css @@ -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; } diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 1db6288a5..c3deebad7 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -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); } diff --git a/modules/extras.py b/modules/extras.py index 162491580..e4eb36639 100644 --- a/modules/extras.py +++ b/modules/extras.py @@ -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 + '
' + 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}
" + 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 diff --git a/modules/merging/convert_sdxl.py b/modules/merging/convert_sdxl.py new file mode 100644 index 000000000..93fc71f5d --- /dev/null +++ b/modules/merging/convert_sdxl.py @@ -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 diff --git a/modules/merging/modules_sdxl.py b/modules/merging/modules_sdxl.py new file mode 100644 index 000000000..61f2cf948 --- /dev/null +++ b/modules/merging/modules_sdxl.py @@ -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 + '
' + 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') diff --git a/modules/sd_checkpoint.py b/modules/sd_checkpoint.py index e5ddd2e85..c20fa00f4 100644 --- a/modules/sd_checkpoint.py +++ b/modules/sd_checkpoint.py @@ -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 diff --git a/modules/sd_models.py b/modules/sd_models.py index 9d0349a3f..6234c2b0b 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -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' diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index 398f608d1..accd6b0ed 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -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 diff --git a/modules/sd_samplers_diffusers.py b/modules/sd_samplers_diffusers.py index 6b2de72aa..a8b77bc90 100644 --- a/modules/sd_samplers_diffusers.py +++ b/modules/sd_samplers_diffusers.py @@ -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 diff --git a/modules/ui_models.py b/modules/ui_models.py index 5d5b452e2..23a39f317 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -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
Specify the components to include
Paths can be relative or absolute

') + 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
') + 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
') + 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('
') + with gr.Row(): + with gr.Column(scale=2): + gr.HTML('Model metadata
') + 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('

Search for models

Select a model from the search results to download

') 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)