diff --git a/src/diffusers/utils/peft_utils.py b/src/diffusers/utils/peft_utils.py index ea6f86798100..31ba150ac7ff 100644 --- a/src/diffusers/utils/peft_utils.py +++ b/src/diffusers/utils/peft_utils.py @@ -150,6 +150,34 @@ def unscale_lora_layers(model, weight: float | None = None): module.set_scale(adapter_name, 1.0) +# Maps the LoRA weight suffix of a text-encoder network alpha key (diffusers naming) +# to the corresponding PEFT module name suffix. PEFT's `alpha_pattern` keys must +# match the end of the PEFT module names (see `peft.utils.get_pattern_key`), e.g. +# the alpha key `text_model.encoder.layers.1.self_attn.to_k_lora.down.weight.alpha` +# belongs to the module `...self_attn.k_proj` -- not its parent block. +_TEXT_ENCODER_ALPHA_SUFFIX_TO_MODULE_SUFFIX = ( + (".to_q_lora.down.weight.alpha", ".q_proj"), + (".to_k_lora.down.weight.alpha", ".k_proj"), + (".to_v_lora.down.weight.alpha", ".v_proj"), + (".to_out_lora.down.weight.alpha", ".out_proj"), + (".lora_linear_layer.down.weight.alpha", ""), + # e.g. `text_model.text_projection.alpha` + (".alpha", ""), +) + + +def _text_encoder_alpha_pattern_key(alpha_key: str) -> str: + """Derive the PEFT `alpha_pattern` key from a text-encoder network alpha key.""" + for suffix, module_suffix in _TEXT_ENCODER_ALPHA_SUFFIX_TO_MODULE_SUFFIX: + if alpha_key.endswith(suffix): + key = alpha_key[: -len(suffix)] + module_suffix + # The loader strips the `text_model.` prefix for flattened text encoders, + # so the pattern key must not require it: `get_pattern_key` matches at the + # end of the module name, making the shorter key match either way. + return key.removeprefix("text_model.") + return alpha_key + + def get_peft_kwargs( rank_dict, network_alpha_dict, peft_state_dict, is_unet=True, model_state_dict=None, adapter_name=None ): @@ -186,7 +214,7 @@ def get_peft_kwargs( for k, v in alpha_pattern.items() } else: - alpha_pattern = {".".join(k.split(".down.")[0].split(".")[:-1]): v for k, v in alpha_pattern.items()} + alpha_pattern = {_text_encoder_alpha_pattern_key(k): v for k, v in alpha_pattern.items()} else: lora_alpha = set(network_alpha_dict.values()).pop() diff --git a/tests/lora/test_peft_utils.py b/tests/lora/test_peft_utils.py index 67b7e2695233..262c4fa41f93 100644 --- a/tests/lora/test_peft_utils.py +++ b/tests/lora/test_peft_utils.py @@ -77,3 +77,38 @@ def test_mixed_ranks_with_per_module_alphas_unchanged(): kwargs = get_peft_kwargs(_rank_dict(module_ranks), network_alphas, _peft_state_dict(module_ranks)) for module in module_ranks: assert _effective_scale(kwargs, module) == 1.0 + + +def test_text_encoder_alpha_pattern_names_lora_modules(): + # Regression test for https://github.com/huggingface/diffusers/issues/14970: + # text-encoder `alpha_pattern` keys must name the LoRA module itself (e.g. `k_proj`), + # not its parent block, otherwise PEFT's end-of-name matching misses and the + # default (most common) alpha is applied to every module. + peft_state_dict = { + "text_model.encoder.layers.0.self_attn.q_proj.lora_A.weight": None, + "text_model.encoder.layers.0.self_attn.q_proj.lora_B.weight": None, + "text_model.encoder.layers.1.self_attn.k_proj.lora_A.weight": None, + "text_model.encoder.layers.1.self_attn.k_proj.lora_B.weight": None, + } + rank_dict = {k.replace(".lora_A.", ".lora_B."): 4 for k in peft_state_dict if ".lora_A." in k} + network_alphas = { + "text_model.encoder.layers.0.self_attn.to_q_lora.down.weight.alpha": 4.0, + "text_model.encoder.layers.1.self_attn.to_k_lora.down.weight.alpha": 8.0, + } + kwargs = get_peft_kwargs(rank_dict, network_alphas, peft_state_dict, is_unet=False) + assert kwargs["lora_alpha"] == 4.0 + assert kwargs["alpha_pattern"] == {"encoder.layers.1.self_attn.k_proj": 8.0} + assert _effective_scale(kwargs, "encoder.layers.1.self_attn.k_proj") == 2.0 + assert _effective_scale(kwargs, "encoder.layers.0.self_attn.q_proj") == 1.0 + + +def test_text_encoder_alpha_pattern_handles_mlp_and_text_projection(): + peft_state_dict = {"m.lora_A.weight": None, "m.lora_B.weight": None} + network_alphas = { + "text_model.encoder.layers.0.mlp.fc1.lora_linear_layer.down.weight.alpha": 8.0, + "text_model.text_projection.alpha": 8.0, + "text_model.encoder.layers.0.mlp.fc2.lora_linear_layer.down.weight.alpha": 4.0, + } + kwargs = get_peft_kwargs({"m.lora_B.weight": 4}, network_alphas, peft_state_dict, is_unet=False) + assert kwargs["lora_alpha"] == 8.0 + assert kwargs["alpha_pattern"] == {"encoder.layers.0.mlp.fc2": 4.0}