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