From d14b54efa3c9dcf84d1fea96965f65911bdd5c97 Mon Sep 17 00:00:00 2001 From: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> Date: Sun, 27 Apr 2025 16:53:22 -0500 Subject: [PATCH] Preliminary support for sd1.5/sdxl OMI lora --- modules/lora/lora_convert.py | 23 +++++++++++------------ modules/lora/lora_load.py | 14 +++++++++++++- 2 files changed, 24 insertions(+), 13 deletions(-) diff --git a/modules/lora/lora_convert.py b/modules/lora/lora_convert.py index c2685aacd..aaef92e43 100644 --- a/modules/lora/lora_convert.py +++ b/modules/lora/lora_convert.py @@ -26,7 +26,7 @@ re_compiled = {} def make_unet_conversion_map() -> Dict[str, str]: unet_conversion_map_layer = [] - for i in range(3): # num_blocks is 3 in sdxl + for i in range(4): # num_blocks is 3 in sdxl # loop over downblocks/upblocks for j in range(2): # loop over resnets/attentions for downblocks @@ -108,7 +108,7 @@ def make_unet_conversion_map() -> Dict[str, str]: class KeyConvert: def __init__(self): self.is_sdxl = True if shared.sd_model_type == "sdxl" else False - self.UNET_CONVERSION_MAP = make_unet_conversion_map() if self.is_sdxl else None + self.UNET_CONVERSION_MAP = make_unet_conversion_map() self.LORA_PREFIX_UNET = "lora_unet_" self.LORA_PREFIX_TEXT_ENCODER = "lora_te_" self.OFT_PREFIX_UNET = "oft_unet_" @@ -117,16 +117,15 @@ class KeyConvert: self.LORA_PREFIX_TEXT_ENCODER2 = "lora_te2_" def __call__(self, key): - if self.is_sdxl: - if "diffusion_model" in key: # Fix NTC Slider naming error - key = key.replace("diffusion_model", "lora_unet") - map_keys = list(self.UNET_CONVERSION_MAP.keys()) # prefix of U-Net modules - map_keys.sort() - search_key = key.replace(self.LORA_PREFIX_UNET, "").replace(self.OFT_PREFIX_UNET, "").replace(self.LORA_PREFIX_TEXT_ENCODER1, "").replace(self.LORA_PREFIX_TEXT_ENCODER2, "") - position = bisect.bisect_right(map_keys, search_key) - map_key = map_keys[position - 1] - if search_key.startswith(map_key): - key = key.replace(map_key, self.UNET_CONVERSION_MAP[map_key]).replace("oft", "lora") # pylint: disable=unsubscriptable-object + if "diffusion_model" in key: # Fix NTC Slider naming error + key = key.replace("diffusion_model", "lora_unet") + map_keys = list(self.UNET_CONVERSION_MAP.keys()) # prefix of U-Net modules + map_keys.sort() + search_key = key.replace(self.LORA_PREFIX_UNET, "").replace(self.OFT_PREFIX_UNET, "").replace(self.LORA_PREFIX_TEXT_ENCODER1, "").replace(self.LORA_PREFIX_TEXT_ENCODER2, "") + position = bisect.bisect_right(map_keys, search_key) + map_key = map_keys[position - 1] + if search_key.startswith(map_key): + key = key.replace(map_key, self.UNET_CONVERSION_MAP[map_key]).replace("oft", "lora") # pylint: disable=unsubscriptable-object if "lycoris" in key and "transformer" in key: key = key.replace("lycoris", "lora_transformer") sd_module = shared.sd_model.network_layer_mapping.get(key, None) diff --git a/modules/lora/lora_load.py b/modules/lora/lora_load.py index 87b792fb2..84a0c01b7 100644 --- a/modules/lora/lora_load.py +++ b/modules/lora/lora_load.py @@ -109,7 +109,19 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: emb_dict[vec_name] = weight bundle_embeddings[emb_name] = emb_dict continue - if len(parts) > 5: # messy handler for diffusers peft lora + if parts[0] in ["clip_l","clip_g","t5","unet","transformer"]: + network_part = [] + while parts[-1] in ["alpha","weight","lora_up","lora_down"]: + network_part.insert(0,parts[-1]) + parts = parts[0:-1] + network_part = ".".join(network_part) + key_network_without_network_parts = "_".join(parts) + if key_network_without_network_parts.startswith("unet") or key_network_without_network_parts.startswith("transformer"): + key_network_without_network_parts = "lora_" + key_network_without_network_parts + key_network_without_network_parts = key_network_without_network_parts.replace("clip_g","lora_te2").replace("clip_l","lora_te") + #TODO Add t5 key support for SD3.5/f1 here? + + elif len(parts) > 5: # messy handler for diffusers peft lora key_network_without_network_parts = '_'.join(parts[:-2]) if not key_network_without_network_parts.startswith('lora_'): key_network_without_network_parts = 'lora_' + key_network_without_network_parts