Skip to content
Open
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
16 changes: 13 additions & 3 deletions src/diffusers/loaders/lora_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -669,11 +669,21 @@ 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. Removing it
# on the first unfused component would desync the set when the adapter
# is still fused into other components.
remaining_merged: set[str] = set()
for component_name in self._lora_loadable_modules:
component_model = getattr(self, component_name, None)
if component_model is not None and issubclass(component_model.__class__, (ModelMixin, PreTrainedModel)):
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,
Expand Down
Loading