diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 160487e88..b227f82f2 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -50,44 +50,45 @@ convert_diffusers_name_to_compvis = lora_convert.convert_diffusers_name_to_compv def assign_network_names_to_compvis_modules(sd_model): if sd_model is None: return + sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility network_layer_mapping = {} if shared.native: - if hasattr(shared.sd_model, 'text_encoder') and shared.sd_model.text_encoder is not None: - for name, module in shared.sd_model.text_encoder.named_modules(): - prefix = "lora_te1_" if hasattr(shared.sd_model, 'text_encoder_2') else "lora_te_" + if hasattr(sd_model, 'text_encoder') and sd_model.text_encoder is not None: + for name, module in sd_model.text_encoder.named_modules(): + prefix = "lora_te1_" if hasattr(sd_model, 'text_encoder_2') else "lora_te_" network_name = prefix + name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - if hasattr(shared.sd_model, 'text_encoder_2'): - for name, module in shared.sd_model.text_encoder_2.named_modules(): + if hasattr(sd_model, 'text_encoder_2'): + for name, module in sd_model.text_encoder_2.named_modules(): network_name = "lora_te2_" + name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - if hasattr(shared.sd_model, 'unet'): - for name, module in shared.sd_model.unet.named_modules(): + if hasattr(sd_model, 'unet'): + for name, module in sd_model.unet.named_modules(): network_name = "lora_unet_" + name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - if hasattr(shared.sd_model, 'transformer'): - for name, module in shared.sd_model.transformer.named_modules(): + if hasattr(sd_model, 'transformer'): + for name, module in sd_model.transformer.named_modules(): network_name = "lora_transformer_" + name.replace(".", "_") network_layer_mapping[network_name] = module if "norm" in network_name and "linear" not in network_name: continue module.network_layer_name = network_name else: - if not hasattr(shared.sd_model, 'cond_stage_model'): + if not hasattr(sd_model, 'cond_stage_model'): sd_model.network_layer_mapping = {} return - for name, module in shared.sd_model.cond_stage_model.wrapped.named_modules(): + for name, module in sd_model.cond_stage_model.wrapped.named_modules(): network_name = name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - for name, module in shared.sd_model.model.named_modules(): + for name, module in sd_model.model.named_modules(): network_name = name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - sd_model.network_layer_mapping = network_layer_mapping + shared.sd_model.network_layer_mapping = network_layer_mapping def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_default_multiplier) -> network.Network: