SegMoE LoRA

This commit is contained in:
AI-Casanova
2024-02-05 22:23:25 -06:00
committed by Vladimir Mandic
parent 4bfa0a5650
commit fe93ea43a7
2 changed files with 28 additions and 5 deletions
+22
View File
@@ -139,6 +139,28 @@ class KeyConvert:
if sd_module is None:
key = key.replace("lora_te1_text_model", "transformer_text_model")
sd_module = shared.sd_model.network_layer_mapping.get(key, None)
# SegMoE begin
expert_key = key + "_experts_0"
expert_module = shared.sd_model.network_layer_mapping.get(expert_key, None)
if expert_module is not None:
sd_module = expert_module
key = expert_key
if sd_module is None:
key = key.replace("_net_", "_experts_0_net_")
sd_module = shared.sd_model.network_layer_mapping.get(key, None)
key = key if isinstance(key, list) else [key]
sd_module = sd_module if isinstance(sd_module, list) else [sd_module]
if "_experts_0" in key[0]:
i = expert_module = 1
while expert_module is not None:
expert_key = key[0].replace("_experts_0", f"_experts_{i}")
expert_module = shared.sd_model.network_layer_mapping.get(expert_key, None)
if expert_module is not None:
key.append(expert_key)
sd_module.append(expert_module)
i += 1
# SegMoe end
return key, sd_module
def diffusers(self, key):
+6 -5
View File
@@ -122,13 +122,14 @@ def load_network(name, network_on_disk) -> network.Network:
key_network_without_network_parts, network_part = key_network.split(".", 1)
# if debug:
# shared.log.debug(f'LoRA load: name="{name}" full={key_network} network={network_part} key={key_network_without_network_parts}')
key, sd_module = convert(key_network_without_network_parts)
if sd_module is None:
key, sd_module = convert(key_network_without_network_parts) # Now returns lists
if sd_module[0] is None:
keys_failed_to_match[key_network] = key
continue
if key not in matched_networks:
matched_networks[key] = network.NetworkWeights(network_key=key_network, sd_key=key, w={}, sd_module=sd_module)
matched_networks[key].w[network_part] = weight
for k, module in zip(key, sd_module):
if k not in matched_networks:
matched_networks[k] = network.NetworkWeights(network_key=key_network, sd_key=k, w={}, sd_module=module)
matched_networks[k].w[network_part] = weight
for key, weights in matched_networks.items():
net_module = None
for nettype in module_types: