Skip to content

Commit 88fa75f

Browse files
committed
fix(lora): only drop adapter from _merged_adapters when unfused from all components
Signed-off-by: Aloys Jehwin <aloysjehwin@gmail.com>
1 parent 6f2010e commit 88fa75f

1 file changed

Lines changed: 13 additions & 3 deletions

File tree

src/diffusers/loaders/lora_base.py

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -669,11 +669,21 @@ def unfuse_lora(self, components: list[str] | None = None, **kwargs):
669669
if issubclass(model.__class__, (ModelMixin, PreTrainedModel)):
670670
for module in model.modules():
671671
if isinstance(module, BaseTunerLayer):
672-
for adapter in set(module.merged_adapters):
673-
if adapter and adapter in self._merged_adapters:
674-
self._merged_adapters = self._merged_adapters - {adapter}
675672
module.unmerge()
676673

674+
# Only remove an adapter from _merged_adapters once it is no longer
675+
# physically merged in any remaining loadable component. Removing it
676+
# on the first unfused component would desync the set when the adapter
677+
# is still fused into other components.
678+
remaining_merged: set[str] = set()
679+
for component_name in self._lora_loadable_modules:
680+
component_model = getattr(self, component_name, None)
681+
if component_model is not None and issubclass(component_model.__class__, (ModelMixin, PreTrainedModel)):
682+
for module in component_model.modules():
683+
if isinstance(module, BaseTunerLayer):
684+
remaining_merged.update(module.merged_adapters)
685+
self._merged_adapters = self._merged_adapters & remaining_merged
686+
677687
def set_adapters(
678688
self,
679689
adapter_names: list[str] | str,

0 commit comments

Comments
 (0)