Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 7.3k
PEFT Integration for Text Encoder to handle multiple alphas/ranks, disable/enable adapters and support for multiple adapters#5147
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Uh oh!
There was an error while loading. Please reload this page.
Changes from all commits
ba24f2ac17634c2a6e53501f6d1d5a150b2961e776cdbe739691368b791885114db139c06c40bd56a14dec87c1940a60280c62ef3c4295c94162ddf1d13f4078a860d78a01d5ecbc7149d650c96f1adcdf890906f8e87f63ba2d4edc83fa0b83fcba74e33a99cb8563c90f85d40a4894ea0595927e3da63d7c567d01a292e836b14cb48405325462db412adcb72ef23e072655bd46ae9724b52b5e6f34371650d4920333f0985d17ece3b0201a15cc080db75ffbac30916c31a5de0f1bc32872e0acb58c1ca4c627c377882fcf174a1f01287b2ccfffd9bcfe9916ac6File filter
Filter by extension
Conversations
Uh oh!
There was an error while loading. Please reload this page.
Jump to
Uh oh!
There was an error while loading. Please reload this page.
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -35,18 +35,23 @@ | ||
| convert_state_dict_to_diffusers, | ||
| convert_state_dict_to_peft, | ||
| deprecate, | ||
| get_adapter_name, | ||
| get_peft_kwargs, | ||
| is_accelerate_available, | ||
| is_omegaconf_available, | ||
| is_peft_available, | ||
| is_transformers_available, | ||
| logging, | ||
| recurse_remove_peft_layers, | ||
| scale_lora_layers, | ||
| set_adapter_layers, | ||
| set_weights_and_activate_adapters, | ||
| ) | ||
| from .utils.import_utils import BACKENDS_MAPPING | ||
| if is_transformers_available(): | ||
| from transformers import CLIPTextModel, CLIPTextModelWithProjection | ||
| from transformers import CLIPTextModel, CLIPTextModelWithProjection, PreTrainedModel | ||
| if is_accelerate_available(): | ||
| from accelerate import init_empty_weights | ||
| @@ -1100,7 +1105,9 @@ class LoraLoaderMixin: | ||
| num_fused_loras = 0 | ||
| use_peft_backend = USE_PEFT_BACKEND | ||
| def load_lora_weights(self, pretrained_model_name_or_path_or_dict: Union[str, Dict[str, torch.Tensor]], **kwargs): | ||
| def load_lora_weights( | ||
| self, pretrained_model_name_or_path_or_dict: Union[str, Dict[str, torch.Tensor]], adapter_name=None, **kwargs | ||
| ): | ||
| """ | ||
| Load LoRA weights specified in `pretrained_model_name_or_path_or_dict` into `self.unet` and | ||
| `self.text_encoder`. | ||
| @@ -1120,6 +1127,9 @@ def load_lora_weights(self, pretrained_model_name_or_path_or_dict: Union[str, Di | ||
| See [`~loaders.LoraLoaderMixin.lora_state_dict`]. | ||
| kwargs (`dict`, *optional*): | ||
| See [`~loaders.LoraLoaderMixin.lora_state_dict`]. | ||
| adapter_name (`str`, *optional*): | ||
| Adapter name to be used for referencing the loaded adapter model. If not specified, it will use | ||
| `default_{i}` where i is the total number of adapters being loaded. | ||
| """ | ||
| # First, ensure that the checkpoint is a compatible one and can be successfully loaded. | ||
| state_dict, network_alphas = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) | ||
| @@ -1143,6 +1153,7 @@ def load_lora_weights(self, pretrained_model_name_or_path_or_dict: Union[str, Di | ||
| text_encoder=self.text_encoder, | ||
| lora_scale=self.lora_scale, | ||
| low_cpu_mem_usage=low_cpu_mem_usage, | ||
| adapter_name=adapter_name, | ||
Member There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Let's keep the private variable at the end. | ||
| _pipeline=self, | ||
| ) | ||
| @@ -1500,6 +1511,7 @@ def load_lora_into_text_encoder( | ||
| prefix=None, | ||
| lora_scale=1.0, | ||
| low_cpu_mem_usage=None, | ||
| adapter_name=None, | ||
| _pipeline=None, | ||
| ): | ||
| """ | ||
| @@ -1523,6 +1535,9 @@ def load_lora_into_text_encoder( | ||
| tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model. | ||
| Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this | ||
| argument to `True` will raise an error. | ||
| adapter_name (`str`, *optional*): | ||
| Adapter name to be used for referencing the loaded adapter model. If not specified, it will use | ||
| `default_{i}` where i is the total number of adapters being loaded. | ||
| """ | ||
| low_cpu_mem_usage = low_cpu_mem_usage if low_cpu_mem_usage is not None else _LOW_CPU_MEM_USAGE_DEFAULT | ||
| @@ -1584,19 +1599,22 @@ def load_lora_into_text_encoder( | ||
| if cls.use_peft_backend: | ||
| from peft import LoraConfig | ||
| lora_rank = list(rank.values())[0] | ||
| # By definition, the scale should be alpha divided by rank. | ||
| # https://github.com/huggingface/peft/blob/ba0477f2985b1ba311b83459d29895c809404e99/src/peft/tuners/lora/layer.py#L71 | ||
| alpha = lora_scale * lora_rank | ||
| lora_config_kwargs = get_peft_kwargs(rank, network_alphas, text_encoder_lora_state_dict) | ||
| target_modules = ["q_proj", "k_proj", "v_proj", "out_proj"] | ||
| if patch_mlp: | ||
| target_modules += ["fc1", "fc2"] | ||
| lora_config = LoraConfig(**lora_config_kwargs) | ||
| # TODO: support multi alpha / rank: https://github.com/huggingface/peft/pull/873 | ||
| lora_config = LoraConfig(r=lora_rank, target_modules=target_modules, lora_alpha=alpha) | ||
| # adapter_name | ||
| if adapter_name is None: | ||
| adapter_name = get_adapter_name(text_encoder) | ||
| text_encoder.load_adapter(adapter_state_dict=text_encoder_lora_state_dict, peft_config=lora_config) | ||
| # inject LoRA layers and load the state dict | ||
| text_encoder.load_adapter( | ||
| adapter_name=adapter_name, | ||
| adapter_state_dict=text_encoder_lora_state_dict, | ||
| peft_config=lora_config, | ||
| ) | ||
| # scale LoRA layers with `lora_scale` | ||
| scale_lora_layers(text_encoder, weight=lora_scale) | ||
| is_model_cpu_offload = False | ||
| is_sequential_cpu_offload = False | ||
| @@ -2178,6 +2196,81 @@ def unfuse_text_encoder_lora(text_encoder): | ||
| self.num_fused_loras -= 1 | ||
| def set_adapter_for_text_encoder( | ||
sayakpaul marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| self, | ||
| adapter_names: Union[List[str], str], | ||
| text_encoder: Optional[PreTrainedModel] = None, | ||
| text_encoder_weights: List[float] = None, | ||
| ): | ||
pacman100 marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| """ | ||
| Sets the adapter layers for the text encoder. | ||
| Args: | ||
| adapter_names (`List[str]` or `str`): | ||
| The names of the adapters to use. | ||
| text_encoder (`torch.nn.Module`, *optional*): | ||
| The text encoder module to set the adapter layers for. If `None`, it will try to get the `text_encoder` | ||
| attribute. | ||
| text_encoder_weights (`List[float]`, *optional*): | ||
| The weights to use for the text encoder. If `None`, the weights are set to `1.0` for all the adapters. | ||
| """ | ||
| if not self.use_peft_backend: | ||
| raise ValueError("PEFT backend is required for this method.") | ||
| def process_weights(adapter_names, weights): | ||
| if weights is None: | ||
| weights = [1.0] * len(adapter_names) | ||
| elif isinstance(weights, float): | ||
| weights = [weights] | ||
| if len(adapter_names) != len(weights): | ||
| raise ValueError( | ||
| f"Length of adapter names {len(adapter_names)} is not equal to the length of the weights {len(weights)}" | ||
| ) | ||
| return weights | ||
| adapter_names = [adapter_names] if isinstance(adapter_names, str) else adapter_names | ||
| text_encoder_weights = process_weights(adapter_names, text_encoder_weights) | ||
| text_encoder = text_encoder or getattr(self, "text_encoder", None) | ||
| if text_encoder is None: | ||
| raise ValueError( | ||
| "The pipeline does not have a default `pipe.text_encoder` class. Please make sure to pass a `text_encoder` instead." | ||
| ) | ||
| set_weights_and_activate_adapters(text_encoder, adapter_names, text_encoder_weights) | ||
| def disable_lora_for_text_encoder(self, text_encoder: Optional[PreTrainedModel] = None): | ||
| """ | ||
| Disables the LoRA layers for the text encoder. | ||
| Args: | ||
| text_encoder (`torch.nn.Module`, *optional*): | ||
| The text encoder module to disable the LoRA layers for. If `None`, it will try to get the | ||
| `text_encoder` attribute. | ||
| """ | ||
| if not self.use_peft_backend: | ||
| raise ValueError("PEFT backend is required for this method.") | ||
| text_encoder = text_encoder or getattr(self, "text_encoder", None) | ||
| if text_encoder is None: | ||
| raise ValueError("Text Encoder not found.") | ||
| set_adapter_layers(text_encoder, enabled=False) | ||
| def enable_lora_for_text_encoder(self, text_encoder: Optional[PreTrainedModel] = None): | ||
| """ | ||
| Enables the LoRA layers for the text encoder. | ||
| Args: | ||
| text_encoder (`torch.nn.Module`, *optional*): | ||
| The text encoder module to enable the LoRA layers for. If `None`, it will try to get the `text_encoder` | ||
| attribute. | ||
| """ | ||
| if not self.use_peft_backend: | ||
| raise ValueError("PEFT backend is required for this method.") | ||
| text_encoder = text_encoder or getattr(self, "text_encoder", None) | ||
| if text_encoder is None: | ||
| raise ValueError("Text Encoder not found.") | ||
| set_adapter_layers(self.text_encoder, enabled=True) | ||
| class FromSingleFileMixin: | ||
| """ | ||
Uh oh!
There was an error while loading. Please reload this page.