Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 29 additions & 1 deletion src/diffusers/utils/peft_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
):
Expand Down Expand Up @@ -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()

Expand Down
35 changes: 35 additions & 0 deletions tests/lora/test_peft_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Loading