diff --git a/modules/lora/lora_convert.py b/modules/lora/lora_convert.py index 8af280f98..48c1b09b3 100644 --- a/modules/lora/lora_convert.py +++ b/modules/lora/lora_convert.py @@ -130,6 +130,12 @@ class KeyConvert: sd_module = shared.sd_model.network_layer_mapping.get(key, None) if sd_module is None: sd_module = shared.sd_model.network_layer_mapping.get(key.replace("guidance", "timestep"), None) # FLUX1 fix + if sd_module is None and key.startswith("lora_te"): + # transformers >=5.6 flattened CLIPTextModel; kohya te keys still carry the text_model wrapper + flat_key = key.replace("_text_model_", "_", 1) + sd_module = shared.sd_model.network_layer_mapping.get(flat_key, None) + if sd_module is not None: + key = flat_key if debug and sd_module is None: raise RuntimeError(f"LoRA key not found in network_layer_mapping: key={key} mapping={shared.sd_model.network_layer_mapping.keys()}") return key, sd_module diff --git a/modules/lora/lora_extract.py b/modules/lora/lora_extract.py index 22c6019fc..3f25824e0 100644 --- a/modules/lora/lora_extract.py +++ b/modules/lora/lora_extract.py @@ -142,14 +142,17 @@ def make_lora(fn, maxrank, auto_rank, rank_ratio, modules, overwrite): if 'te' in modules and getattr(shared.sd_model, 'text_encoder', None) is not None: task = progress.add_task(description="te1 decompose", total=len(list(shared.sd_model.text_encoder.named_modules()))) + # transformers >=5.6 flattened CLIPTextModel; kohya naming keeps the text_model wrapper + flattened_clip = 'CLIPTextModel' in shared.sd_model.text_encoder.__class__.__name__ and not hasattr(shared.sd_model.text_encoder, 'text_model') for name, module in shared.sd_model.text_encoder.named_modules(): progress.update(task, advance=1) weights_backup = getattr(module, "network_weights_backup", None) if weights_backup is None or getattr(module, "network_current_names", None) is None: continue prefix = "lora_te1_" if hasattr(shared.sd_model, 'text_encoder_2') else "lora_te_" + key_name = f'text_model.{name}' if flattened_clip else name module.svdhandler = SVDHandler(maxrank, rank_ratio) - module.svdhandler.network_name = prefix + name.replace(".", "_") + module.svdhandler.network_name = prefix + key_name.replace(".", "_") with devices.inference_context(): module.svdhandler.decompose(module.weight, weights_backup) progress.remove_task(task)