diff --git a/src/diffusers/loaders/lora_base.py b/src/diffusers/loaders/lora_base.py index d4c88d35924f..e39e96a6159d 100644 --- a/src/diffusers/loaders/lora_base.py +++ b/src/diffusers/loaders/lora_base.py @@ -669,11 +669,19 @@ def unfuse_lora(self, components: list[str] | None = None, **kwargs): if issubclass(model.__class__, (ModelMixin, PreTrainedModel)): for module in model.modules(): if isinstance(module, BaseTunerLayer): - for adapter in set(module.merged_adapters): - if adapter and adapter in self._merged_adapters: - self._merged_adapters = self._merged_adapters - {adapter} module.unmerge() + # Only remove an adapter from _merged_adapters once it is no longer + # physically merged in any remaining loadable component. + remaining_merged: set[str] = set() + for component_name in self._lora_loadable_modules: + component_model = getattr(self, component_name, None) + if isinstance(component_model, nn.Module): + for module in component_model.modules(): + if isinstance(module, BaseTunerLayer): + remaining_merged.update(module.merged_adapters) + self._merged_adapters = self._merged_adapters & remaining_merged + def set_adapters( self, adapter_names: list[str] | str, diff --git a/tests/lora/test_lora_loader_utils.py b/tests/lora/test_lora_loader_utils.py index b35ffea80768..bc1f59ca92b3 100644 --- a/tests/lora/test_lora_loader_utils.py +++ b/tests/lora/test_lora_loader_utils.py @@ -17,9 +17,16 @@ import pytest import torch +import torch.nn as nn +from peft import LoraConfig +from peft.tuners.tuners_utils import BaseTunerLayer from safetensors.torch import save_file +from diffusers.configuration_utils import ConfigMixin from diffusers.loaders import StableDiffusionLoraLoaderMixin, lora_base +from diffusers.loaders.lora_base import LoraBaseMixin +from diffusers.loaders.peft import PeftAdapterMixin +from diffusers.models.modeling_utils import ModelMixin LORA_KEY = "unet.test.lora_A.weight" @@ -100,3 +107,46 @@ def test_local_directory_with_multiple_files_warns_and_uses_first(tmp_path, monk assert weight_name == first_path.name assert "contains more than one weights file" in caplog.text + + +def test_unfuse_lora_partial_components_keeps_merged_adapters_in_sync(): + """Regression test for gh-14214. + + Unfusing only a subset of components must keep _merged_adapters in sync + with the adapters still physically fused in the remaining components. + """ + + class TinyModel(ModelMixin, ConfigMixin, PeftAdapterMixin): + config_name = "config.json" + def __init__(self): + super().__init__() + self.linear = nn.Linear(8, 8) + + class FakePipeline(LoraBaseMixin): + _lora_loadable_modules = ["unet", "text_encoder"] + def __init__(self, unet, text_encoder): + self._merged_adapters = set() + self.unet, self.text_encoder = unet, text_encoder + + unet = TinyModel() + text_encoder = TinyModel() + config = LoraConfig(r=4, lora_alpha=4, target_modules=["linear"], init_lora_weights=False) + unet.add_adapter(config, adapter_name="adapter") + text_encoder.add_adapter(config, adapter_name="adapter") + + pipe = FakePipeline(unet, text_encoder) + pipe.fuse_lora(components=["unet", "text_encoder"], adapter_names=["adapter"]) + assert pipe.num_fused_loras == 1 + + pipe.unfuse_lora(components=["text_encoder"]) + assert "adapter" in pipe.fused_loras, "adapter should remain tracked while unet is still fused" + assert pipe.num_fused_loras == 1 + + unet_still_merged = any( + isinstance(m, BaseTunerLayer) and len(m.merged_adapters) > 0 + for m in unet.modules() + ) + assert unet_still_merged, "unet should still be physically merged at the PEFT level" + + pipe.unfuse_lora(components=["unet"]) + assert pipe.num_fused_loras == 0